diff --git a/.circleci/config.yml b/.circleci/config.yml index d8dc40433dc..e9c295aca08 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -1095,50 +1095,6 @@ jobs: - google_generate_content_endpoint_coverage.xml - google_generate_content_endpoint_coverage - llm_responses_api_testing: - docker: - - *python312_image - working_directory: ~/project - resource_class: large - environment: - REQUEST_TIMEOUT: "180" - - steps: - - checkout - - skip_if_unrelated_changes - - setup_google_dns - - install_uv - - install_rust - - restore_cache: - keys: - - v1-uv-cache-{{ checksum "uv.lock" }} - - run: - name: Install Dependencies - command: | - uv sync --frozen --all-groups --all-extras --python 3.12 - - save_cache: - paths: - - ~/.cache/uv - key: v1-uv-cache-{{ checksum "uv.lock" }} - # Run pytest and generate JUnit XML report - - run: - name: Run tests - command: | - mkdir -p test-results - TEST_FILES=$(circleci tests glob "tests/llm_responses_api_testing/**/test_*.py") - echo "$TEST_FILES" | circleci tests run \ - --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ - -v \ - --junitxml=test-results/junit.xml \ - --durations=5 \ - -n 8 \ - --reruns 1 --only-rerun Timeout" - no_output_timeout: 15m - - # Store test results - - store_test_results: - path: test-results ocr_testing: docker: - *python312_image @@ -2074,98 +2030,6 @@ jobs: # Store test results - store_test_results: path: test-results - proxy_spend_accuracy_tests: - machine: - image: ubuntu-2204:2024.04.1 - resource_class: large - working_directory: ~/project - steps: - - checkout - - run: - name: Generate LiteLLM master key - command: | - key="$(openssl rand -hex 16)" - printf 'export LITELLM_MASTER_KEY=sk-%s\n' "$key" >> "$BASH_ENV" - - skip_if_unrelated_changes - - setup_google_dns - - install_uv - - install_rust - - run: - name: Install Dependencies - command: | - uv sync --frozen --all-groups --all-extras --python 3.12 - - start_postgres - - start_redis - - start_fake_openai_endpoint - - attach_workspace: - at: ~/project - - run: - name: Load Docker Database Image - command: | - zstd -d litellm-docker-database.tar.zst --stdout | docker load - docker images | grep litellm-docker-database - - run: - name: Run Docker container - # Point the proxy at the job-local Redis (start_redis) instead of the - # shared remote Redis. The Redis transaction buffer uses a single - # global pod-lock key (cronjob_lock:db_spend_update_job) and a single - # global buffer list (litellm_spend_update_buffer); sharing those - # across concurrent CI pipelines causes spend flushes to stall or - # land in the wrong DB, which is what makes this test flaky. - command: | - docker run -d \ - -p 4000:4000 \ - -e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \ - -e REDIS_HOST=host.docker.internal \ - -e REDIS_PORT=6379 \ - -e LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" \ - -e OPENAI_API_KEY=$OPENAI_API_KEY \ - -e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \ - -e LITELLM_LICENSE=$LITELLM_LICENSE \ - -e AWS_ACCESS_KEY_ID=$AWS_ACCESS_KEY_ID \ - -e AWS_SECRET_ACCESS_KEY=$AWS_SECRET_ACCESS_KEY \ - -e USE_DDTRACE=True \ - -e DD_API_KEY=$DD_API_KEY \ - -e DD_SITE=$DD_SITE \ - -e AWS_REGION_NAME=$AWS_REGION_NAME \ - -e PROXY_BATCH_WRITE_AT=2 \ - -e LITELLM_LOG=ERROR \ - --add-host host.docker.internal:host-gateway \ - --name my-app \ - -v $(pwd)/litellm/proxy/example_config_yaml/spend_tracking_config.yaml:/app/config.yaml \ - litellm-docker-database:ci \ - --config /app/config.yaml \ - --port 4000 - - run: - name: Start outputting logs - command: docker logs -f my-app - background: true - - wait_for_service: - url: http://localhost:4000 - timeout: "300" - - run: - name: Run tests - command: | - mkdir -p test-results - TEST_FILES=$(circleci tests glob "tests/spend_tracking_tests/**/test_*.py") - echo "$TEST_FILES" | circleci tests run \ - --verbose \ - --command="tr ' ' '\\n' | awk '/\\.py/ {print; next} {sub(/\\.[A-Z][^.]*$/, \"\"); gsub(/\\./, \"/\"); print \$0 \".py\"}' | xargs uv run --no-sync python -m pytest --tb=short \ - -vv \ - --junitxml=test-results/junit.xml \ - --durations=5" - no_output_timeout: 15m - - store_test_results: - path: test-results - - run: - name: Stop and remove first container - when: always - command: | - docker stop my-app - docker rm my-app - docker stop redis-cache - docker rm redis-cache - proxy_multi_instance_tests: machine: image: ubuntu-2204:2024.04.1 @@ -3591,9 +3455,6 @@ workflows: - proxy_logging_guardrails_model_info_tests: requires: - build_docker_database_image - - proxy_spend_accuracy_tests: - requires: - - build_docker_database_image - proxy_multi_instance_tests: requires: - build_docker_database_image @@ -3611,7 +3472,6 @@ workflows: - realtime_translation_testing - guardrails_testing - google_generate_content_endpoint_testing - - llm_responses_api_testing - ocr_testing - search_testing - batches_testing diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index ebb925c8b68..e5847e5fbea 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -143,9 +143,9 @@ jobs: run: | diff -u model_prices_and_context_window.json litellm/model_prices_and_context_window_backup.json if [ -z "$RUST_BRIDGE_ARTIFACT" ]; then - .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime + .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime --extra utils else - .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime --no-install-project + .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra caching --extra extra_proxy --extra proxy-runtime --extra utils --no-install-project uv pip install --no-deps --python .venv/bin/python rust-bridge-dist/*.whl cp rust-bridge-dist/litellm/rust_bridge/_native.abi3.so litellm/rust_bridge/_native.abi3.so uv run --no-sync python -c "import importlib.metadata; import litellm.rust_bridge._native; print(importlib.metadata.version('litellm'))" diff --git a/.github/workflows/publish-basedpyright-base-counts.yml b/.github/workflows/publish-basedpyright-base-counts.yml deleted file mode 100644 index 34f60d25980..00000000000 --- a/.github/workflows/publish-basedpyright-base-counts.yml +++ /dev/null @@ -1,65 +0,0 @@ -name: Publish basedpyright base counts - -# Every commit on main can become a future merge-base. -# Publishing its per-rule basedpyright counts as an artifact lets -# scripts/type_check_gate.py download them in seconds instead of paying a -# 60-110s second basedpyright pass on every fresh worktree or moved merge-base. -# No concurrency group on purpose: runs must never cancel each other, because -# every sha's artifact matters (any of them can become a merge-base). - -on: - push: - branches: - - main - workflow_dispatch: - inputs: - ref: - description: "Ref to compute and publish base counts for (defaults to the workflow run's commit)" - required: false - -permissions: - contents: read - -jobs: - publish: - runs-on: ubuntu-latest - timeout-minutes: 20 - - steps: - - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - ref: ${{ inputs.ref || github.sha }} - clean: true - 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: Cache the Rust build - uses: ./.github/actions/cache-cargo-build - - - name: Cache Prisma binaries - uses: ./.github/actions/cache-prisma-binaries - - # The gate provisions its own measurement env (.venv-typecheck: a frozen - # uv sync of its canonical dependency groups plus a generated Prisma - # client), so no install step here can drift from what local runs measure. - - name: Emit basedpyright counts for HEAD - run: | - python scripts/type_check_gate.py --emit-counts-dir "$RUNNER_TEMP/basedpyright-counts" - counts_file=$(ls "$RUNNER_TEMP"/basedpyright-counts/basedpyright-counts-*.json) - echo "COUNTS_ARTIFACT_NAME=$(basename "$counts_file" .json)" >> "$GITHUB_ENV" - - - name: Upload counts artifact - uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1 - with: - name: ${{ env.COUNTS_ARTIFACT_NAME }} - path: ${{ runner.temp }}/basedpyright-counts/ - if-no-files-found: error diff --git a/.github/workflows/publish-lint-base-counts.yml b/.github/workflows/publish-lint-base-counts.yml new file mode 100644 index 00000000000..7af00e9c842 --- /dev/null +++ b/.github/workflows/publish-lint-base-counts.yml @@ -0,0 +1,96 @@ +name: Publish lint base counts + +# Every commit on main can become a future merge-base. +# Publishing its per-rule counts for each lint gate (strict ruff, type discipline, +# test quality, basedpyright) as an artifact lets the gates download them through +# scripts/lint_base_counts.py in seconds instead of scanning the merge-base tree in +# a throwaway worktree on every fresh checkout or moved merge-base. +# No concurrency group on purpose: runs must never cancel each other, because +# every sha's artifact matters (any of them can become a merge-base). + +on: + push: + branches: + - main + workflow_dispatch: + inputs: + ref: + description: "Ref to compute and publish base counts for (defaults to the workflow run's commit)" + required: false + +permissions: + contents: read + +jobs: + publish: + runs-on: ubuntu-latest + timeout-minutes: 20 + strategy: + fail-fast: false + matrix: + include: + - checker: ruff-strict + script: scripts/ruff_strict_gate.py + - checker: type-discipline + script: scripts/type_discipline_gate.py + - checker: test-quality + script: scripts/test_quality_gate.py + - checker: basedpyright + script: scripts/type_check_gate.py + + steps: + - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + ref: ${{ inputs.ref || github.sha }} + clean: true + 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: Cache the Rust build + if: matrix.checker == 'basedpyright' + uses: ./.github/actions/cache-cargo-build + + - name: Cache Prisma binaries + if: matrix.checker == 'basedpyright' + uses: ./.github/actions/cache-prisma-binaries + + # The three source scanners only need the pinned dev tools (ruff and the + # stdlib checkers), the same versions test-linting.yml's lint job runs. + - name: Install the dev tools + if: matrix.checker != 'basedpyright' + run: | + uv sync --frozen --only-group dev --no-install-project + + - name: Emit ${{ matrix.checker }} counts for HEAD + if: matrix.checker != 'basedpyright' + run: | + uv run --no-sync python ${{ matrix.script }} --emit-counts-dir "$RUNNER_TEMP/lint-counts" + + # The basedpyright gate provisions its own measurement env (.venv-typecheck: + # a frozen uv sync of its canonical dependency groups plus a generated Prisma + # client), so no install step here can drift from what local runs measure. + - name: Emit basedpyright counts for HEAD + if: matrix.checker == 'basedpyright' + run: | + python scripts/type_check_gate.py --emit-counts-dir "$RUNNER_TEMP/lint-counts" + + - name: Name the counts artifact + run: | + counts_file=$(ls "$RUNNER_TEMP"/lint-counts/${{ matrix.checker }}-counts-*.json) + echo "COUNTS_ARTIFACT_NAME=$(basename "$counts_file" .json)" >> "$GITHUB_ENV" + + - name: Upload counts artifact + uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1 + with: + name: ${{ env.COUNTS_ARTIFACT_NAME }} + path: ${{ runner.temp }}/lint-counts/ + if-no-files-found: error diff --git a/.github/workflows/test-code-quality.yml b/.github/workflows/test-code-quality.yml index 26a2c427f32..8a3a9107b51 100644 --- a/.github/workflows/test-code-quality.yml +++ b/.github/workflows/test-code-quality.yml @@ -89,6 +89,9 @@ jobs: - name: test_workflow_job_name_collisions run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_workflow_job_name_collisions.py + - name: test_unit_passed_gate + run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_unit_passed_gate.py + - name: test_e2e_changed_gate run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_e2e_changed_gate.py tests/code_coverage_tests/test_e2e_idp_stack.py diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml index bed94a69873..db2fc969ace 100644 --- a/.github/workflows/test-linting.yml +++ b/.github/workflows/test-linting.yml @@ -17,9 +17,9 @@ jobs: lint: runs-on: ubuntu-latest timeout-minutes: 15 - # actions: read lets scripts/type_check_gate.py download the base-counts - # artifact published by publish-basedpyright-base-counts.yml instead of - # re-running basedpyright over the merge-base tree. + # actions: read lets the four lint gates download the base-counts artifacts + # published by publish-lint-base-counts.yml instead of re-scanning the + # merge-base tree in a throwaway worktree. permissions: contents: read pull-requests: read @@ -144,18 +144,24 @@ jobs: run: | uv run --no-sync ruff check --config ruff-tests.toml tests - - name: Check strict-rule budget (delta vs base) + - name: Check strict ruff rules (delta vs merge-base counts) if: steps.changes.outputs.decision != 'skip' + env: + GH_TOKEN: ${{ github.token }} run: | uv run --no-sync python scripts/ruff_strict_gate.py --base "$GATE_BASE_SHA" - - name: Check type-discipline budget (mutable collections / casts / type guards / kwargs / unexplained suppressions, delta vs base) + - name: Check type discipline (mutable collections / casts / type guards / kwargs / unexplained suppressions, delta vs merge-base counts) if: steps.changes.outputs.decision != 'skip' + env: + GH_TOKEN: ${{ github.token }} run: | uv run --no-sync python scripts/type_discipline_gate.py --base "$GATE_BASE_SHA" - - name: Check test-quality budget (zero-assert / mock-echo tests, sys.path.insert, raw env writes, litellm global mutation, credential-gated skips, conftest snapshot inventory, delta vs base) + - name: Check test quality (zero-assert / mock-echo tests, sys.path.insert, raw env writes, litellm global mutation, credential-gated skips, conftest snapshot inventory, delta vs merge-base counts) if: steps.changes.outputs.decision != 'skip' + env: + GH_TOKEN: ${{ github.token }} run: | uv run --no-sync python scripts/test_quality_gate.py --base "$GATE_BASE_SHA" @@ -164,7 +170,7 @@ jobs: run: | uv run --no-sync python -c "import openai; print(f'OpenAI version: {openai.__version__}')" - - name: Check basedpyright budget (delta vs base) + - name: Check basedpyright (delta vs merge-base counts) if: steps.changes.outputs.decision != 'skip' env: GH_TOKEN: ${{ github.token }} @@ -206,40 +212,6 @@ jobs: run: | uv run --no-sync python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1) - # Intentionally NON-GATING. This job turns red when a *-budget.json ceiling is - # raised (or a rule/budget is dropped) so a loosening is obvious in review, but it - # must be kept OUT of the branch-protection required-checks list so a justified - # bump can still be merged by a human who has seen and accepted the red. - budget-ratchet: - runs-on: ubuntu-latest - timeout-minutes: 5 - permissions: - contents: read - - steps: - - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - fetch-depth: 1 - persist-credentials: false - - - name: Fetch ratchet base - env: - BASE_SHA: ${{ github.event.pull_request.base.sha }} - run: | - retry() { "$@" || { sleep 15; "$@"; } || { sleep 30; "$@"; }; } - retry git fetch --no-tags --depth=1 origin "$BASE_SHA" - - - name: Set up Python - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 - with: - python-version: "3.12" - - - name: Ratchet check (budgets may only decrease; non-gating) - env: - BASE_SHA: ${{ github.event.pull_request.base.sha }} - run: | - python scripts/budget_ratchet_check.py --base "$BASE_SHA" - secret-scan: runs-on: ubuntu-latest timeout-minutes: 5 diff --git a/.github/workflows/test-mcp-dependency-resolution.yml b/.github/workflows/test-mcp-dependency-resolution.yml index 1772bdeabc6..3f0d1260e49 100644 --- a/.github/workflows/test-mcp-dependency-resolution.yml +++ b/.github/workflows/test-mcp-dependency-resolution.yml @@ -13,6 +13,9 @@ on: - "rust-toolchain.toml" - "litellm-rust/**" - "litellm/__init__.py" + - "litellm/_version.py" + - "litellm/_lazy_imports_registry.py" + - "scripts/build_core_distribution.py" - "litellm/proxy/proxy_server.py" - "litellm/**/*mcp*" - "litellm/**/*mcp*/**" @@ -86,6 +89,7 @@ jobs: uv pip check --python ".venv-$extra" if [ "$extra" = core ]; then checker=("$GITHUB_WORKSPACE/tests/base_sdk_tests/check_base_sdk_install.py") + (cd "$RUNNER_TEMP" && "$GITHUB_WORKSPACE/.venv-$extra/bin/python" -I "$GITHUB_WORKSPACE/tests/base_sdk_tests/check_sdk_http.py") else checker=("$GITHUB_WORKSPACE/scripts/check_mcp_sdk_install.py" --extra "$extra") fi @@ -115,3 +119,63 @@ jobs: fi (cd "$RUNNER_TEMP" && "$GITHUB_WORKSPACE/.venv-lowest-$extra/bin/python" "${checker[@]}") done + + core-distribution: + permissions: + contents: read + id-token: write + runs-on: ubuntu-latest + timeout-minutes: 30 + env: + UV_PYTHON: ${{ matrix.python-version }} + LITELLM_LOCAL_MODEL_COST_MAP: "True" + strategy: + fail-fast: false + matrix: + python-version: ["3.10", "3.11", "3.12", "3.13", "3.14"] + steps: + - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + persist-credentials: false + - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: ${{ matrix.python-version }} + - uses: ./.github/actions/setup-uv-with-retries + with: + version: "0.10.9" + - uses: ./.github/actions/cache-cargo-build + - name: Build the core wheel and source distribution + run: >- + uv run --no-project --with coverage==7.14.0 python -m coverage run --rcfile=/dev/null + --branch --include="*/scripts/build_core_distribution.py" + scripts/build_core_distribution.py --out-dir dist/core + - name: Verify metadata, resources, and independent source rebuild + env: + CORE_DISTRIBUTION_DIR: ${{ github.workspace }}/dist/core + run: >- + uv run --no-project --with pytest==9.0.3 --with pytest-cov==5.0.0 --with coverage==7.14.0 --with 'tomli==2.4.1; python_version < "3.11"' + python -m pytest tests/base_sdk_tests/test_core_distribution.py -v + --cov-config=/dev/null --cov=scripts.build_core_distribution --cov-branch --cov-append --cov-report= + - name: Verify isolated core installations + run: | + wheel=$(realpath dist/core/litellm_core-*.whl) + for resolution in highest lowest-direct; do + uv venv --python ${{ matrix.python-version }} ".venv-core-$resolution" + uv pip install --python ".venv-core-$resolution" --resolution "$resolution" "$wheel" + uv pip check --python ".venv-core-$resolution" + (cd "$RUNNER_TEMP" && "$GITHUB_WORKSPACE/.venv-core-$resolution/bin/python" -I "$GITHUB_WORKSPACE/tests/base_sdk_tests/check_base_sdk_install.py" --profile core) + (cd "$RUNNER_TEMP" && "$GITHUB_WORKSPACE/.venv-core-$resolution/bin/python" -I "$GITHUB_WORKSPACE/tests/base_sdk_tests/check_sdk_http.py") + done + uv pip install --python .venv-core-highest coverage==7.14.0 + .venv-core-highest/bin/python -I -m coverage run --rcfile=/dev/null --append --branch --include="*/tests/base_sdk_tests/check_base_sdk_install.py" tests/base_sdk_tests/check_base_sdk_install.py --profile core + uv pip install --python .venv-core-highest boto3 tokenizers huggingface-hub + (cd "$RUNNER_TEMP" && "$GITHUB_WORKSPACE/.venv-core-highest/bin/python" -I "$GITHUB_WORKSPACE/tests/base_sdk_tests/check_base_sdk_install.py" --profile dependencies) + (cd "$RUNNER_TEMP" && "$GITHUB_WORKSPACE/.venv-core-highest/bin/python" -I "$GITHUB_WORKSPACE/tests/base_sdk_tests/check_sdk_http.py") + uv run --no-project --with coverage==7.14.0 python -m coverage xml --rcfile=/dev/null -o coverage-core.xml + - name: Upload core packaging coverage + uses: codecov/codecov-action@0fb7174895f61a3b6b78fc075e0cd60383518dac # v5.5.5 + with: + use_oidc: true + files: coverage-core.xml + flags: core-packaging + fail_ci_if_error: false diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index c07487d5b49..a9f75b5f87e 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -98,6 +98,9 @@ jobs: rust-test: runs-on: ubuntu-latest timeout-minutes: 30 + # Trim debuginfo to keep the job's target dir within the runner's disk + env: + CARGO_PROFILE_DEV_DEBUG: line-tables-only defaults: run: working-directory: litellm-rust diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml deleted file mode 100644 index ba88464d652..00000000000 --- a/.github/workflows/test-unit-proxy-db.yml +++ /dev/null @@ -1,223 +0,0 @@ -name: "Unit Tests: Proxy DB Operations" - -on: - pull_request: - branches: - - main - - "litellm_**" - push: - branches: - - main - -permissions: - contents: read - -concurrency: - group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }} - cancel-in-progress: ${{ github.event_name == 'pull_request' }} - -# Semantic matrix: each shard groups tests by concern (auth, server, logging, …) -# rather than alphabetical letter ranges. Adding a new test file means adding it -# to whichever group it belongs to, not reshuffling slices. -# -# Design targets: -# * Every shard runs in <= 7 minutes of wall-clock on the default runner. -# Most of a shard's time is pytest plugin load + xdist worker imports + -# pytest-cov instrumentation, not the tests themselves. Keeping per-shard -# work low and matching worker count to runner cores is what controls it. -# * `timeout` bounds the pytest step only. Checkout, dependency install, and -# Prisma client generation draw on a separate allowance in the base -# workflow, so slow setup shows up as a slow job rather than as a -# cancelled shard whose tests were passing. -# * workers: 4 matches the 4-core ubuntu-latest runner. -n 8 on 4 cores -# oversubscribes 2x and workers fight for CPU during their cold-start -# imports (measured ~441% CPU for -n 8 locally, i.e. ~55% effective). -# * test_key_generate_prisma.py stays serial (workers=0) — it has event-loop -# conflicts with the logging worker when run in parallel. -# * test_proxy_utils.py runs as a single shard with --dist=worksteal so -# xdist balances its 188 parametrized cases across workers instead of -# pinning the whole file to one worker (the default --dist=loadscope -# behavior for single-file targets). -jobs: - lens-python-310: - name: Lens Python 3.10 - permissions: - contents: read - id-token: write - pull-requests: write - uses: ./.github/workflows/_test-unit-base.yml - with: - python-version: "3.10" - test-path: tests/unit/proxy/lens/test_inference.py - workers: 0 - reruns: 0 - timeout-minutes: 5 - artifact-name: lens-python-310 - - # Fast guard — fails the workflow when a test directory or file inside a sharded - # tree is claimed by no shard. The semantic-shard design has no catch-all bucket, - # so an unassigned child runs nowhere; assert_ci_coverage.py holds the tree list - # and reads the same test-path keys the coverage census does. - assert-shard-coverage: - runs-on: ubuntu-latest - timeout-minutes: 2 - permissions: - contents: read - steps: - - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - persist-credentials: false - - name: Assert every test directory and file is claimed by a shard - run: python3 .github/scripts/assert_ci_coverage.py --shards - - proxy-db: - needs: assert-shard-coverage - # Display only the semantic shard name in the checks UI instead of GHA's - # default "proxy-db (key-generation, tests/unit/proxy/…, 0, loadscope, 20)" - # which includes every matrix field and gets truncated past the test-path. - name: ${{ matrix.test-group }} - permissions: - contents: read - id-token: write - pull-requests: write - strategy: - fail-fast: false - matrix: - include: - # Must run serially — event-loop conflict with the logging worker. - - test-group: key-generation - test-path: >- - tests/unit/proxy/management_endpoints/test_key_generate_prisma.py - workers: 0 - dist: loadscope - timeout: 20 - - # ---- auth: split into 2 shards ---- - - test-group: auth-checks - test-path: >- - tests/unit/proxy/auth/test_auth_checks.py - tests/unit/proxy/auth/test_user_api_key_auth.py - tests/unit/proxy/test_credential_slot_registry.py - tests/unit/proxy/test_deprecated_key_grace_period.py - workers: 4 - dist: loadscope - timeout: 15 - - test-group: jwt-and-keys - test-path: >- - tests/unit/proxy/auth/test_jwt.py - tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py - tests/unit/proxy/test_proxy_custom_auth.py - workers: 4 - dist: loadscope - timeout: 15 - - # ---- test_proxy_utils.py, single shard, worksteal distribution ---- - - test-group: proxy-utils - test-path: >- - tests/unit/proxy/test_proxy_utils.py - workers: 4 - dist: worksteal - timeout: 15 - - # ---- proxy server: split into 2 shards ---- - - test-group: proxy-server-core - test-path: >- - tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py - tests/unit/proxy/test__lazy_features.py - tests/unit/proxy/test_aproxy_startup.py - tests/unit/proxy/test_proxy_server.py - workers: 4 - dist: loadscope - timeout: 15 - - test-group: proxy-runtime - test-path: >- - tests/unit/proxy/auth/test_multipart_bypass_repro.py - tests/unit/proxy/auth/test_proxy_routes.py - tests/unit/proxy/middleware/test_request_size_limit_middleware.py - tests/unit/proxy/test_proxy_config_unit_test.py - tests/unit/proxy/test_proxy_token_counter.py - tests/unit/proxy/test_server_root_path.py - workers: 4 - dist: loadscope - timeout: 15 - - # ---- logging: split into 2 shards ---- - - test-group: custom-logging - test-path: >- - tests/proxy_unit_tests/test_proxy_custom_logger.py - tests/unit/proxy/test_custom_callback_input.py - tests/unit/proxy/test_custom_logger_s3_gcs.py - workers: 4 - dist: loadscope - timeout: 15 - - test-group: logging-misc - test-path: >- - tests/unit/proxy/management_helpers/test_audit_logs_proxy.py - tests/unit/proxy/spend_tracking/test_search_api_logging.py - tests/unit/proxy/test_proxy_reject_logging.py - workers: 4 - dist: loadscope - timeout: 15 - - - test-group: db-and-spend - test-path: >- - tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py - tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py - tests/unit/proxy/db/test_update_daily_tag_spend.py - tests/unit/proxy/test_db_schema_changes.py - tests/unit/proxy/test_prisma_client_backoff_retry.py - tests/unit/proxy/test_update_spend.py - tests/unit/skills/test_skills_db.py - workers: 4 - dist: loadscope - timeout: 15 - - # ---- guardrails + budget + hooks: split into 2 ---- - - test-group: guardrails-hooks - test-path: >- - tests/unit/proxy/hooks/test_banned_keyword_list.py - tests/unit/proxy/test_proxy_setting_guardrails.py - tests/unit/proxy/test_unit_test_proxy_hooks.py - workers: 4 - dist: loadscope - timeout: 15 - - test-group: budgets - test-path: >- - tests/unit/proxy/auth/test_default_end_user_budget_simple.py - tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py - tests/unit/proxy/test_zero_cost_model_budget_bypass.py - workers: 4 - dist: loadscope - timeout: 15 - - - test-group: endpoints-and-responses - test-path: >- - tests/proxy_unit_tests/test_proxy_exception_mapping.py - tests/unit/proxy/lens - tests/unit/proxy/auth/test_models_fallback_endpoint.py - tests/unit/proxy/common_utils/test_check_batch_cost.py - tests/unit/proxy/common_utils/test_check_responses_cost.py - tests/unit/proxy/common_utils/test_realtime_cache.py - tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py - tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py - tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py - tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py - tests/unit/proxy/response_polling - tests/unit/proxy/test_custom_tokenizer_bug.py - tests/unit/proxy/test_get_favicon.py - tests/unit/proxy/test_get_image.py - tests/unit/proxy/test_prompt_test_endpoint.py - tests/unit/proxy/test_reducto_ocr_route.py - tests/unit/proxy/test_response_polling_pre_call_checks.py - tests/unit/proxy/test_ui_path_detection.py - workers: 4 - dist: loadscope - timeout: 15 - uses: ./.github/workflows/_test-unit-base.yml - with: - test-path: ${{ matrix.test-path }} - workers: ${{ matrix.workers }} - reruns: 2 - timeout-minutes: ${{ matrix.timeout }} - dist: ${{ matrix.dist }} - artifact-name: proxy-db-${{ matrix.test-group }} diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 54fdc6b43a2..21b82db44b7 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -28,11 +28,6 @@ concurrency: # defaults. An absent matrix key renders as an empty string, which is not a # number, so a partially-specified entry would fail the call rather than fall # back to the default. -# -# tests/unit/proxy keeps its own caller (test-unit-proxy-db.yml): it is already -# a matrix and carries a shard-coverage guard that reads that file by name. -# Folding it in here is a follow-up, together with generalising that guard into -# assert_ci_coverage.py. jobs: rust-bridge: name: Build the Rust bridge @@ -346,6 +341,16 @@ jobs: timeout-minutes: 60 job-timeout-minutes: 100 + - shard: mcp-elicitation + artifact-name: mcp-elicitation + test-path: >- + tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py + tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py + workers: 2 + reruns: 0 + timeout-minutes: 20 + job-timeout-minutes: 60 + - shard: proxy-infra artifact-name: proxy-infra test-path: >- @@ -498,3 +503,211 @@ jobs: job-timeout-minutes: ${{ matrix.job-timeout-minutes }} dist: ${{ matrix.dist || 'loadscope' }} artifact-name: ${{ matrix.artifact-name }} + + lens-python-310: + name: Lens Python 3.10 + permissions: + contents: read + id-token: write + pull-requests: write + uses: ./.github/workflows/_test-unit-base.yml + with: + python-version: "3.10" + test-path: tests/unit/proxy/lens/test_inference.py + workers: 0 + reruns: 0 + timeout-minutes: 5 + artifact-name: lens-python-310 + + # Fast guard — fails the workflow when a test directory or file inside a sharded + # tree is claimed by no shard. The semantic-shard design has no catch-all bucket, + # so an unassigned child runs nowhere; assert_ci_coverage.py holds the tree list + # and reads the same test-path keys the coverage census does. + assert-shard-coverage: + runs-on: ubuntu-latest + timeout-minutes: 2 + permissions: + contents: read + steps: + - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + persist-credentials: false + - name: Assert every test directory and file is claimed by a shard + run: python3 .github/scripts/assert_ci_coverage.py --shards + + proxy-db: + needs: assert-shard-coverage + # Display only the semantic shard name in the checks UI instead of GHA's + # default "proxy-db (key-generation, tests/unit/proxy/…, 0, loadscope, 20)" + # which includes every matrix field and gets truncated past the test-path. + name: ${{ matrix.test-group }} + permissions: + contents: read + id-token: write + pull-requests: write + strategy: + fail-fast: false + matrix: + include: + # Must run serially — event-loop conflict with the logging worker. + - test-group: key-generation + test-path: >- + tests/unit/proxy/management_endpoints/test_key_generate_prisma.py + workers: 0 + dist: loadscope + timeout: 20 + + # ---- auth: split into 2 shards ---- + - test-group: auth-checks + test-path: >- + tests/unit/proxy/auth/test_auth_checks.py + tests/unit/proxy/auth/test_user_api_key_auth.py + tests/unit/proxy/test_credential_slot_registry.py + tests/unit/proxy/test_deprecated_key_grace_period.py + workers: 4 + dist: loadscope + timeout: 15 + - test-group: jwt-and-keys + test-path: >- + tests/unit/proxy/auth/test_jwt.py + tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py + tests/unit/proxy/test_proxy_custom_auth.py + workers: 4 + dist: loadscope + timeout: 15 + + # ---- test_proxy_utils.py, single shard, worksteal distribution ---- + - test-group: proxy-utils + test-path: >- + tests/unit/proxy/test_proxy_utils.py + workers: 4 + dist: worksteal + timeout: 15 + + # ---- proxy server: split into 2 shards ---- + - test-group: proxy-server-core + test-path: >- + tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py + tests/unit/proxy/test__lazy_features.py + tests/unit/proxy/test_aproxy_startup.py + tests/unit/proxy/test_proxy_server.py + workers: 4 + dist: loadscope + timeout: 15 + - test-group: proxy-runtime + test-path: >- + tests/unit/proxy/auth/test_multipart_bypass_repro.py + tests/unit/proxy/auth/test_proxy_routes.py + tests/unit/proxy/middleware/test_request_size_limit_middleware.py + tests/unit/proxy/test_proxy_config_unit_test.py + tests/unit/proxy/test_proxy_token_counter.py + tests/unit/proxy/test_server_root_path.py + workers: 4 + dist: loadscope + timeout: 15 + + - test-group: mcp-oauth + test-path: >- + tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py + tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py + tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py + tests/unit/proxy/_experimental/mcp_server/outbound_credentials + workers: 4 + dist: loadscope + timeout: 15 + + # ---- logging: split into 2 shards ---- + - test-group: custom-logging + test-path: >- + tests/proxy_unit_tests/test_proxy_custom_logger.py + tests/unit/proxy/test_custom_callback_input.py + tests/unit/proxy/test_custom_logger_s3_gcs.py + workers: 4 + dist: loadscope + timeout: 15 + - test-group: logging-misc + test-path: >- + tests/unit/proxy/management_helpers/test_audit_logs_proxy.py + tests/unit/proxy/spend_tracking/test_search_api_logging.py + tests/unit/proxy/test_proxy_reject_logging.py + workers: 4 + dist: loadscope + timeout: 15 + + - test-group: db-and-spend + test-path: >- + tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py + tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py + tests/unit/proxy/db/test_update_daily_tag_spend.py + tests/unit/proxy/test_db_schema_changes.py + tests/unit/proxy/test_prisma_client_backoff_retry.py + tests/unit/proxy/test_update_spend.py + tests/unit/skills/test_skills_db.py + workers: 4 + dist: loadscope + timeout: 15 + + # ---- guardrails + budget + hooks: split into 2 ---- + - test-group: guardrails-hooks + test-path: >- + tests/unit/proxy/hooks/test_banned_keyword_list.py + tests/unit/proxy/test_proxy_setting_guardrails.py + tests/unit/proxy/test_unit_test_proxy_hooks.py + workers: 4 + dist: loadscope + timeout: 15 + - test-group: budgets + test-path: >- + tests/unit/proxy/auth/test_default_end_user_budget_simple.py + tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py + tests/unit/proxy/test_zero_cost_model_budget_bypass.py + workers: 4 + dist: loadscope + timeout: 15 + + - test-group: endpoints-and-responses + test-path: >- + tests/proxy_unit_tests/test_proxy_exception_mapping.py + tests/unit/proxy/lens + tests/unit/proxy/auth/test_models_fallback_endpoint.py + tests/unit/proxy/common_utils/test_check_batch_cost.py + tests/unit/proxy/common_utils/test_check_responses_cost.py + tests/unit/proxy/common_utils/test_realtime_cache.py + tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py + tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py + tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py + tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py + tests/unit/proxy/response_polling + tests/unit/proxy/test_custom_tokenizer_bug.py + tests/unit/proxy/test_get_favicon.py + tests/unit/proxy/test_get_image.py + tests/unit/proxy/test_prompt_test_endpoint.py + tests/unit/proxy/test_reducto_ocr_route.py + tests/unit/proxy/test_response_polling_pre_call_checks.py + tests/unit/proxy/test_ui_path_detection.py + workers: 4 + dist: loadscope + timeout: 15 + uses: ./.github/workflows/_test-unit-base.yml + with: + test-path: ${{ matrix.test-path }} + workers: ${{ matrix.workers }} + reruns: 2 + timeout-minutes: ${{ matrix.timeout }} + dist: ${{ matrix.dist }} + artifact-name: proxy-db-${{ matrix.test-group }} + + unit-passed: + name: unit passed + needs: [rust-bridge, unit, lens-python-310, assert-shard-coverage, proxy-db] + if: always() + runs-on: ubuntu-latest + timeout-minutes: 2 + permissions: {} + steps: + - name: Require every unit job to succeed + env: + NEEDS: ${{ toJSON(needs) }} + run: | + jq -r 'to_entries[] | "\(.key): \(.value.result)"' <<< "$NEEDS" + jq -e 'all(.[]; .result == "success")' <<< "$NEEDS" > /dev/null diff --git a/AGENTS.md b/AGENTS.md index 948a46d4172..8ec1e75bd0a 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -54,13 +54,13 @@ Don't hesitate to use values in .env to get needed API keys and other secrets, a Python max line length is 120, not 88 -Never edit or commit `ruff-strict-budget.json`, `type-discipline-budget.json`, `basedpyright-code-budget.json`, or `test-quality-budget.json` on a PR branch, and don't run `make lint-budget-update` there. A scheduled Devin automation lowers the limits on the default branch in its own PR by exactly what landed since the last ratchet, so concurrent PRs don't fight over the same `"limit"` lines. Keep the hosted automation's target in sync when the repository default changes. If your branch already carries a budget edit, drop it before opening the PR +The four lint gates (`scripts/ruff_strict_gate.py`, `scripts/type_discipline_gate.py`, `scripts/test_quality_gate.py`, `scripts/type_check_gate.py`) compare each rule's codebase count on your branch against the count at its merge-base with the default branch, and a rule may not grow. There are no budget files to edit or ratchet: when a gate fails, fix the violations the branch introduced or remove at least as many of that rule elsewhere in the tree. The one exception is `reportAny` / `reportExplicitAny`, which share a fixed codebase-wide cap in `ANY_CAPS` in `scripts/type_check_gate.py` because Any spreads past the lines you touch. A branch may add Anys while the total stays under that cap. Only lower `ANY_CAPS`, and only in its own PR to the default branch, never raise it on a feature branch `make check` (f.k.a. `make pre-commit`, which still works identically as an alias) saves its complete output to a log file in .git (overwriting previous logs) and prints that path as its first and last output lines. To inspect a run, read or grep that log instead of re-running the multi-minute checks just to see a different slice -`make check`, `make lint`, `scripts/pre_commit_lint.sh`, and the standalone budget gates (`scripts/ruff_strict_gate.py`, `scripts/type_discipline_gate.py`, `scripts/type_check_gate.py`) each hold one of 2 machine-wide slots, so when other sessions or worktrees on the same box are already running heavy work, yours prints "all N machine-wide slots are busy; queueing" and then stays quiet until a slot frees. Give the command a long timeout and let it wait rather than killing it, retrying it, or assuming it hung. Don't change the # of machine-wide slots or make it unlimited by setting `LITELLM_GATE_SLOTS=0` +`make check`, `make lint`, `scripts/pre_commit_lint.sh`, and the standalone lint gates (`scripts/ruff_strict_gate.py`, `scripts/type_discipline_gate.py`, `scripts/type_check_gate.py`) each hold one of 2 machine-wide slots, so when other sessions or worktrees on the same box are already running heavy work, yours prints "all N machine-wide slots are busy; queueing" and then stays quiet until a slot frees. Give the command a long timeout and let it wait rather than killing it, retrying it, or assuming it hung. Don't change the # of machine-wide slots or make it unlimited by setting `LITELLM_GATE_SLOTS=0` -If you're trying to create a new function that relies on untyped stuff, instead of adding more Any's and pushing `reportAny` / `reportExplicitAny` closer to their basedpyright ceilings, just validate it in the caller with Pydantic (a model or `TypeAdapter` that returns the typed thing or raises will do) and then pass the now typed variable in +If you're trying to create a new function that relies on untyped stuff, instead of adding more Any's and pushing `reportAny` / `reportExplicitAny` closer to their cap in `ANY_CAPS`, just validate it in the caller with Pydantic (a model or `TypeAdapter` that returns the typed thing or raises will do) and then pass the now typed variable in If you get an LIT001 fail, refactor the code to follow functional programming best practices rather than introducing mutable sequences or sets. For example, build values in one shot with comprehensions or generators wrapped in `tuple()` / `frozenset()` instead of seeding an empty `list`/`set` and appending to it over time. Ideally, `# mutable-ok` is never used; reach for it only as a last resort when an immutable rewrite is impossible, and always pair it with a real reason. Plain `dict` is allowed: most of the Python ecosystem takes and returns dicts, and converting to `MappingProxyType` at every boundary costs more than it protects. Never deep copy a value just to hand out an immutable view; a defensive copy of a self-referential or large object is worse than the mutation it guards against @@ -97,6 +97,7 @@ Follow these coding conventions for new/updated code (a three-line fix in a lega - Annotate every variable with `: Final` (LIT010). Unpacking and walrus targets cannot carry the annotation, so they are implicitly final. Don't rebind them. Never rebind or mutate function parameters (LIT011); `self`/`cls` attribute stores are the exception. If rebinding or in-place mutation is truly unavoidable, suppress with `# rebind-ok: ` - Qualify every TypedDict field with `ReadOnly[...]` (LIT012), which nests freely with `Required` / `NotRequired` / `Annotated` in any order. If making the key writable is truly unavoidable, suppress with `# writable-ok: ` - Comprehensions take at most one `for` clause and one `if` clause (LIT014); split stacked clauses into a helper generator, a named intermediate, or a plain loop. Suppress with `# comprehension-ok: ` only when unavoidable + - Every pydantic model must be frozen (LIT015), set via `model_config = ConfigDict(frozen=True)`, a dict-literal `model_config`, an inner `class Config`, or the class keywords. Subclasses inherit it unless they override it. Replace in-place field writes with `model_copy(update=...)`. If making the model mutable is truly unavoidable, suppress with `# frozen-ok: ` on the `class` line - Use dependency injection - Fully typed; no `Any` or coarse types like `dict[str, Any]` or just `dict`. Every function parameter must be strongly typed - Use tagged unions + match @@ -129,4 +130,4 @@ Before implementing: Ask yourself: "Would a senior engineer say this is overcomplicated?" If yes, simplify -Before requesting maintainer review, verify the current PR tip passes required CI and code coverage, meets Greptile confidence of at least 4/5, and has acceptable Veria and Bugbot reviews. Inspect warnings and findings, fix actionable issues, and rerun the affected checks and reviewers after changes. Record evidence for any false positive or unavailable review; never treat a pending or missing bot result as a pass. Do not lower coverage thresholds or lint budgets to satisfy a check +Before requesting maintainer review, verify the current PR tip passes required CI and code coverage, meets Greptile confidence of at least 4/5, and has acceptable Veria and Bugbot reviews. Inspect warnings and findings, fix actionable issues, and rerun the affected checks and reviewers after changes. Record evidence for any false positive or unavailable review; never treat a pending or missing bot result as a pass. Do not lower coverage thresholds or raise `ANY_CAPS` to satisfy a check diff --git a/Makefile b/Makefile index 25b966e2b9a..b8614ca4eec 100644 --- a/Makefile +++ b/Makefile @@ -6,9 +6,8 @@ test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \ test-rust-extension rust-sqlx-prepare lens-dev \ info lint lint-inner lint-dev lint-checks format \ - lint-basedpyright lint-e2e-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \ - lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \ - lint-test-quality lint-test-quality-budget-update \ + lint-basedpyright lint-e2e-basedpyright lint-type-discipline \ + lint-ruff-strict lint-gate lint-test-quality \ install-dev install-proxy-dev install-test-deps install-hooks \ install-helm-unittest check-circular-imports check-import-safety check check-inner pre-commit \ lint-install lint-fetch-base bootstrap @@ -32,13 +31,10 @@ help: @echo " make lint-ruff - Run Ruff linting only" @echo " make lint-basedpyright - Run basedpyright strict, gated by per-rule error counts" @echo " make lint-e2e-basedpyright - Run basedpyright over tests/e2e and tests/e2e_harness (zero errors allowed)" - @echo " make lint-basedpyright-budget-update - Ratchet basedpyright limits down by what this branch fixed" @echo " make lint-format - Check ruff format formatting (matches CI)" - @echo " make lint-ruff-budget - Gate the codebase total of each strict ruff rule against its limit" + @echo " make lint-ruff-strict - Gate each strict ruff rule's codebase total against its merge-base count" @echo " make lint-gate - Strict ruff gate in CI-parity mode (fetches the default branch, simulates the merge)" - @echo " make lint-ruff-budget-update - Ratchet ruff-strict-budget.json limits down by what this branch fixed" - @echo " make lint-test-quality - Gate the test suite against test-quality-budget.json" - @echo " make lint-budget-update - Ratchet all budgets down (ruff + type-discipline + test quality + basedpyright)" + @echo " make lint-test-quality - Gate the test suite's TQ counts against their merge-base counts" @echo " make check-circular-imports - Check for circular imports" @echo " make check-import-safety - Check import safety" @echo " make test - Run all tests" @@ -143,7 +139,7 @@ lint-fetch-base: # Mirror test-linting.yml's lint job environment: the proxy-dev group plus a generated # Prisma client, so `basedpyright tests/e2e` resolves the same modules CI does. The -# budget gate itself no longer measures here (scripts/type_check_gate.py provisions its +# basedpyright gate itself no longer measures here (scripts/type_check_gate.py provisions its # own .venv-typecheck). --inexact tops up the venv instead of pruning the proxy extras # gen:api and the running proxy need. lint-install: @@ -213,25 +209,20 @@ lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE) lint-e2e-basedpyright: $(LINT_E2E_DEP_INSTALL) $(UV_RUN) basedpyright tests/e2e tests/e2e_harness -# Type-discipline budget (mutable collections / casts / type guards / kwargs / +# Type-discipline gate (mutable collections / casts / type guards / kwargs / # unexplained suppressions), the test-linting.yml step `make lint` used to omit. lint-type-discipline: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE) $(UV_RUN) python scripts/type_discipline_gate.py --base "$(BASE_REF)" -# Test-quality budget (zero-assert / mock-echo tests, sys.path.insert, raw env writes, +# Test-quality gate (zero-assert / mock-echo tests, sys.path.insert, raw env writes, # litellm module-global mutation, credential-gated skips, conftest snapshot # inventory), counted across tests/ the same delta-vs-base way. lint-test-quality: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE) $(UV_RUN) python scripts/test_quality_gate.py --base "$(BASE_REF)" -# --update lowers each limit by what this branch fixed since its branch point, so -# it needs the base ref fetched to resolve the merge-base. -lint-basedpyright-budget-update: install-dev - $(UV_RUN) python scripts/type_check_gate.py --update --base "$(BASE_REF)" - lint-format: format-check -lint-ruff-budget: install-dev +lint-ruff-strict: install-dev $(UV_RUN) python scripts/ruff_strict_gate.py --base "$(BASE_REF)" # Strict gate, invoked the same way CI does in test-linting.yml so a local pass @@ -239,18 +230,6 @@ lint-ruff-budget: install-dev lint-gate: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE) $(UV_RUN) python scripts/ruff_strict_gate.py --base "$(BASE_REF)" -lint-ruff-budget-update: install-dev - $(UV_RUN) python scripts/ruff_strict_gate.py --update --base "$(BASE_REF)" - -lint-type-discipline-budget-update: install-dev - $(UV_RUN) python scripts/type_discipline_gate.py --update --base "$(BASE_REF)" - -lint-test-quality-budget-update: install-dev - $(UV_RUN) python scripts/test_quality_gate.py --update --base "$(BASE_REF)" - -# Ratchet all budgets in one shot (ruff strict + type-discipline + test quality + basedpyright) -lint-budget-update: lint-ruff-budget-update lint-type-discipline-budget-update lint-test-quality-budget-update lint-basedpyright-budget-update - check-circular-imports: $(LINT_DEP_INSTALL) cd litellm && $(UV_RUN) python ../tests/documentation_tests/test_circular_imports.py && cd .. @@ -260,7 +239,7 @@ check-import-safety: $(LINT_DEP_INSTALL) # Combined linting, isomorphic to test-linting.yml's lint job so a local pass means a # green CI lint: it installs the same env (proxy-dev + generated Prisma client) and then # runs the diff-scoped ruff format check, whole-tree ruff check, the strict-rule / -# type-discipline / basedpyright budgets as a delta vs the base, then the circular-import +# type-discipline / basedpyright gates as a delta vs the base, then the circular-import # and import-safety checks. Steps that compare against the base resolve it the same way CI # does (merge-base with origin's current default branch). Setup (env sync, Prisma client, # base fetch) runs once up front; the checks themselves are independent, so a sub-make diff --git a/README.md b/README.md index 7ffc44854bb..141e8232de8 100644 --- a/README.md +++ b/README.md @@ -90,6 +90,24 @@ Managing LLM calls across providers gets complicated fast — different SDKs, au uv add litellm ``` +An independent `litellm-core` distribution provides the Python SDK with the same +`import litellm` API and runtime dependencies. It has no optional extras, CLI entry +points, or bundled dashboard. Install one SDK distribution per environment because +`litellm` and `litellm-core` own overlapping Python files. Use `litellm` for the +proxy, CLI, and optional extras + +Build and install core from this checkout while its release integration is pending: + +```shell +python scripts/build_core_distribution.py --out-dir dist/core +python -m pip install dist/core/litellm_core-*.whl +``` + +The builder requires Git, `uv` and the Rust build toolchain. It reads the release version +from the root `pyproject.toml`, leaves source files unchanged, and produces a wheel +and a self-contained sdist in `dist/core`. Run installation in a fresh environment +without `litellm` + ```python from litellm import completion import os @@ -401,7 +419,7 @@ Set `LITELLM_PROXY_API_BASE` and `LITELLM_PROXY_API_KEY` and every model call th | [Vercel AI Gateway (`vercel_ai_gateway`)](https://docs.litellm.ai/docs/providers/vercel_ai_gateway) | ✅ | ✅ | ✅ | | | | | | | | | [VLLM (`vllm`)](https://docs.litellm.ai/docs/providers/vllm) | ✅ | ✅ | ✅ | | | | | | | | | [Volcengine (`volcengine`)](https://docs.litellm.ai/docs/providers/volcano) | ✅ | ✅ | ✅ | | | | | | | | -| [Voyage AI (`voyage`)](https://docs.litellm.ai/docs/providers/voyage) | | | | ✅ | | | | | | | +| [VoyageAI by MongoDB (`voyage`)](https://docs.litellm.ai/docs/providers/voyage) | | | | ✅ | | | | | | | | [WandB Inference (`wandb`)](https://docs.litellm.ai/docs/providers/wandb_inference) | ✅ | ✅ | ✅ | | | | | | | | | [Watsonx Text (`watsonx_text`)](https://docs.litellm.ai/docs/providers/watsonx) | ✅ | ✅ | ✅ | | | | | | | | | [xAI (`xai`)](https://docs.litellm.ai/docs/providers/xai) | ✅ | ✅ | ✅ | | | | | | | | diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json deleted file mode 100644 index 92dc89eb0b8..00000000000 --- a/basedpyright-code-budget.json +++ /dev/null @@ -1,146 +0,0 @@ -{ - "reportAny": { - "limit": 13429 - }, - "reportArgumentType": { - "limit": 2198 - }, - "reportAssignmentType": { - "limit": 319 - }, - "reportAttributeAccessIssue": { - "limit": 480 - }, - "reportCallIssue": { - "limit": 112 - }, - "reportConstantRedefinition": { - "limit": 40 - }, - "reportDeprecated": { - "limit": 209 - }, - "reportDuplicateImport": { - "limit": 19 - }, - "reportExplicitAny": { - "limit": 3369 - }, - "reportFunctionMemberAccess": { - "limit": 7 - }, - "reportGeneralTypeIssues": { - "limit": 101 - }, - "reportIncompatibleMethodOverride": { - "limit": 56 - }, - "reportIncompatibleVariableOverride": { - "limit": 8 - }, - "reportInconsistentOverload": { - "limit": 12 - }, - "reportIndexIssue": { - "limit": 24 - }, - "reportInvalidTypeForm": { - "limit": 30 - }, - "reportInvalidTypeVarUse": { - "limit": 1 - }, - "reportMatchNotExhaustive": { - "limit": 0 - }, - "reportMissingParameterType": { - "limit": 5570 - }, - "reportMissingTypeArgument": { - "limit": 15281 - }, - "reportMissingTypeStubs": { - "limit": 40 - }, - "reportOperatorIssue": { - "limit": 0 - }, - "reportOptionalCall": { - "limit": 0 - }, - "reportOptionalIterable": { - "limit": 0 - }, - "reportOptionalMemberAccess": { - "limit": 0 - }, - "reportOptionalOperand": { - "limit": 0 - }, - "reportOptionalSubscript": { - "limit": 0 - }, - "reportPossiblyUnboundVariable": { - "limit": 56 - }, - "reportPrivateUsage": { - "limit": 1804 - }, - "reportRedeclaration": { - "limit": 8 - }, - "reportReturnType": { - "limit": 180 - }, - "reportTypedDictNotRequiredAccess": { - "limit": 22 - }, - "reportUndefinedVariable": { - "limit": 0 - }, - "reportUnknownArgumentType": { - "limit": 44802 - }, - "reportUnknownLambdaType": { - "limit": 109 - }, - "reportUnknownMemberType": { - "limit": 38269 - }, - "reportUnknownParameterType": { - "limit": 19584 - }, - "reportUnknownVariableType": { - "limit": 29814 - }, - "reportUnnecessaryCast": { - "limit": 110 - }, - "reportUnnecessaryComparison": { - "limit": 687 - }, - "reportUnnecessaryContains": { - "limit": 4 - }, - "reportUnnecessaryIsInstance": { - "limit": 816 - }, - "reportUntypedBaseClass": { - "limit": 0 - }, - "reportUntypedFunctionDecorator": { - "limit": 27 - }, - "reportUnusedClass": { - "limit": 21 - }, - "reportUnusedFunction": { - "limit": 136 - }, - "reportUnusedImport": { - "limit": 542 - }, - "reportUnusedVariable": { - "limit": 137 - } -} diff --git a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json index 4fb926a9658..996e440d77d 100644 --- a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json +++ b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json @@ -6436,6 +6436,133 @@ ], "title": "litellm_zero_cost_requests rate", "type": "timeseries" + }, + { + "collapsed": false, + "gridPos": { + "h": 1, + "w": 24, + "x": 0, + "y": 438 + }, + "id": 112, + "panels": [], + "title": "Project model rate limits", + "type": "row" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "description": "Configured rate limit for the Project on the requested model in the current window, by rate_limit_type", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "custom": { + "drawStyle": "line", + "fillOpacity": 10, + "lineWidth": 1, + "showPoints": "never", + "spanNulls": false + }, + "unit": "short" + }, + "overrides": [] + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 439 + }, + "id": 113, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "max by (project_id, project_alias, requested_model, rate_limit_type) (litellm_project_model_rate_limit_allowed_metric)", + "legendFormat": "{{project_alias}} ({{project_id}}) / {{requested_model}} / {{rate_limit_type}}", + "range": true, + "refId": "A" + } + ], + "title": "litellm_project_model_rate_limit_allowed_metric", + "type": "timeseries" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "description": "Requests or tokens the Project has consumed on the requested model in the current rate limit window, by rate_limit_type", + "fieldConfig": { + "defaults": { + "color": { + "mode": "palette-classic" + }, + "custom": { + "drawStyle": "line", + "fillOpacity": 10, + "lineWidth": 1, + "showPoints": "never", + "spanNulls": false + }, + "unit": "short" + }, + "overrides": [] + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 439 + }, + "id": 114, + "options": { + "legend": { + "calcs": [], + "displayMode": "list", + "placement": "bottom", + "showLegend": true + }, + "tooltip": { + "mode": "multi", + "sort": "desc" + } + }, + "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "max by (project_id, project_alias, requested_model, rate_limit_type) (litellm_project_model_rate_limit_used_metric)", + "legendFormat": "{{project_alias}} ({{project_id}}) / {{requested_model}} / {{rate_limit_type}}", + "range": true, + "refId": "A" + } + ], + "title": "litellm_project_model_rate_limit_used_metric", + "type": "timeseries" } ], "preload": false, diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 84e319c8c37..2b2d0a82431 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -732,6 +732,7 @@ class CheckBatchCost: another pod claimed it. Raises on results-fetch or cost-computation failures so the caller can leave the job unprocessed and retry it on a later poll. + """ from litellm.batches.batch_utils import ( count_error_file_failed_requests, @@ -815,15 +816,14 @@ class CheckBatchCost: custom_llm_provider=custom_llm_provider, ) - # CheckBatchCost bypasses async_post_call_success_hook, so convert raw - # output/error file IDs to managed base64 IDs before the DB write here. - managed_files_hook = self.proxy_logging_obj.get_proxy_hook("managed_files") - if managed_files_hook is not None: + from litellm.proxy.openai_files_endpoints.common_utils import ManagedBatchOutputFileWriter + + managed_files_hook: Final = self.proxy_logging_obj.get_proxy_hook("managed_files") + if isinstance(managed_files_hook, ManagedBatchOutputFileWriter): + managed_file_writer: Final = managed_files_hook from litellm.proxy._types import UserAPIKeyAuth - managed_file_model_name = self._get_managed_file_model_name( - job=job, deployment_info=deployment_info - ) + managed_file_model_name = self._get_managed_file_model_name(job=job, deployment_info=deployment_info) _minimal_auth = UserAPIKeyAuth( user_id=job.created_by or "default-user-id", team_id=getattr(job, "team_id", None), @@ -832,17 +832,19 @@ class CheckBatchCost: _raw_file_id = cast(str | None, getattr(response, _file_attr, None)) if _raw_file_id and not _is_base64_encoded_unified_file_id(_raw_file_id): try: - _unified_file_id = managed_files_hook.get_unified_output_file_id( + _unified_file_id = managed_file_writer.get_unified_output_file_id( output_file_id=_raw_file_id, model_id=model_id, model_name=managed_file_model_name, ) - await managed_files_hook.store_unified_file_id( - file_id=_unified_file_id, - file_object=None, + await managed_file_writer.store_batch_output_file( + unified_file_id=_unified_file_id, + provider_file_id=_raw_file_id, + model_id=model_id, + model_name=managed_file_model_name, + owner=_minimal_auth, litellm_parent_otel_span=None, - model_mappings={model_id: _raw_file_id}, - user_api_key_dict=_minimal_auth, + size_bytes=len(content_bytes) if _file_attr == "output_file_id" else None, ) setattr(response, _file_attr, _unified_file_id) verbose_proxy_logger.info( diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py index 738b9ebe1e0..7dd1fd7a1ee 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -6,7 +6,7 @@ same route are non-inference and free. """ from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Dict, Final, Optional, Protocol, cast +from typing import TYPE_CHECKING, Dict, Final, Protocol, cast import litellm from litellm._logging import verbose_proxy_logger @@ -47,6 +47,10 @@ def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_Manag return ManagedObjectRepository(prisma_client).table +def _is_response_gone_at_provider(error: Exception, provider_response_id: str) -> bool: + return getattr(error, "status_code", None) == 404 and provider_response_id in str(error) + + class CheckResponsesCost: def __init__( self, @@ -61,10 +65,15 @@ class CheckResponsesCost: self.prisma_client: PrismaClient = prisma_client self.llm_router: Router = llm_router + def _resolve_deployment(self, response_id: str) -> bool: + model_id: str | None = ResponsesAPIRequestUtils.get_model_id_from_response_id(response_id) + return model_id is not None and self.llm_router.get_deployment(model_id=model_id) is not None + async def _get_response( self, response_id: str, litellm_metadata: Dict[str, str], + via_router: bool, ) -> ResponsesAPIResponse: """Fetch the upstream response, using deployment credentials when available. @@ -75,8 +84,7 @@ class CheckResponsesCost: sees provider env vars, so it fails for every deployment whose credentials live in the config; the row then never leaves ``queued``. """ - model_id: Optional[str] = ResponsesAPIRequestUtils.get_model_id_from_response_id(response_id) - if model_id is None or self.llm_router.get_deployment(model_id=model_id) is None: + if not via_router: return await litellm.aget_responses(response_id=response_id, litellm_metadata=litellm_metadata) router_response = await self.llm_router.aget_responses( response_id=response_id, litellm_metadata=litellm_metadata @@ -140,6 +148,8 @@ class CheckResponsesCost: - Cost is tracked by the get-responses call, billed because the poll is stamped with BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN - Mark responses in a terminal state as complete in the database + - Mark responses the provider no longer has (404 through a resolved + deployment) as stale_expired """ try: await self._cleanup_stale_managed_objects() @@ -159,6 +169,7 @@ class CheckResponsesCost: verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check") completed_jobs: Final[list[_ManagedObjectRow]] = [] + expired_jobs: Final[list[_ManagedObjectRow]] = [] for job in jobs: unified_object_id = job.unified_object_id @@ -171,31 +182,48 @@ class CheckResponsesCost: # Get the stored response object to extract model information stored_response = job.file_object model_name = stored_response.get("model", None) - + # Decrypt the response ID responses_id_security, _, _ = ResponsesIDSecurity()._decrypt_response_id(unified_object_id) - + # Prepare metadata with model information for cost tracking litellm_metadata = { "user_api_key_user_id": job.created_by or "default-user-id", INTERNAL_CALL_ORIGIN_METADATA_KEY: BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, } - + # Add model information if available if model_name: litellm_metadata["model"] = model_name litellm_metadata["model_group"] = model_name # Use same value for model_group - + via_router = self._resolve_deployment(responses_id_security) + except Exception as e: + verbose_proxy_logger.warning( + f"Skipping job {unified_object_id} due to error: {e}" + ) + continue + + provider_response_id = ResponsesAPIRequestUtils.decode_responses_api_response_id( + responses_id_security + ).get("response_id", responses_id_security) + try: response = await self._get_response( response_id=responses_id_security, litellm_metadata=litellm_metadata, + via_router=via_router, ) - + verbose_proxy_logger.debug( f"Response {unified_object_id} status: {response.status}, model: {model_name}" ) - + except Exception as e: + if via_router and _is_response_gone_at_provider(e, provider_response_id): + verbose_proxy_logger.info( + f"Response {unified_object_id} no longer available at provider (404), marking stale_expired: {e}" + ) + expired_jobs.append(job) + continue verbose_proxy_logger.warning( f"Skipping job {unified_object_id} due to error: {e}" ) @@ -210,6 +238,7 @@ class CheckResponsesCost: # Mark completed jobs in the database if len(completed_jobs) > 0: await _managed_object_table(self.prisma_client).update_many( + # bounded-ok: at most MAX_OBJECTS_PER_POLL_CYCLE rows per cycle, the find_many take above where={"id": {"in": [job.id for job in completed_jobs]}}, data={"status": "completed"}, ) @@ -217,3 +246,13 @@ class CheckResponsesCost: f"Marked {len(completed_jobs)} response jobs as completed" ) + if len(expired_jobs) > 0: + await _managed_object_table(self.prisma_client).update_many( + # bounded-ok: at most MAX_OBJECTS_PER_POLL_CYCLE rows per cycle, the find_many take above + where={"id": {"in": [job.id for job in expired_jobs]}}, + data={"status": "stale_expired"}, + ) + verbose_proxy_logger.info( + f"Marked {len(expired_jobs)} response jobs as stale_expired" + ) + diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index d0149831026..906938b2f1a 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -1,9 +1,11 @@ # What is this? ## This hook is used to check for LiteLLM managed files in the request body, and replace them with model-specific file id +import asyncio import base64 import json -from collections.abc import Iterator, Mapping, Sequence +import time +from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence from types import MappingProxyType from typing import ( TYPE_CHECKING, @@ -23,21 +25,25 @@ from uuid import NAMESPACE_URL, uuid5 import httpx from fastapi import HTTPException from pydantic import ValidationError -from typing_extensions import ReadOnly +from typing_extensions import ReadOnly, Unpack import litellm from litellm import Router, verbose_logger from litellm._internal_context import with_service_target from litellm._uuid import uuid from litellm.caching.caching import DualCache -from litellm.constants import MAX_FILE_LIST_LIMIT +from litellm.constants import ( + BATCH_OUTPUT_FILE_FALLBACK_RETRY_AFTER_SECONDS, + BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS, + MAX_FILE_LIST_LIMIT, +) from litellm.files.types import FileRetrieveCallOptions, FileRetrieveProvider from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.hidden_params import HIDDEN_PARAMS_ATTR from litellm.litellm_core_utils.prompt_templates.common_utils import ( extract_file_metadata, ) -from openai import AsyncOpenAI +from openai import APIConnectionError, AsyncOpenAI from openai.types.file_deleted import FileDeleted from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend @@ -55,7 +61,10 @@ from litellm.proxy._types import ( ProxyException, UserAPIKeyAuth, ) -from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup +from litellm.proxy.litellm_pre_call_utils import ( + LiteLLMProxyRequestSetup, + sanitize_for_log, +) from litellm.proxy.openai_files_endpoints.common_utils import ( BATCH_CREATE_HIDDEN_PARAM, FILE_LIST_CONTINUATION_CHUNK_SIZE, @@ -77,6 +86,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import ( request_tags_from_metadata, ) +from litellm.repositories.managed_file_repository import ManagedFileRepository from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccessIssue] AllMessageValues, AsyncCursorPage, @@ -96,6 +106,17 @@ from litellm.types.utils import ( SpecialEnums, ) + +class _ManagedFileRetrieve(Protocol): + async def __call__( + self, + *, + file_id: str, + _litellm_internal_model_credentials: Mapping[str, object] | None = None, + **kwargs: Unpack[FileRetrieveCallOptions], + ) -> OpenAIFileObject: ... + + if TYPE_CHECKING: from opentelemetry.trace import Span as _Span from prisma.models import ( @@ -146,6 +167,91 @@ def _parse_managed_file_object(raw_file_object: object, unified_file_id: str) -> return None +def _batch_output_file_object( + unified_file_id: str, raw_file_id: str, size_bytes: int, *, fallback: bool +) -> OpenAIFileObject: + filename: Final = raw_file_id.rsplit("/", 1)[-1] or raw_file_id + return OpenAIFileObject( + id=unified_file_id, + object="file", + purpose="batch_output", + filename=filename, + created_at=int(time.time()), + bytes=size_bytes, + status="processed", + litellm_details_fallback=True if fallback else None, + ) + + +def _public_file_object(file_object: OpenAIFileObject, unified_file_id: str) -> OpenAIFileObject: + return file_object.model_copy(update={"id": unified_file_id, "litellm_details_fallback": None}) + + +def _is_transient_file_retrieve_error(error: Exception) -> bool: + status_code: Final = getattr(error, "status_code", None) + if isinstance(status_code, int): + return status_code in {408, 429} or status_code >= 500 + return isinstance(error, (httpx.TransportError, APIConnectionError, asyncio.TimeoutError)) + + +def _proxy_llm_router() -> Router | None: + import litellm.proxy.proxy_server as proxy_server_module + + return cast(Router | None, getattr(proxy_server_module, "llm_router", None)) + + +def _provider_file_retrieve_credentials( + *, + llm_router: Router | None, + model_id: str | None, +) -> Mapping[str, object] | None: + if llm_router is None or model_id is None: + return None + try: + credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id) + except Exception as error: + verbose_logger.warning( + "Failed to retrieve credentials for provider file " + f"model_id={sanitize_for_log(model_id)}: {sanitize_for_log(error)}" + ) + return None + return cast(Mapping[str, object], credentials) if credentials else None + + +_PROVIDER_FILE_RETRIEVE_PROVIDERS: Final[frozenset[str]] = frozenset( + { + "openai", + "azure", + "gemini", + "vertex_ai", + "bedrock", + "hosted_vllm", + "litellm_proxy", + "manus", + "anthropic", + "mistral", + "xai", + } +) + + +def _model_name_file_retrieve_provider(model_name: str | None) -> FileRetrieveProvider | None: + if model_name is None: + return None + provider, separator, _ = model_name.partition("/") + if not separator or provider not in _PROVIDER_FILE_RETRIEVE_PROVIDERS: + return None + return cast(FileRetrieveProvider, provider) + + +def _has_provider_file_retrieve_route( + *, + router_credentials: Mapping[str, object] | None, + model_name: str | None, +) -> bool: + return bool(router_credentials) or _model_name_file_retrieve_provider(model_name) is not None + + class _ManagedFileRow(Protocol): unified_file_id: str file_object: OpenAIFileObject @@ -161,6 +267,8 @@ class _ManagedFileRow(Protocol): class _ManagedFileTableActions(Protocol): async def find_first(self, where: Mapping[str, object]) -> Optional[_ManagedFileRow]: ... + async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ... + async def find_many( self, where: Mapping[str, object], @@ -253,13 +361,21 @@ def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, s _MANAGED_FILES_TARGET: Final = "managed_files" +_PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS: Final = (0.5, 1.0, 2.0) class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): # Class variables or attributes - def __init__(self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient): + def __init__( + self, + internal_usage_cache: InternalUsageCache, + prisma_client: PrismaClient, + *, + sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, + ): self.internal_usage_cache = internal_usage_cache self.prisma_client = prisma_client + self._sleep = sleep @staticmethod def _get_prometheus_logger(): @@ -595,14 +711,231 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): for file_id in provider_file_ids: model_name = decode_model_from_file_id(file_id) raw_file_id = get_original_file_id(file_id) - await self.store_unified_file_id( - file_id=file_id, - file_object=None, + await self.store_batch_output_file( + unified_file_id=file_id, + provider_file_id=raw_file_id, + model_id=model_name or None, + model_name=model_name, litellm_parent_otel_span=litellm_parent_otel_span, - model_mappings={model_name: raw_file_id} if model_name else {}, - user_api_key_dict=owner_identity, + owner=owner_identity, ) + async def _afile_retrieve_with_retries( + self, + *, + provider_file_id: str, + call_options: Mapping[str, object], + internal_model_credentials: Mapping[str, object] | None = None, + ) -> OpenAIFileObject: + """Retrieve provider file details with transient retries and SDK retries disabled.""" + retrieve_file: Final = cast(_ManagedFileRetrieve, litellm.afile_retrieve) + retrieve_options: Final = cast( + FileRetrieveCallOptions, + {**call_options, "max_retries": 0}, + ) + for attempt in range(len(_PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS) + 1): + if attempt > 0: + await self._sleep(_PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS[attempt - 1]) + try: + if internal_model_credentials is None: + return await retrieve_file( + file_id=provider_file_id, + **retrieve_options, + ) + return await retrieve_file( + file_id=provider_file_id, + _litellm_internal_model_credentials=MappingProxyType(dict(internal_model_credentials)), + **retrieve_options, + ) + except Exception as error: + if not _is_transient_file_retrieve_error(error) or attempt == len( + _PROVIDER_FILE_RETRIEVE_RETRY_DELAYS_SECONDS + ): + raise + raise RuntimeError("Provider file retrieve retry loop ended without a result") + + async def _fetch_provider_file_object( + self, + *, + unified_file_id: str, + provider_file_id: str, + model_id: str | None, + model_name: str | None, + llm_router: Router | None = None, + raise_on_failure: bool = False, + allow_default_provider: bool = False, + ) -> tuple[OpenAIFileObject | None, bool]: + """Fetch provider file details through the configured route under a total timeout.""" + route_llm_router: Final = llm_router if llm_router is not None else _proxy_llm_router() + router_credentials: Final = _provider_file_retrieve_credentials( + llm_router=route_llm_router, + model_id=model_id, + ) + model_name_provider: Final = _model_name_file_retrieve_provider(model_name) + fetch_route_available: Final = _has_provider_file_retrieve_route( + router_credentials=router_credentials, + model_name=model_name, + ) + default_provider_route: Final = allow_default_provider and route_llm_router is not None + if not fetch_route_available and not default_provider_route: + return None, False + + try: + if router_credentials is not None: + provider_file_object: Final = await asyncio.wait_for( + self._afile_retrieve_with_retries( + provider_file_id=provider_file_id, + call_options=router_credentials, + internal_model_credentials=router_credentials, + ), + timeout=BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS, + ) + return provider_file_object.model_copy(update={"id": unified_file_id}), True + if model_name_provider is not None: + provider_file_object_by_model_name: Final = await asyncio.wait_for( + self._afile_retrieve_with_retries( + provider_file_id=provider_file_id, + call_options={"custom_llm_provider": model_name_provider}, + ), + timeout=BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS, + ) + return ( + provider_file_object_by_model_name.model_copy(update={"id": unified_file_id}), + True, + ) + if default_provider_route: + provider_file_object_by_default_route: Final = await asyncio.wait_for( + self._afile_retrieve_with_retries( + provider_file_id=provider_file_id, + call_options=router_credentials or {}, + ), + timeout=BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS, + ) + return ( + provider_file_object_by_default_route.model_copy(update={"id": unified_file_id}), + True, + ) + return None, False + except Exception as error: + verbose_logger.warning( + "Failed to retrieve batch file object for " + f"provider_file_id={sanitize_for_log(provider_file_id)}: " + f"{type(error).__name__} {sanitize_for_log(error)}" + ) + if raise_on_failure: + if isinstance(error, TimeoutError) and not str(error): + raise TimeoutError( + "Provider file retrieve timed out " + f"after {BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS} seconds" + ) from error + raise + return None, True + + async def _save_refreshed_file_object( + self, + stored: LiteLLM_ManagedFileTable, + file_object: OpenAIFileObject, + ) -> None: + if not await ManagedFileRepository(self.prisma_client).update_file_object(stored.unified_file_id, file_object): + return + refreshed_row: Final = stored.model_copy(update={"file_object": file_object}) + await self.internal_usage_cache.async_set_cache( + key=stored.unified_file_id, + value=refreshed_row.model_dump(), + litellm_parent_otel_span=None, + ) + + async def store_batch_output_file( + self, + *, + unified_file_id: str, + provider_file_id: str, + model_id: str | None, + model_name: str | None = None, + owner: UserAPIKeyAuth, + litellm_parent_otel_span: Span | None, + size_bytes: int | None = None, + fetch_provider_details: bool = True, + ) -> None: + """Register batch output or error file metadata, optionally fetching provider details.""" + stored_file: Final = await self.get_unified_file_id(unified_file_id, litellm_parent_otel_span) + stored_object: Final = stored_file.file_object if stored_file is not None else None + if not fetch_provider_details: + if stored_file is not None: + return + router_credentials: Final = _provider_file_retrieve_credentials( + llm_router=_proxy_llm_router(), + model_id=model_id, + ) + file_object_without_provider_details: Final = _batch_output_file_object( + unified_file_id, + provider_file_id, + size_bytes or 0, + fallback=_has_provider_file_retrieve_route( + router_credentials=router_credentials, + model_name=model_name, + ), + ) + await self.store_unified_file_id( + file_id=unified_file_id, + file_object=file_object_without_provider_details, + litellm_parent_otel_span=litellm_parent_otel_span, + model_mappings={model_id: provider_file_id} if model_id else {}, + user_api_key_dict=owner, + ) + return + if stored_object is not None and not stored_object.litellm_details_fallback: + return + + fallback_written_recently: Final = ( + stored_object is not None + and time.time() - stored_object.created_at < BATCH_OUTPUT_FILE_FALLBACK_RETRY_AFTER_SECONDS + ) + provider_fetch_result: Final = ( + (None, True) + if fallback_written_recently + else await self._fetch_provider_file_object( + unified_file_id=unified_file_id, + provider_file_id=provider_file_id, + model_id=model_id, + model_name=model_name, + ) + ) + provider_object, fetch_route_available = provider_fetch_result + if ( + stored_file is not None + and provider_object is None + and (size_bytes is None or (stored_object is not None and stored_object.bytes == size_bytes)) + ): + return + + file_object: Final = ( + provider_object + if provider_object is not None + else ( + stored_object.model_copy(update={"bytes": size_bytes}) + if stored_object is not None and size_bytes is not None + else _batch_output_file_object( + unified_file_id, + provider_file_id, + size_bytes or 0, + fallback=fetch_route_available, + ) + ) + ) + + if stored_file is not None: + await self._save_refreshed_file_object(stored_file, file_object) + return + + await self.store_unified_file_id( + file_id=unified_file_id, + file_object=file_object, + litellm_parent_otel_span=litellm_parent_otel_span, + model_mappings={model_id: provider_file_id} if model_id else {}, + user_api_key_dict=owner, + ) + async def list_user_batches( self, user_api_key_dict: UserAPIKeyAuth, @@ -745,6 +1078,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): user_api_key_dict=user_api_key_dict, db_batch_object=row, unified_batch_id=_is_base64_encoded_unified_file_id(row.unified_object_id), + fetch_provider_details=False, ) except Exception as e: verbose_logger.warning(f"Failed to resolve managed file ids for batch {row.unified_object_id}: {e}") @@ -806,7 +1140,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): } ) return [ - parsed_file_object.model_copy(update={"id": row.unified_file_id}) + _public_file_object(parsed_file_object, row.unified_file_id) for row in file_ids if (parsed_file_object := _parse_managed_file_object(row.file_object, row.unified_file_id)) is not None ] @@ -1474,42 +1808,13 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) setattr(response, file_attr, unified_file_id) - # Use llm_router credentials when available. Without credentials, - # Azure and other auth-required providers return 500/401. - file_object = None - try: - # Import module and use getattr for better testability with mocks - import litellm.proxy.proxy_server as proxy_server_module - - _llm_router = getattr(proxy_server_module, "llm_router", None) - if _llm_router is not None and model_id: - _creds = _llm_router.get_deployment_credentials_with_provider(model_id) or {} - file_object = await litellm.afile_retrieve( - file_id=provider_file_id, - **cast(FileRetrieveCallOptions, _creds), - ) - else: - file_object = await litellm.afile_retrieve( - custom_llm_provider=cast( - FileRetrieveProvider, - model_name.split("/")[0] if model_name and "/" in model_name else "openai", - ), - file_id=provider_file_id, - ) - verbose_logger.debug( - f"Successfully retrieved file object for {file_attr}={provider_file_id}" - ) - except Exception as e: - verbose_logger.warning( - f"Failed to retrieve file object for {file_attr}={provider_file_id}: {str(e)}. Storing with None and will fetch on-demand." - ) - - await self.store_unified_file_id( - file_id=unified_file_id, - file_object=file_object, + await self.store_batch_output_file( + unified_file_id=unified_file_id, + provider_file_id=provider_file_id, + model_id=model_id, + model_name=resolved_model_name, + owner=user_api_key_dict, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, - model_mappings={model_id: provider_file_id}, - user_api_key_dict=user_api_key_dict, ) request_metadata: Final = data.get("litellm_metadata") await self.store_unified_object_id( @@ -1612,6 +1917,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): async def afile_retrieve( self, file_id: str, litellm_parent_otel_span: Optional[Span], llm_router: Optional[Router] = None ) -> OpenAIFileObject: + """Return public details for a managed file ID, refreshing a basic entry when possible.""" stored_file_object = await self.get_unified_file_id(file_id, litellm_parent_otel_span) # Case 1 : This is not a managed file @@ -1621,30 +1927,57 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): # Case 2: Managed file and the file object exists in the database # The stored file_object has the raw provider ID. Replace with the unified ID # so callers see a consistent ID (matching Case 3 which does response.id = file_id). - if stored_file_object and stored_file_object.file_object: - # Use model_copy to ensure the ID update persists (Pydantic v2 compatibility) - response = stored_file_object.file_object.model_copy(update={"id": file_id}) - return response + if stored_file_object.file_object is not None: + file_object: Final = stored_file_object.file_object + if file_object.litellm_details_fallback and stored_file_object.model_mappings: + try: + model_id, provider_file_id = next(iter(stored_file_object.model_mappings.items())) + refreshed_file_object, _ = await self._fetch_provider_file_object( + unified_file_id=file_id, + provider_file_id=provider_file_id, + model_id=model_id, + model_name=model_id, + llm_router=llm_router, + ) + if refreshed_file_object is None: + return _public_file_object(file_object, file_id) + await self._save_refreshed_file_object(stored_file_object, refreshed_file_object) + return _public_file_object(refreshed_file_object, file_id) + except Exception as error: + verbose_logger.warning( + "Failed to refresh batch file object for " + f"file_id={sanitize_for_log(file_id)}: {sanitize_for_log(error)}" + ) + return _public_file_object(file_object, file_id) # Case 3: Managed file exists in the database but not the file object (for. e.g the batch task might not have run) # So we fetch the file object from the provider. We deliberately do not store the result to avoid interfering with batch cost tracking code. - if not llm_router: + model_mapping: Final = next(iter(stored_file_object.model_mappings.items()), None) + if model_mapping is None: raise Exception( - f"LiteLLM Managed File object with id={file_id} has no file_object " - f"and llm_router is required to fetch from provider" + f"LiteLLM Managed File object with id={file_id} has no file_object and no provider route to fetch it" ) + model_id, model_file_id = model_mapping try: - model_id, model_file_id = next(iter(stored_file_object.model_mappings.items())) - credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id) or {} - response = await litellm.afile_retrieve( - file_id=model_file_id, - **cast(FileRetrieveCallOptions, credentials), + response, fetch_route_available = await self._fetch_provider_file_object( + unified_file_id=file_id, + provider_file_id=model_file_id, + model_id=model_id, + model_name=model_id, + llm_router=llm_router, + raise_on_failure=True, + allow_default_provider=True, ) - response.id = file_id # Replace with unified ID - return response except Exception as e: raise Exception(f"Failed to retrieve file {file_id} from provider: {str(e)}") from e + if not fetch_route_available: + raise Exception( + f"LiteLLM Managed File object with id={file_id} has no file_object and no provider route to fetch it" + ) + if response is None: + raise ValueError("Provider file details could not be retrieved") + return _public_file_object(response, file_id) async def afile_list( self, @@ -1707,7 +2040,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): **cursor_args, ) matches.extend( - parsed_file_object.model_copy(update={"id": row.unified_file_id}) + _public_file_object(parsed_file_object, row.unified_file_id) for row in chunk if (parsed_file_object := _parse_managed_file_object(row.file_object, row.unified_file_id)) is not None and (purpose is None or parsed_file_object.purpose == purpose) diff --git a/enterprise/pyproject.toml b/enterprise/pyproject.toml index 08b049f0157..6b55abae873 100644 --- a/enterprise/pyproject.toml +++ b/enterprise/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-enterprise" -version = "0.1.74" +version = "0.1.75" description = "Package for LiteLLM Enterprise features" readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.1.74" +version = "0.1.75" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-enterprise==", diff --git a/litellm-proxy-extras/litellm_proxy_extras/utils.py b/litellm-proxy-extras/litellm_proxy_extras/utils.py index 1fd292b8137..32ef52680e6 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/utils.py +++ b/litellm-proxy-extras/litellm_proxy_extras/utils.py @@ -11,7 +11,7 @@ import time from collections.abc import Callable from dataclasses import dataclass, replace from pathlib import Path -from typing import TYPE_CHECKING, Final, Optional, Union +from typing import TYPE_CHECKING, Final, Optional, Protocol, Union from urllib.parse import unquote, urlsplit from litellm_proxy_extras import prisma_toolchain @@ -35,6 +35,12 @@ if TYPE_CHECKING: import psycopg.sql +class LensCheckConnect(Protocol): + def __call__( + self, conninfo: str, *, connect_timeout: int, autocommit: bool + ) -> "psycopg.Connection[tuple[object, ...]]": ... + + def str_to_bool(value: Optional[str]) -> bool: if value is None: return False @@ -243,11 +249,9 @@ def _redact_credentials(text: str) -> str: passwords: Final = sorted(_configured_database_passwords(), key=len, reverse=True) alternation: Final = "|".join(re.escape(password) for password in passwords) password_pattern: Final = ( - re.compile(rf"(?P:|password=)(?:{alternation})(?=@|&|$|[\s'\"\]),])", re.IGNORECASE) - if passwords - else None + re.compile(rf"(?{_REDACTED}", text) if password_pattern is not None else text + result: Final = password_pattern.sub(_REDACTED, text) if password_pattern is not None else text return _secret_shape_redactor()(result) @@ -690,7 +694,7 @@ class ProxyExtrasDBManager: ) @staticmethod - def raise_if_lens_rename_pending() -> None: + def raise_if_lens_rename_pending(connect: LensCheckConnect | None = None) -> None: database_url: Final = os.environ.get("DATABASE_URL") if not database_url: return @@ -698,8 +702,9 @@ class ProxyExtrasDBManager: import psycopg except ImportError as exc: raise RuntimeError("Install psycopg to verify Lens data safety before prisma db push.") from exc + open_connection: Final = connect if connect is not None else psycopg.connect try: - with psycopg.connect( + with open_connection( ProxyExtrasDBManager._strip_prisma_query_params(database_url), connect_timeout=10, autocommit=True ) as connection: legacy: Final = connection.execute( @@ -710,7 +715,8 @@ class ProxyExtrasDBManager: ).fetchone() except psycopg.Error as exc: raise RuntimeError( - "Cannot verify Lens data safety; refusing prisma db push. Check database connectivity and psycopg installation." + "Cannot verify Lens data safety; refusing prisma db push. " + f"The database check failed: {_redact_credentials(str(exc)).strip()}" ) from exc if legacy is not None: raise RuntimeError( @@ -1152,7 +1158,7 @@ class ProxyExtrasDBManager: if next_budget.attempts_left < budget.attempts_left: time.sleep(random.randrange(5, 15)) - budget = next_budget # rebind-ok: the loop carries the budget from one migrate deploy pass to the next + budget = next_budget raise RuntimeError( f"Database migration failed after {MAX_MIGRATE_DEPLOY_ATTEMPTS} " diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index 00984dbbf82..221d5012360 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-proxy-extras" -version = "0.4.106" +version = "0.4.107" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." readme = "README.md" requires-python = ">=3.9" @@ -30,7 +30,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.4.106" +version = "0.4.107" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-proxy-extras==", diff --git a/litellm-rust/.agents/skills/rust-string-enums/SKILL.md b/litellm-rust/.agents/skills/rust-string-enums/SKILL.md index fc6988639d0..5fc067790a0 100644 --- a/litellm-rust/.agents/skills/rust-string-enums/SKILL.md +++ b/litellm-rust/.agents/skills/rust-string-enums/SKILL.md @@ -35,7 +35,7 @@ pub enum EventType { } ``` -Use `#[strum(serialize_all = "snake_case")]` or another supported case style when it exactly matches the contract. Use explicit variant spellings otherwise. With multiple accepted aliases, set `to_string` to the existing canonical output: Strum Display otherwise selects the longest `serialize` spelling +Spell each variant explicitly with `#[strum(serialize = "...")]`; do not use `serialize_all`. With multiple accepted aliases, set `to_string` to the existing canonical output: Strum Display otherwise selects the longest `serialize` spelling Use only the derives the contract needs. A deserialize-only type should remain deserialize-only. Do not add an unknown variant to a closed enum, or derive Serde for a type that currently has no serialization contract @@ -43,7 +43,7 @@ Plain Serde derives with `rename` or `rename_all` remain appropriate for closed ## Preserve behavior during migration -Read the type, its callers, and existing tests before changing it. Preserve canonical output, accepted aliases, case sensitivity, whitespace handling, unknown values, malformed-input rejection, and existing public conversion APIs. Keep conversions needed by callers or compatibility even when Serde no longer uses them +Read the type, its callers, and existing tests before changing it. Preserve canonical output, accepted aliases, case sensitivity, whitespace handling, unknown values, malformed-input rejection, and existing public conversion contracts. Wrappers that only delegate to a Strum-derived trait are removed rather than kept: callers use `FromStr` and `From for &'static str` directly. Keep conversions that add behavior needed by callers or compatibility even when Serde no longer uses them Use the workspace dependencies and enable `serde_with.workspace = true` in a crate only when needed. Check the versions and enabled features in `Cargo.toml` and `Cargo.lock` rather than upgrading dependencies for this refactor diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index 442dfd8e957..c22f9575f74 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -4,6 +4,10 @@ For diagnostic tracing changes, follow [.agents/skills/rust-tracing/SKILL.md](.a For string-valued enums and their Serde conversions, follow [.agents/skills/rust-string-enums/SKILL.md](.agents/skills/rust-string-enums/SKILL.md) +For fieldless enums, derive `strum::VariantArray` and use `VARIANTS` instead of a hand-listed `ALL` array; derive Strum string conversions instead of hand-written variant-to-string matches + +Use the derived conversions directly (`<&'static str>::from(x)` / `.into()`, `str::parse`) with no `as_str`/`parse` wrapper that only delegates, and spell each variant with explicit `#[strum(serialize = "...")]` instead of `serialize_all` + ## Test placement - Never create a `tests.rs` (or `test.rs`) file under `src/`, and never `#[path = "tests.rs"] mod tests;` @@ -17,6 +21,9 @@ For string-valued enums and their Serde conversions, follow [.agents/skills/rust Use [`#[rstest]`](https://docs.rs/rstest/latest/rstest/attr.rstest.html) for new and updated tests and [`#[fixture]`](https://docs.rs/rstest/latest/rstest/attr.fixture.html) for reusable setup, injected through typed test arguments. Express input variations as named `#[case::name(...)]` cases instead of loops or duplicated tests so each failure identifies its case. Keep behavior assertions in the test body and fixtures focused on setup. Use the workspace `rstest` dependency +- Never loop over inputs (`for`, `.iter().for_each`, `.all`) inside a test body; give each input its own `#[case::name(...)]`, or use `#[values(...)]` for a cross product +- Exception: a test pinning a Rust table against a repo-owned data file (for example `include_str!` of a JSON config) may iterate that file's entries + ## Error definitions - A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs` diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 987491d514b..9d511774b04 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3485,6 +3485,7 @@ dependencies = [ "rstest", "serde", "serde_json", + "strum", "thiserror 2.0.19", "tokio", ] @@ -3692,7 +3693,9 @@ dependencies = [ "litellm-auth-types", "rstest", "serde", + "serde_json", "serde_yaml_ng", + "strum", "tempfile", "thiserror 2.0.19", ] @@ -3990,6 +3993,7 @@ dependencies = [ "rustls-native-certs", "serde", "serde_json", + "strum", "tempfile", "thiserror 2.0.19", "tokio", @@ -4248,6 +4252,7 @@ dependencies = [ name = "litellm-llms-types" version = "0.1.0" dependencies = [ + "indexmap 2.14.0", "macro_rules_attribute", "rstest", "schemars 1.2.2", @@ -4538,6 +4543,7 @@ dependencies = [ "serde", "serde_json", "sha2 0.10.9", + "strum", "tar", "target-lexicon", "tempfile", @@ -6512,6 +6518,7 @@ checksum = "687274d293b6cdc6e73e0fee520bf2049650090d7164f87672d212a3c530cf4a" dependencies = [ "chrono", "dyn-clone", + "indexmap 2.14.0", "ref-cast", "schemars_derive 1.2.2", "serde", diff --git a/litellm-rust/crates/cache/Cargo.toml b/litellm-rust/crates/cache/Cargo.toml index f18dbd9cb26..51bbe12a85e 100644 --- a/litellm-rust/crates/cache/Cargo.toml +++ b/litellm-rust/crates/cache/Cargo.toml @@ -8,6 +8,7 @@ repository.workspace = true [dependencies] serde.workspace = true serde_json = { workspace = true, features = ["preserve_order"] } +strum.workspace = true thiserror.workspace = true [dev-dependencies] diff --git a/litellm-rust/crates/cache/src/cache_type.rs b/litellm-rust/crates/cache/src/cache_type.rs index 22d8c8c7cb5..80d040db487 100644 --- a/litellm-rust/crates/cache/src/cache_type.rs +++ b/litellm-rust/crates/cache/src/cache_type.rs @@ -1,57 +1,44 @@ use serde::{Deserialize, Serialize}; -#[derive(Clone, Copy, Debug, Deserialize, Serialize, PartialEq, Eq, Hash)] +#[derive( + Clone, + Copy, + Debug, + Deserialize, + Serialize, + PartialEq, + Eq, + Hash, + strum::EnumString, + strum::IntoStaticStr, + strum::VariantArray, +)] pub enum CacheType { #[serde(rename = "local")] + #[strum(serialize = "local")] Local, #[serde(rename = "redis")] + #[strum(serialize = "redis")] Redis, #[serde(rename = "redis-semantic")] + #[strum(serialize = "redis-semantic")] RedisSemantic, #[serde(rename = "valkey-semantic")] + #[strum(serialize = "valkey-semantic")] ValkeySemantic, #[serde(rename = "s3")] + #[strum(serialize = "s3")] S3, #[serde(rename = "disk")] + #[strum(serialize = "disk")] Disk, #[serde(rename = "qdrant-semantic")] + #[strum(serialize = "qdrant-semantic")] QdrantSemantic, #[serde(rename = "azure-blob")] + #[strum(serialize = "azure-blob")] AzureBlob, #[serde(rename = "gcs")] + #[strum(serialize = "gcs")] Gcs, } - -impl CacheType { - pub const ALL: [Self; 9] = [ - Self::Local, - Self::Redis, - Self::RedisSemantic, - Self::ValkeySemantic, - Self::S3, - Self::Disk, - Self::QdrantSemantic, - Self::AzureBlob, - Self::Gcs, - ]; - - pub const fn as_python_name(self) -> &'static str { - match self { - Self::Local => "local", - Self::Redis => "redis", - Self::RedisSemantic => "redis-semantic", - Self::ValkeySemantic => "valkey-semantic", - Self::S3 => "s3", - Self::Disk => "disk", - Self::QdrantSemantic => "qdrant-semantic", - Self::AzureBlob => "azure-blob", - Self::Gcs => "gcs", - } - } - - pub fn from_python_name(value: &str) -> Option { - Self::ALL - .into_iter() - .find(|cache_type| cache_type.as_python_name() == value) - } -} diff --git a/litellm-rust/crates/cache/tests/cache_type.rs b/litellm-rust/crates/cache/tests/cache_type.rs index 24aaba8e5fd..49b3e8b27ad 100644 --- a/litellm-rust/crates/cache/tests/cache_type.rs +++ b/litellm-rust/crates/cache/tests/cache_type.rs @@ -1,5 +1,6 @@ use litellm_cache::CacheType; use rstest::rstest; +use strum::VariantArray; #[rstest] #[case(CacheType::Local, "local")] @@ -15,16 +16,16 @@ fn every_python_cache_type_has_one_round_trip_identity( #[case] cache_type: CacheType, #[case] name: &str, ) { - assert_eq!(cache_type.as_python_name(), name); - assert_eq!(CacheType::from_python_name(name), Some(cache_type)); + assert_eq!(<&'static str>::from(cache_type), name); + assert_eq!(name.parse::().ok(), Some(cache_type)); assert_eq!( serde_json::to_value(cache_type).unwrap(), serde_json::Value::from(name) ); assert_eq!( - CacheType::ALL + CacheType::VARIANTS .iter() - .filter(|candidate| candidate.as_python_name() == name) + .filter(|candidate| <&'static str>::from(**candidate) == name) .count(), 1 ); @@ -33,7 +34,10 @@ fn every_python_cache_type_has_one_round_trip_identity( #[rstest] fn python_cache_types_are_listed_in_python_order() { assert_eq!( - CacheType::ALL.map(CacheType::as_python_name), + CacheType::VARIANTS + .iter() + .map(|cache_type| <&'static str>::from(*cache_type)) + .collect::>(), [ "local", "redis", @@ -52,5 +56,5 @@ fn python_cache_types_are_listed_in_python_order() { #[case::unknown("memcached")] #[case::case_sensitive("Redis")] fn unknown_python_names_have_no_cache_type(#[case] name: &str) { - assert_eq!(CacheType::from_python_name(name), None); + assert!(name.parse::().is_err()); } diff --git a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs index 70b7ecb45fd..dd3da8991cc 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs @@ -936,9 +936,6 @@ mod payload_tests { /// The payload phases of `Logging` on top of `StubLogger`, with `pre_call` handing the /// payload to the case's `on_pre_call`. const PAYLOAD_LOGGER: &CStr = c" -class Request: - pass - class PayloadLogger(StubLogger): def update_from_kwargs(self, **update): self.update = update @@ -954,7 +951,7 @@ class PayloadLogger(StubLogger): self.record('post_call', None) self.post = (original_response, api_key, additional_args) -request = Request() +bound = {} kwargs = {} logger = PayloadLogger() on_pre_call = lambda additional_args: None @@ -1225,10 +1222,10 @@ on_pre_call = lambda args: observed.append( def check(): assert observed == [(True, True)], observed ")] - #[case::request_attribute_behind_an_omitted_keyword(c" + #[case::bound_value_behind_an_omitted_keyword(c" document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'} pages = [0] -request.document = document +bound['document'] = document kwargs = {'pages': pages} observed = [] on_pre_call = lambda args: observed.append( diff --git a/litellm-rust/crates/callbacks-legacy-python/src/call.rs b/litellm-rust/crates/callbacks-legacy-python/src/call.rs index 8e7645cae7b..8ea67292ae7 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/call.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/call.rs @@ -13,21 +13,21 @@ use pyo3::{ pub struct PublicCall { args: Py, kwargs: Py, - request: Py, + bound: Py, } impl PublicCall { /// Copies the keyword arguments once, so the legacy path's rewrites never reach the /// caller's own dict while every value keeps its identity. pub fn capture( - request: &Bound<'_, PyAny>, + bound: &Bound<'_, PyDict>, args: &Bound<'_, PyTuple>, kwargs: &Bound<'_, PyDict>, ) -> PyResult { Ok(Self { args: args.clone().unbind(), kwargs: kwargs.copy()?.unbind(), - request: request.clone().unbind(), + bound: bound.clone().unbind(), }) } @@ -55,31 +55,30 @@ impl PublicCall { py: Python<'py>, name: &str, ) -> PyResult>> { - lookup(self.kwargs.bind(py), self.request.bind(py), name) + lookup(self.kwargs.bind(py), self.bound.bind(py), name) } pub(crate) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.args)?; visit.call(&self.kwargs)?; - visit.call(&self.request) + visit.call(&self.bound) } } #[cfg(test)] mod tests { use super::*; + use crate::test_support::{local, local_dict}; fn capture<'py>(py: Python<'py>, source: &std::ffi::CStr) -> (PublicCall, Bound<'py, PyDict>) { let locals = PyDict::new(py); py.run(source, Some(&locals), Some(&locals)).unwrap(); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap(); + let call = PublicCall::capture( + &local_dict(&locals, "bound"), + &PyTuple::empty(py), + &local_dict(&locals, "kwargs"), + ) + .unwrap(); (call, locals) } @@ -91,25 +90,22 @@ mod tests { py, c" pages = [0] -class Request: - pass -request = Request() +bound = {'pages': [1]} kwargs = {'pages': pages} ", ); - let caller = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); + let caller = local_dict(&locals, "kwargs"); call.kwargs() .bind(py) .set_item("litellm_call_id", "call") .unwrap(); assert!(!caller.contains("litellm_call_id").unwrap()); - let pages = locals.get_item("pages").unwrap().unwrap(); - assert!(call.lookup(py, "pages").unwrap().unwrap().is(&pages)); + assert!( + call.lookup(py, "pages") + .unwrap() + .unwrap() + .is(local(&locals, "pages")) + ); }); } } diff --git a/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs b/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs index f252a45b562..a7b28f1b0ca 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs @@ -173,21 +173,23 @@ pub(crate) fn local<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, locals.get_item(name).unwrap().unwrap() } -/// A legacy call over the namespace's `kwargs` (or none) and `request` (or `None`). +pub(crate) fn local_dict<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyDict> { + local(locals, name).cast_into().unwrap() +} + +/// A legacy call over the namespace's `kwargs` and `bound` dicts, each empty when absent. pub(crate) fn legacy_call( py: Python<'_>, locals: &Bound<'_, PyDict>, asynchronous: bool, ) -> LegacyLogging { - let request = locals - .get_item("request") - .unwrap() - .unwrap_or_else(|| py.None().into_bound(py)); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .map(|kwargs| kwargs.cast_into::().unwrap()) - .unwrap_or_else(|| PyDict::new(py)); - let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap(); + let dict = |name: &str| { + locals + .get_item(name) + .unwrap() + .map(|value| value.cast_into::().unwrap()) + .unwrap_or_else(|| PyDict::new(py)) + }; + let call = PublicCall::capture(&dict("bound"), &PyTuple::empty(py), &dict("kwargs")).unwrap(); LegacyLogging::new(py, crate::LoggingOperation::Ocr, call, asynchronous) } diff --git a/litellm-rust/crates/config/Cargo.toml b/litellm-rust/crates/config/Cargo.toml index 36bd68fe2a0..e23739997e2 100644 --- a/litellm-rust/crates/config/Cargo.toml +++ b/litellm-rust/crates/config/Cargo.toml @@ -9,8 +9,10 @@ repository.workspace = true litellm-auth-types.workspace = true serde.workspace = true serde_yaml_ng = "0.10.0" +strum.workspace = true thiserror.workspace = true [dev-dependencies] rstest.workspace = true +serde_json.workspace = true tempfile.workspace = true diff --git a/litellm-rust/crates/config/src/mcp.rs b/litellm-rust/crates/config/src/mcp.rs index 522744e4fe4..cbe992f9956 100644 --- a/litellm-rust/crates/config/src/mcp.rs +++ b/litellm-rust/crates/config/src/mcp.rs @@ -41,57 +41,43 @@ impl fmt::Debug for McpServer { } } -#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq, strum::IntoStaticStr)] #[serde(rename_all = "snake_case")] pub enum McpTransport { #[default] + #[strum(serialize = "http")] Http, + #[strum(serialize = "sse")] Sse, + #[strum(serialize = "stdio")] Stdio, } -impl McpTransport { - pub fn as_str(self) -> &'static str { - match self { - Self::Http => "http", - Self::Sse => "sse", - Self::Stdio => "stdio", - } - } -} - -#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq, strum::IntoStaticStr)] #[serde(rename_all = "snake_case")] pub enum McpAuth { + #[strum(serialize = "none")] None, + #[strum(serialize = "api_key")] ApiKey, + #[strum(serialize = "bearer_token")] BearerToken, + #[strum(serialize = "basic")] Basic, + #[strum(serialize = "authorization")] Authorization, + #[strum(serialize = "token")] Token, + #[strum(serialize = "oauth2")] Oauth2, + #[strum(serialize = "aws_sigv4")] AwsSigv4, + #[strum(serialize = "oauth2_token_exchange")] Oauth2TokenExchange, + #[strum(serialize = "oauth2_id_jag")] Oauth2IdJag, + #[strum(serialize = "true_passthrough")] TruePassthrough, + #[strum(serialize = "oauth_delegate")] OauthDelegate, } - -impl McpAuth { - pub fn as_str(self) -> &'static str { - match self { - Self::None => "none", - Self::ApiKey => "api_key", - Self::BearerToken => "bearer_token", - Self::Basic => "basic", - Self::Authorization => "authorization", - Self::Token => "token", - Self::Oauth2 => "oauth2", - Self::AwsSigv4 => "aws_sigv4", - Self::Oauth2TokenExchange => "oauth2_token_exchange", - Self::Oauth2IdJag => "oauth2_id_jag", - Self::TruePassthrough => "true_passthrough", - Self::OauthDelegate => "oauth_delegate", - } - } -} diff --git a/litellm-rust/crates/config/tests/mcp.rs b/litellm-rust/crates/config/tests/mcp.rs new file mode 100644 index 00000000000..9ac6ad02bbb --- /dev/null +++ b/litellm-rust/crates/config/tests/mcp.rs @@ -0,0 +1,39 @@ +use litellm_config::{McpAuth, McpTransport}; +use rstest::rstest; +use serde_json::json; + +#[rstest] +#[case::http(McpTransport::Http, "http")] +#[case::sse(McpTransport::Sse, "sse")] +#[case::stdio(McpTransport::Stdio, "stdio")] +fn mcp_transport_as_str_matches_the_serde_spelling( + #[case] transport: McpTransport, + #[case] name: &str, +) { + assert_eq!(<&'static str>::from(transport), name); + assert_eq!( + serde_json::from_value::(json!(name)).unwrap(), + transport + ); +} + +#[rstest] +#[case::none(McpAuth::None, "none")] +#[case::api_key(McpAuth::ApiKey, "api_key")] +#[case::bearer_token(McpAuth::BearerToken, "bearer_token")] +#[case::basic(McpAuth::Basic, "basic")] +#[case::authorization(McpAuth::Authorization, "authorization")] +#[case::token(McpAuth::Token, "token")] +#[case::oauth2(McpAuth::Oauth2, "oauth2")] +#[case::aws_sigv4(McpAuth::AwsSigv4, "aws_sigv4")] +#[case::oauth2_token_exchange(McpAuth::Oauth2TokenExchange, "oauth2_token_exchange")] +#[case::oauth2_id_jag(McpAuth::Oauth2IdJag, "oauth2_id_jag")] +#[case::true_passthrough(McpAuth::TruePassthrough, "true_passthrough")] +#[case::oauth_delegate(McpAuth::OauthDelegate, "oauth_delegate")] +fn mcp_auth_as_str_matches_the_serde_spelling(#[case] auth: McpAuth, #[case] name: &str) { + assert_eq!(<&'static str>::from(auth), name); + assert_eq!( + serde_json::from_value::(json!(name)).unwrap(), + auth + ); +} diff --git a/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs b/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs index e71ecdffa44..a2bb462c2a9 100644 --- a/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs +++ b/litellm-rust/crates/core-utils/src/get_llm_provider_logic.rs @@ -7,17 +7,26 @@ pub struct CustomLlmProvider<'a> { } #[derive(Debug, Clone, Copy, PartialEq, Eq, EnumString, IntoStaticStr)] -#[strum(serialize_all = "snake_case")] pub enum LlmProviders { + #[strum(serialize = "anthropic")] Anthropic, + #[strum(serialize = "aws_textract")] AwsTextract, + #[strum(serialize = "azure_ai")] AzureAi, + #[strum(serialize = "bedrock")] Bedrock, + #[strum(serialize = "cohere")] Cohere, + #[strum(serialize = "mistral")] Mistral, + #[strum(serialize = "openai")] Openai, + #[strum(serialize = "openai_like")] OpenaiLike, + #[strum(serialize = "reducto")] Reducto, + #[strum(serialize = "vertex_ai")] VertexAi, } diff --git a/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs b/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs index 10a0d719e9d..dcb98ed17f9 100644 --- a/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs +++ b/litellm-rust/crates/core-utils/src/prompt_templates/factory.rs @@ -17,18 +17,13 @@ pub const EMPTY_TEXT_PLACEHOLDER: &str = "[System: Empty message content sanitised to satisfy protocol]"; #[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq)] -#[strum(serialize_all = "snake_case")] pub enum TurnRole { + #[strum(serialize = "user")] User, + #[strum(serialize = "assistant")] Assistant, } -impl TurnRole { - pub fn as_str(self) -> &'static str { - self.into() - } -} - #[derive(Clone, Debug, PartialEq, Eq)] pub struct Turn { pub role: TurnRole, diff --git a/litellm-rust/crates/gateway-mcp/src/configured.rs b/litellm-rust/crates/gateway-mcp/src/configured.rs index 06a1779a274..bb4f5594f1f 100644 --- a/litellm-rust/crates/gateway-mcp/src/configured.rs +++ b/litellm-rust/crates/gateway-mcp/src/configured.rs @@ -49,8 +49,8 @@ fn info(name: &str, config: &McpServer) -> ServerInfo { let identity = format!( "{name}|{}|{}|{}|{}", config.url.as_ref().map_or("", SecretValue::expose), - config.transport.as_str(), - config.auth_type.map_or("", McpAuth::as_str), + <&'static str>::from(config.transport), + config.auth_type.map_or("", <&'static str>::from), config.alias.as_deref().unwrap_or("") ); ServerInfo { diff --git a/litellm-rust/crates/host-python/src/argument.rs b/litellm-rust/crates/host-python/src/argument.rs index 13214e0e9ed..78879036f7e 100644 --- a/litellm-rust/crates/host-python/src/argument.rs +++ b/litellm-rust/crates/host-python/src/argument.rs @@ -2,78 +2,162 @@ use pyo3::{prelude::*, types::PyDict}; pub fn lookup<'py>( kwargs: &Bound<'py, PyDict>, - request: &Bound<'py, PyAny>, + bound: &Bound<'py, PyDict>, name: &str, ) -> PyResult>> { - if let Some(value) = kwargs.get_item(name)? { - return Ok(Some(value)); + match kwargs.get_item(name)? { + Some(value) => Ok(Some(value)), + None => bound.get_item(name), } - if let Ok(bound) = request.cast::() { - return bound.get_item(name); - } - request.getattr_opt(name) +} + +pub fn present<'py>( + kwargs: &Bound<'py, PyDict>, + bound: &Bound<'py, PyDict>, + name: &str, +) -> PyResult>> { + Ok(lookup(kwargs, bound, name)?.filter(|value| !value.is_none())) +} + +/// What `original_function(*args, **kwargs)` sees: the signature base with the keyword +/// dict laid over it, so a rewritten keyword wins and a deleted keyword falls back to the +/// signature default. +pub fn effective_py_args<'py>( + base: &Bound<'py, PyDict>, + kwargs: &Bound<'py, PyDict>, +) -> PyResult> { + // Shallow copy, like the Python path: nested values stay shared with the caller. + let merged = base.copy()?; + merged.update(kwargs.as_mapping())?; + Ok(merged) } #[cfg(test)] mod tests { + use rstest::rstest; + use super::*; - #[test] - fn lookup_prefers_the_keyword_even_when_none_and_falls_back_to_the_request() { + fn dicts<'py>( + py: Python<'py>, + kwargs: &str, + bound: &str, + ) -> (Bound<'py, PyDict>, Bound<'py, PyDict>) { + let eval = |source: &str| { + py.eval(&std::ffi::CString::new(source).unwrap(), None, None) + .unwrap() + .cast_into::() + .unwrap() + }; + (eval(kwargs), eval(bound)) + } + + #[rstest] + #[case::keyword_wins( + "{'api_key': 'keyword'}", + "{'api_key': 'bound'}", + Some(Some("keyword")) + )] + #[case::explicit_none_wins("{'api_key': None}", "{'api_key': 'bound'}", Some(None))] + #[case::bound_fallback("{}", "{'api_key': 'bound'}", Some(Some("bound")))] + #[case::missing("{}", "{}", None)] + fn lookup_prefers_the_keyword_and_falls_back_to_bound( + #[case] kwargs: &str, + #[case] bound: &str, + #[case] expected: Option>, + ) { crate::initialize_python(); Python::attach(|py| { - let locals = PyDict::new(py); - py.run( - c" -key = object() -document = {'type': 'document_url'} -class Request: - api_key = 'from-request' - api_base = 'from-request' - document = document -request = Request() -kwargs = {'api_key': key, 'api_base': None} -", - Some(&locals), - Some(&locals), - ) - .unwrap(); - let item = |name: &str| locals.get_item(name).unwrap().unwrap(); - let kwargs = item("kwargs").cast_into::().unwrap(); - let request = item("request"); - let find = |name: &str| lookup(&kwargs, &request, name).unwrap(); - assert!(find("api_key").unwrap().is(item("key"))); - assert!(find("api_base").unwrap().is_none()); - assert!(find("document").unwrap().is(item("document"))); - assert!(find("model").is_none()); + let (kwargs, bound) = dicts(py, kwargs, bound); + let value = lookup(&kwargs, &bound, "api_key") + .unwrap() + .map(|value| value.extract::>().unwrap()); + assert_eq!(value, expected.map(|value| value.map(str::to_owned))); }); } - #[rstest::rstest] - #[case::prepared_value("{'api_key': 'replacement'}", Some("replacement"))] - #[case::explicit_none("{'api_key': None}", None)] - #[case::bound_fallback("{}", Some("original"))] - fn prepared_mapping_overrides_bound_values( - #[case] source: &str, + #[rstest] + #[case::explicit_none_hides_bound("{'api_key': None}", "{'api_key': 'bound'}", None)] + #[case::bound_none("{}", "{'api_key': None}", None)] + #[case::bound_value("{}", "{'api_key': 'bound'}", Some("bound"))] + fn present_treats_none_as_unset( + #[case] kwargs: &str, + #[case] bound: &str, #[case] expected: Option<&str>, ) { crate::initialize_python(); Python::attach(|py| { + let (kwargs, bound) = dicts(py, kwargs, bound); + let value = present(&kwargs, &bound, "api_key") + .unwrap() + .map(|value| value.extract::().unwrap()); + assert_eq!(value.as_deref(), expected); + }); + } + + #[rstest] + fn lookup_returns_the_callers_object() { + crate::initialize_python(); + Python::attach(|py| { + let document = PyDict::new(py); let bound = PyDict::new(py); - bound.set_item("api_key", "original").unwrap(); - let source = std::ffi::CString::new(source).unwrap(); - let prepared = py - .eval(&source, None, None) - .unwrap() - .cast_into::() - .unwrap(); - let value = lookup(&prepared, bound.as_any(), "api_key") + bound.set_item("document", &document).unwrap(); + let found = lookup(&PyDict::new(py), &bound, "document") .unwrap() .unwrap(); + assert!(found.is(&document)); + }); + } + + fn dict<'py>(py: Python<'py>, source: &str) -> Bound<'py, PyDict> { + py.eval(&std::ffi::CString::new(source).unwrap(), None, None) + .unwrap() + .cast_into::() + .unwrap() + } + + #[rstest] + #[case::keyword_wins("{'api_key': 'base'}", "{'api_key': 'keyword'}", Some(Some("keyword")))] + #[case::explicit_none_wins("{'api_key': 'base'}", "{'api_key': None}", Some(None))] + #[case::base_default("{'api_key': 'base'}", "{}", Some(Some("base")))] + #[case::keyword_only("{}", "{'api_key': 'keyword'}", Some(Some("keyword")))] + #[case::missing("{}", "{}", None)] + fn effective_lays_the_keywords_over_the_base( + #[case] base: &str, + #[case] kwargs: &str, + #[case] expected: Option>, + ) { + crate::initialize_python(); + Python::attach(|py| { + let merged = effective_py_args(&dict(py, base), &dict(py, kwargs)).unwrap(); + let value = merged + .get_item("api_key") + .unwrap() + .map(|value| value.extract::>().unwrap()); + assert_eq!(value, expected.map(|value| value.map(str::to_owned))); + }); + } + + #[rstest] + fn effective_leaves_both_inputs_untouched_and_keeps_object_identity() { + crate::initialize_python(); + Python::attach(|py| { + let document = PyDict::new(py); + let base = dict(py, "{'model': 'base', 'pages': None}"); + let kwargs = PyDict::new(py); + kwargs.set_item("document", &document).unwrap(); + let merged = effective_py_args(&base, &kwargs).unwrap(); + merged.set_item("model", "merged").unwrap(); assert_eq!( - value.extract::>().unwrap().as_deref(), - expected + base.get_item("model") + .unwrap() + .unwrap() + .extract::() + .unwrap(), + "base" ); + assert!(!kwargs.contains("model").unwrap()); + assert!(merged.get_item("document").unwrap().unwrap().is(&document)); }); } } diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 00543f64085..95bc03c2174 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -19,7 +19,7 @@ mod owned; mod runtime; mod services; -pub use argument::lookup; +pub use argument::{effective_py_args, lookup, present}; pub use binding::PythonBinding; pub use conversion_cache::{FromPythonCache, ToPythonCache}; pub use driver::{CallOptions, run_call}; diff --git a/litellm-rust/crates/http/Cargo.toml b/litellm-rust/crates/http/Cargo.toml index 3d880be49f9..7d64467f2f7 100644 --- a/litellm-rust/crates/http/Cargo.toml +++ b/litellm-rust/crates/http/Cargo.toml @@ -21,6 +21,7 @@ rustls-native-certs.workspace = true tokio-tungstenite.workspace = true serde.workspace = true serde_json.workspace = true +strum.workspace = true thiserror.workspace = true tokio.workspace = true tracing.workspace = true diff --git a/litellm-rust/crates/http/src/tls.rs b/litellm-rust/crates/http/src/tls.rs index c58076607e4..0b1c44bf6fe 100644 --- a/litellm-rust/crates/http/src/tls.rs +++ b/litellm-rust/crates/http/src/tls.rs @@ -42,30 +42,28 @@ impl KeyExchangeGroup { } } -#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)] +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord, strum::EnumString)] +#[strum( + parse_err_ty = Unsupported, + parse_err_fn = unsupported_cipher_token +)] pub enum Tls12CipherSuite { + #[strum(serialize = "ECDHE-ECDSA-AES128-GCM-SHA256")] EcdheEcdsaAes128Gcm, + #[strum(serialize = "ECDHE-ECDSA-AES256-GCM-SHA384")] EcdheEcdsaAes256Gcm, + #[strum(serialize = "ECDHE-ECDSA-CHACHA20-POLY1305")] EcdheEcdsaChacha20, + #[strum(serialize = "ECDHE-RSA-AES128-GCM-SHA256")] EcdheRsaAes128Gcm, + #[strum(serialize = "ECDHE-RSA-AES256-GCM-SHA384")] EcdheRsaAes256Gcm, + #[strum(serialize = "ECDHE-RSA-CHACHA20-POLY1305")] EcdheRsaChacha20, } -impl FromStr for Tls12CipherSuite { - type Err = Unsupported; - - fn from_str(name: &str) -> Result { - match name { - "ECDHE-ECDSA-AES128-GCM-SHA256" => Ok(Self::EcdheEcdsaAes128Gcm), - "ECDHE-ECDSA-AES256-GCM-SHA384" => Ok(Self::EcdheEcdsaAes256Gcm), - "ECDHE-ECDSA-CHACHA20-POLY1305" => Ok(Self::EcdheEcdsaChacha20), - "ECDHE-RSA-AES128-GCM-SHA256" => Ok(Self::EcdheRsaAes128Gcm), - "ECDHE-RSA-AES256-GCM-SHA384" => Ok(Self::EcdheRsaAes256Gcm), - "ECDHE-RSA-CHACHA20-POLY1305" => Ok(Self::EcdheRsaChacha20), - _ => Err(Unsupported::CipherToken(name.to_owned())), - } - } +fn unsupported_cipher_token(name: &str) -> Unsupported { + Unsupported::CipherToken(name.to_owned()) } impl Tls12CipherSuite { @@ -367,6 +365,39 @@ mod tests { ); } + #[rstest] + #[case::ecdhe_ecdsa_aes128( + "ECDHE-ECDSA-AES128-GCM-SHA256", + Tls12CipherSuite::EcdheEcdsaAes128Gcm + )] + #[case::ecdhe_ecdsa_aes256( + "ECDHE-ECDSA-AES256-GCM-SHA384", + Tls12CipherSuite::EcdheEcdsaAes256Gcm + )] + #[case::ecdhe_ecdsa_chacha20( + "ECDHE-ECDSA-CHACHA20-POLY1305", + Tls12CipherSuite::EcdheEcdsaChacha20 + )] + #[case::ecdhe_rsa_aes128("ECDHE-RSA-AES128-GCM-SHA256", Tls12CipherSuite::EcdheRsaAes128Gcm)] + #[case::ecdhe_rsa_aes256("ECDHE-RSA-AES256-GCM-SHA384", Tls12CipherSuite::EcdheRsaAes256Gcm)] + #[case::ecdhe_rsa_chacha20("ECDHE-RSA-CHACHA20-POLY1305", Tls12CipherSuite::EcdheRsaChacha20)] + fn tls12_cipher_suite_parses_each_openssl_name( + #[case] name: &str, + #[case] expected: Tls12CipherSuite, + ) { + assert_eq!(name.parse::().unwrap(), expected); + } + + #[rstest] + #[case::unknown("AES128-SHA")] + #[case::case_sensitive("ecdhe-rsa-aes256-gcm-sha384")] + fn unsupported_cipher_token_keeps_the_verbatim_name(#[case] name: &str) { + assert_eq!( + name.parse::().unwrap_err(), + Unsupported::CipherToken(name.to_owned()) + ); + } + #[test] fn named_suites_are_the_only_tls12_suites_offered_and_tls13_stays() { let tls = ClientConfig::try_from(&config(HttpSettings { diff --git a/litellm-rust/crates/lens/tests/receiver.rs b/litellm-rust/crates/lens/tests/receiver.rs index bda1cecbbb3..2af3c67955a 100644 --- a/litellm-rust/crates/lens/tests/receiver.rs +++ b/litellm-rust/crates/lens/tests/receiver.rs @@ -10,7 +10,10 @@ use rstest::rstest; use serde_json::json; use sha2::{Digest, Sha256}; use std::{ - sync::{Arc, atomic::Ordering}, + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, time::Duration, }; use wiremock::{ @@ -116,6 +119,107 @@ async fn agent_picker_query_preserves_scope_through_the_internal_read_route() { assert_eq!(response.json::().await.unwrap(), result); } +#[rstest] +#[case::list_paid("list", None, 0.75)] +#[case::list_zero("list", None, 0.0)] +#[case::detail_paid("trace", None, 0.75)] +#[case::paged_zero("trace", Some(1), 0.0)] +#[tokio::test] +async fn internal_reads_refresh_delayed_gateway_amounts( + #[case] operation: &str, + #[case] page_size: Option, + #[case] cost: f64, +) { + let store = MockServer::start().await; + let start_ms = (unix_seconds() as i64 - 600) * 1000; + let available = Arc::new(AtomicBool::new(false)); + Mock::given(body_string_contains("FROM agent_traces_by_key")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"data": [{ + "trace_id": "trace", "trace_ref": "ref", "team_id": "team", + "api_key_hash": "key", "user_id": "owner", "name": "model call", + "service": "agent", "input_preview": "", "status": "STATUS_CODE_OK", + "start_ms": start_ms, "duration_ms": 1, "span_count": 1, + "agent_count": 0, "agent_invocations": 0, "llm_calls": 1, + "tool_calls": 0, "input_tokens": 1, "output_tokens": 1, + "models": ["test-model"], "error_count": 0, "request_ids": [] + }]}))) + .mount(&store) + .await; + Mock::given(body_string_contains("o.SpanId AS span_id")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"data": [{ + "trace_id": "trace", "span_id": "span", "parent_span_id": "", + "name": "model call", "type": "llm", "agent": "", + "status": "STATUS_CODE_OK", "status_message": "", "error_truncated": 0, + "start_ns": start_ms * 1_000_000, "duration_ns": 1_000_000, + "service": "agent", "input_preview": "", "model": "test-model", + "input_tokens": 1, "output_tokens": 1, "litellm_request_id": "", + "call_keys": ["provider_response:response"], "call_evidence": "complete", + "team_id": "team", "api_key_hash": "key", "user_id": "owner" + }]}))) + .expect(2) + .mount(&store) + .await; + let spend_available = available.clone(); + Mock::given(body_string_contains("FROM spend_logs FINAL")) + .respond_with(move |_: &wiremock::Request| { + let rows = if spend_available.load(Ordering::Acquire) { + json!([{ + "request_id": "request", "litellm_call_id": "", + "response_id": "response", "upstream_response_id": "", + "trace_id": "", "span_id": "", "team_id": "team", + "api_key": "key", "user": "owner", "spend": cost, + "start_ms": start_ms + }]) + } else { + json!([]) + }; + ResponseTemplate::new(200).set_body_json(json!({"data": rows})) + }) + .expect(2) + .mount(&store) + .await; + let server = serve(&store.uri(), true).await; + let scope = json!({"all_teams": 0, "user_id": "owner", "team_ids": []}); + let (request, summary_path) = if operation == "list" { + ( + json!({ + "operation": operation, "scope": scope, "start_ms": start_ms, + "end_ms": start_ms + 1000, "cursor": null, "limit": 50 + }), + "/data/0", + ) + } else { + ( + json!({ + "operation": operation, "scope": scope, "trace_id": "trace", + "trace_ref": "ref", "cursor": null, "page_size": page_size + }), + "/summary", + ) + }; + let client = http_client().unwrap(); + for expected in [None, Some(cost)] { + if expected.is_some() { + available.store(true, Ordering::Release); + tokio::time::sleep(litellm_traces_cache::LIVE_TTL + Duration::from_millis(200)).await; + } + let response = client + .post(format!("{}/internal/read", server.url)) + .bearer_auth(SERVICE_TOKEN) + .json(&request) + .send() + .await + .unwrap(); + assert_eq!(response.status(), 200); + let body = response.json::().await.unwrap(); + let summary = body.pointer(summary_path).unwrap(); + assert_eq!(summary["spend"], json!(expected)); + assert_eq!(summary["priced_calls"], u64::from(expected.is_some())); + assert_eq!(summary["llm_calls"], 1); + assert!(!body.to_string().contains("gateway_spend_pending")); + } +} + #[rstest] #[tokio::test] async fn feedback_summary_query_preserves_scope_through_the_internal_read_route() { diff --git a/litellm-rust/crates/llms-types/AGENTS.md b/litellm-rust/crates/llms-types/AGENTS.md index 0d418822d06..e5bd83b1d6e 100644 --- a/litellm-rust/crates/llms-types/AGENTS.md +++ b/litellm-rust/crates/llms-types/AGENTS.md @@ -67,3 +67,9 @@ The same ownership rule applies to Messages, Responses, Chat Completions, OCR, a - Follow the workspace test-placement and `rstest` rules - Test provider transformations, header policy, and stream execution in their owning crates - Do not test import locations or Rust source structure as substitutes for behavior + +# references + +## json_schema.rs + +- https://json-schema.org/draft/2020-12/json-schema-core diff --git a/litellm-rust/crates/llms-types/Cargo.toml b/litellm-rust/crates/llms-types/Cargo.toml index 2d880b87faf..4f190f34026 100644 --- a/litellm-rust/crates/llms-types/Cargo.toml +++ b/litellm-rust/crates/llms-types/Cargo.toml @@ -6,9 +6,10 @@ license.workspace = true repository.workspace = true [features] -schema = ["dep:schemars"] +schema = ["dep:schemars", "schemars/indexmap2"] [dependencies] +indexmap = { version = "2.14.0", features = ["serde"] } macro_rules_attribute.workspace = true schemars = { workspace = true, optional = true } serde.workspace = true diff --git a/litellm-rust/crates/llms-types/src/formats/AGENTS.md b/litellm-rust/crates/llms-types/src/formats/AGENTS.md new file mode 100644 index 00000000000..e49c286f058 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/AGENTS.md @@ -0,0 +1,36 @@ +# references + +These links describe shared format fields and discriminators. The official API reference is authoritative, and SDK sources only clarify shapes the reference leaves implicit. These contracts are partial typed projections, not exhaustive upstream schemas + +## chat_completions.rs and chat_completions/content.rs + +- https://developers.openai.com/api/reference/resources/chat + +SDK wire definitions: + +- https://github.com/openai/openai-python/blob/main/src/openai/types/chat/chat_completion_content_part_param.py + +## responses/output.rs and responses/streaming_websocket.rs + +- https://developers.openai.com/api/reference/resources/responses +- https://developers.openai.com/api/reference/resources/responses/streaming-events + +SDK wire definitions: + +- https://github.com/openai/openai-python/blob/main/src/openai/types/responses/response.py +- https://github.com/openai/openai-python/blob/main/src/openai/types/responses/response_output_item.py +- https://github.com/openai/openai-python/blob/main/src/openai/types/responses/mcp_tool_call_error.py + +## batches.rs + +- https://developers.openai.com/api/reference/resources/batches + +## audio_transcription.rs + +- https://developers.openai.com/api/reference/resources/audio/subresources/transcriptions + +## ocr.rs + +- https://docs.mistral.ai/api/endpoint/ocr + +Messages references live in `messages/AGENTS.md`. `LiteLLMOcrResponse` follows the Mistral OCR response shape, and `OcrBoundingBox` describes the corner coordinates providers copy into the normalized `bbox`. Normalized `tables` and `keyValuePairs` stay open because providers pass through different native shapes diff --git a/litellm-rust/crates/llms-types/src/formats/audio_transcription.rs b/litellm-rust/crates/llms-types/src/formats/audio_transcription.rs index e00ecb0b5fb..624eb7be4e4 100644 --- a/litellm-rust/crates/llms-types/src/formats/audio_transcription.rs +++ b/litellm-rust/crates/llms-types/src/formats/audio_transcription.rs @@ -1,6 +1,6 @@ use serde_json::Value; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct AudioTranscriptionResponseData { pub text: String, } diff --git a/litellm-rust/crates/llms-types/src/formats/batches.rs b/litellm-rust/crates/llms-types/src/formats/batches.rs index 9749b042a36..0390be04dec 100644 --- a/litellm-rust/crates/llms-types/src/formats/batches.rs +++ b/litellm-rust/crates/llms-types/src/formats/batches.rs @@ -1,4 +1,4 @@ -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Copy, Eq)] #[serde(rename_all = "snake_case")] pub enum BatchStatus { @@ -7,7 +7,7 @@ pub enum BatchStatus { Completed, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Eq)] pub struct BatchRequestCounts { pub total: u64, @@ -15,7 +15,7 @@ pub struct BatchRequestCounts { pub failed: u64, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Eq)] pub struct BatchResponse { pub id: String, diff --git a/litellm-rust/crates/llms-types/src/formats/chat_completions.rs b/litellm-rust/crates/llms-types/src/formats/chat_completions.rs index 31b5046469a..8f154dd7fb5 100644 --- a/litellm-rust/crates/llms-types/src/formats/chat_completions.rs +++ b/litellm-rust/crates/llms-types/src/formats/chat_completions.rs @@ -1,51 +1,41 @@ use serde_json::{Map, Value}; use strum::IntoStaticStr; +mod content; +pub use content::{ + ChatContentPart, ChatFile, ChatInputAudio, ChatLogprobs, ChatMediaUrl, ChatMediaUrlParameters, + ChatTokenLogprob, ChatTopLogprob, ChatVideoMetadata, PromptCacheBreakpoint, PromptCacheMode, +}; + /// Reasoning effort level accepted or applied by the model. -#[macro_rules_attribute::apply(wire_type)] -#[derive(Copy, Eq, IntoStaticStr)] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq, IntoStaticStr, strum::EnumString, strum::VariantArray)] #[serde(rename_all = "snake_case")] -#[strum(serialize_all = "snake_case")] pub enum ReasoningEffort { + #[strum(serialize = "none")] None, + #[strum(serialize = "minimal")] Minimal, + #[strum(serialize = "low")] Low, + #[strum(serialize = "medium")] Medium, + #[strum(serialize = "high")] High, + #[strum(serialize = "xhigh")] Xhigh, + #[strum(serialize = "max")] Max, } -impl ReasoningEffort { - pub const ALL: [Self; 7] = [ - Self::None, - Self::Minimal, - Self::Low, - Self::Medium, - Self::High, - Self::Xhigh, - Self::Max, - ]; - - pub fn as_str(self) -> &'static str { - self.into() - } - - pub fn parse(value: &str) -> Option { - Self::ALL - .into_iter() - .find(|effort| effort.as_str() == value) - } -} - -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(untagged)] pub enum ChatMessageContent { Text(String), Parts(Vec), } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct ChatMessage { pub role: String, #[serde(default, skip_serializing_if = "Option::is_none")] @@ -56,7 +46,7 @@ pub struct ChatMessage { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct ChatCompletionToolCallFunctionChunk { #[serde(default, skip_serializing_if = "Option::is_none")] pub name: Option, @@ -65,7 +55,7 @@ pub struct ChatCompletionToolCallFunctionChunk { pub provider_specific_fields: Option>, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct ChatCompletionToolCallChunk { #[serde(default, skip_serializing_if = "Option::is_none")] pub id: Option, @@ -75,7 +65,7 @@ pub struct ChatCompletionToolCallChunk { pub index: i64, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(tag = "type", rename_all = "snake_case")] pub enum ChatCompletionThinkingBlock { Thinking { @@ -96,7 +86,7 @@ pub enum ChatCompletionThinkingBlock { /// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python /// path reports so cost tracking sees the same numbers on either path. -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct PromptTokensDetails { pub cached_tokens: u64, @@ -104,7 +94,7 @@ pub struct PromptTokensDetails { pub text_tokens: u64, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct ChatCompletionsUsage { pub prompt_tokens: u64, @@ -113,7 +103,7 @@ pub struct ChatCompletionsUsage { pub prompt_tokens_details: PromptTokensDetails, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct ChatCompletionsChoiceMessage { pub role: String, // Whether an empty turn is `None` or `""` is the provider's choice, not a @@ -123,7 +113,7 @@ pub struct ChatCompletionsChoiceMessage { pub content: Option, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct ChatCompletionsChoice { pub index: u64, pub message: ChatCompletionsChoiceMessage, @@ -135,7 +125,7 @@ pub struct ChatCompletionsChoice { /// There is deliberately no `id`: Python mints the `chatcmpl-…` id on the /// `ModelResponse` it already created, and echoing the provider's own id here /// would change it. Pinned by `response_carries_no_id` in the Anthropic chat transformation tests. -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct ChatCompletionsResponse { pub created: u64, pub model: String, @@ -143,7 +133,7 @@ pub struct ChatCompletionsResponse { pub usage: ChatCompletionsUsage, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct ChatCompletionDelta { #[serde(default, skip_serializing_if = "Option::is_none")] @@ -162,7 +152,7 @@ pub struct ChatCompletionDelta { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct ChatCompletionStreamingChoice { pub index: u64, pub delta: ChatCompletionDelta, @@ -172,7 +162,7 @@ pub struct ChatCompletionStreamingChoice { pub logprobs: Option, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct ChatCompletionChunk { pub id: String, pub created: u64, @@ -189,6 +179,7 @@ pub struct ChatCompletionChunk { #[cfg(test)] mod tests { use rstest::rstest; + use strum::VariantArray; use super::*; @@ -207,10 +198,10 @@ mod tests { ) { assert_eq!( serde_json::to_value(effort).unwrap(), - Value::String(effort.as_str().to_string()) + Value::String(<&'static str>::from(effort).to_string()) ); - assert_eq!(ReasoningEffort::parse(effort.as_str()), Some(effort)); - assert!(ReasoningEffort::ALL.contains(&effort)); + assert_eq!(<&'static str>::from(effort).parse(), Ok(effort)); + assert!(ReasoningEffort::VARIANTS.contains(&effort)); } #[rstest] @@ -218,6 +209,6 @@ mod tests { #[case::uppercase("HIGH")] #[case::empty("")] fn reasoning_effort_parse_rejects(#[case] value: &str) { - assert_eq!(ReasoningEffort::parse(value), None); + assert!(value.parse::().is_err()); } } diff --git a/litellm-rust/crates/llms-types/src/formats/chat_completions/content.rs b/litellm-rust/crates/llms-types/src/formats/chat_completions/content.rs new file mode 100644 index 00000000000..d508bb672ab --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/chat_completions/content.rs @@ -0,0 +1,155 @@ +use serde_json::{Map, Value}; + +use crate::formats::messages::{CacheControl, CitationsConfig, ContentSource}; + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ChatContentPart { + Text { + text: String, + #[serde(skip_serializing_if = "Option::is_none")] + cache_control: Option, + #[serde(skip_serializing_if = "Option::is_none")] + prompt_cache_breakpoint: Option, + #[serde(flatten)] + extra: Map, + }, + ImageUrl { + image_url: ChatMediaUrl, + #[serde(skip_serializing_if = "Option::is_none")] + prompt_cache_breakpoint: Option, + #[serde(flatten)] + extra: Map, + }, + VideoUrl { + video_url: ChatMediaUrl, + #[serde(flatten)] + extra: Map, + }, + InputAudio { + input_audio: ChatInputAudio, + #[serde(skip_serializing_if = "Option::is_none")] + prompt_cache_breakpoint: Option, + #[serde(flatten)] + extra: Map, + }, + File { + file: Box, + #[serde(skip_serializing_if = "Option::is_none")] + prompt_cache_breakpoint: Option, + #[serde(flatten)] + extra: Map, + }, + Document { + source: Box, + #[serde(skip_serializing_if = "Option::is_none")] + title: Option, + #[serde(skip_serializing_if = "Option::is_none")] + context: Option, + #[serde(skip_serializing_if = "Option::is_none")] + citations: Option, + #[serde(flatten)] + extra: Map, + }, + Refusal { + refusal: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct PromptCacheBreakpoint { + pub mode: PromptCacheMode, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum PromptCacheMode { + Explicit, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(untagged)] +pub enum ChatMediaUrl { + Url(String), + Parameters(Box), +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ChatMediaUrlParameters { + pub url: String, + pub detail: Option, + pub format: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ChatInputAudio { + pub data: String, + pub format: String, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ChatFile { + pub file_data: Option, + pub file_id: Option, + pub filename: Option, + pub format: Option, + pub detail: Option, + pub video_metadata: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ChatVideoMetadata { + pub fps: Option, + pub start_offset: Option, + pub end_offset: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ChatLogprobs { + pub content: Option>, + pub refusal: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ChatTokenLogprob { + pub token: String, + pub logprob: serde_json::Number, + pub bytes: Option>, + pub top_logprobs: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ChatTopLogprob { + pub token: String, + pub logprob: serde_json::Number, + pub bytes: Option>, + #[serde(flatten)] + pub extra: Map, +} diff --git a/litellm-rust/crates/llms-types/src/formats/messages/AGENTS.md b/litellm-rust/crates/llms-types/src/formats/messages/AGENTS.md index 72558688269..20a5632f267 100644 --- a/litellm-rust/crates/llms-types/src/formats/messages/AGENTS.md +++ b/litellm-rust/crates/llms-types/src/formats/messages/AGENTS.md @@ -3,3 +3,17 @@ This directory owns shared Messages API data contracts and serialization: reques Adapter contracts and execution inputs such as `MessagesTransformContext` belong in `llms/src/base_llm/messages`. Provider rewriting and interpretation belong in `llms/src//messages`. Call envelopes, live streams, and call orchestration belong in `inference-messages` Represent web-search results, encrypted-content fields, and thinking configuration as data here. Decisions to flatten results, remove encrypted content, select thinking budgets, or require beta headers belong to provider transformations. Data-shape validation belongs here, while model capability checks and request adaptation do not + +# references + +These upstream contracts document fields and discriminators. They are documentation links, not saved golden snapshots. Beta-only blocks and tools (MCP, compaction, tool changes, advisor, fallback, browser state, and toolsets) follow the beta reference + +## Messages bodies, content, tools, and streaming + +- https://platform.claude.com/docs/en/api/http/messages/create +- https://platform.claude.com/docs/en/api/messages/create.md +- https://platform.claude.com/docs/en/api/beta/messages/create.md +- https://platform.claude.com/docs/en/build-with-claude/streaming.md +- https://platform.claude.com/docs/en/build-with-claude/context-editing + +Anthropic-compatible hosts document their deviations in the `llms/src//messages` guides. Copy a host reference into `../../providers/AGENTS.md` only when that host gets a typed extension in `providers` diff --git a/litellm-rust/crates/llms-types/src/formats/messages/content.rs b/litellm-rust/crates/llms-types/src/formats/messages/content.rs new file mode 100644 index 00000000000..595ba68c103 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/messages/content.rs @@ -0,0 +1,712 @@ +use serde_json::{Map, Value}; + +use super::{CacheControl, MessagesToolParam}; +use crate::json_schema::JsonSchema; + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ContentSource { + Base64 { + media_type: String, + data: String, + #[serde(flatten)] + extra: Map, + }, + Url { + url: String, + #[serde(flatten)] + extra: Map, + }, + File { + file_id: String, + #[serde(flatten)] + extra: Map, + }, + Text { + media_type: String, + data: String, + #[serde(flatten)] + extra: Map, + }, + Content { + content: BlockContent, + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(untagged)] +pub enum BlockContent { + Text(String), + Blocks(Vec), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ToolCaller { + Direct { + #[serde(flatten)] + extra: Map, + }, + #[serde(rename = "code_execution_20250825")] + CodeExecution { + tool_id: String, + #[serde(flatten)] + extra: Map, + }, + #[serde(rename = "code_execution_20260120")] + CodeExecution20260120 { + tool_id: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct CitationsConfig { + pub enabled: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct PageCitation { + pub cited_text: String, + pub document_index: u64, + pub document_title: Option, + pub start_page_number: u64, + pub end_page_number: u64, + pub file_id: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct CharCitation { + pub cited_text: String, + pub document_index: u64, + pub document_title: Option, + pub start_char_index: u64, + pub end_char_index: u64, + pub file_id: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ContentBlockCitation { + pub cited_text: String, + pub document_index: u64, + pub document_title: Option, + pub start_block_index: u64, + pub end_block_index: u64, + pub file_id: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct WebSearchCitation { + pub cited_text: String, + pub url: String, + pub encrypted_index: String, + pub title: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct SearchResultCitation { + pub cited_text: String, + pub search_result_index: u64, + pub source: String, + pub title: Option, + pub start_block_index: u64, + pub end_block_index: u64, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum Citation { + PageLocation(PageCitation), + CharLocation(CharCitation), + WebSearchResultLocation(WebSearchCitation), + ContentBlockLocation(ContentBlockCitation), + SearchResultLocation(SearchResultCitation), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum MessagesContentPart { + Text(TextBlock), + Image(ImageBlock), + Document(Box), + SearchResult(SearchResultBlock), + Thinking(ThinkingBlock), + RedactedThinking(RedactedThinkingBlock), + ToolUse(ToolUseBlock), + ToolResult(ToolResultBlock), + ToolReference(ToolReferenceBlock), + BrowserState(BrowserStateBlock), + ServerToolUse(ServerToolUseBlock), + WebSearchToolResult(ServerToolResultBlock), + WebFetchToolResult(ServerToolResultBlock), + CodeExecutionToolResult(ServerToolResultBlock), + BashCodeExecutionToolResult(ServerToolResultBlock), + TextEditorCodeExecutionToolResult( + ServerToolResultBlock, + ), + ToolSearchToolResult(ServerToolResultBlock), + AdvisorToolResult(ServerToolResultBlock), + ContainerUpload(ContainerUploadBlock), + McpToolUse(McpToolUseBlock), + McpToolResult(McpToolResultBlock), + McpToolListing(McpToolListingBlock), + Compaction(CompactionBlock), + ToolAddition(ToolChangeBlock), + ToolRemoval(ToolChangeBlock), + Fallback(FallbackBlock), +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct TextBlock { + pub text: String, + pub citations: Option>, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ImageBlock { + pub source: ContentSource, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct DocumentBlock { + pub source: ContentSource, + pub title: Option, + pub context: Option, + pub citations: Option, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct SearchResultBlock { + pub source: String, + pub title: String, + pub content: Vec, + pub citations: Option, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ThinkingBlock { + pub thinking: String, + pub signature: String, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct RedactedThinkingBlock { + pub data: String, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ToolUseBlock { + pub id: String, + pub name: String, + pub input: Map, + pub caller: Option, + pub toolset_name: Option, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ToolResultBlock { + pub tool_use_id: String, + pub content: Option, + pub is_error: Option, + pub toolset_name: Option, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ToolReferenceBlock { + pub tool_name: String, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct BrowserStateBlock { + pub tabs: Vec, + pub state_changes: Option>, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct BrowserTab { + pub tab_id: String, + pub title: String, + pub url: String, + pub active: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum BrowserStateChange { + TabOpened { + tab_id: String, + #[serde(flatten)] + extra: Map, + }, + DownloadStarted { + download_id: String, + url: String, + #[serde(flatten)] + extra: Map, + }, + DownloadCompleted { + download_id: String, + url: String, + #[serde(skip_serializing_if = "Option::is_none")] + path: Option, + #[serde(skip_serializing_if = "Option::is_none")] + size_bytes: Option, + #[serde(flatten)] + extra: Map, + }, + DownloadFailed { + download_id: String, + url: String, + #[serde(skip_serializing_if = "Option::is_none")] + error: Option, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ServerToolUseBlock { + pub id: String, + pub name: String, + pub input: Map, + pub caller: Option, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ServerToolResultBlock { + pub tool_use_id: String, + pub content: C, + pub caller: Option, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ServerToolError { + pub error_code: String, + pub error_message: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(untagged)] +pub enum WebSearchToolResultContent { + Error(WebSearchResultError), + Results(Vec), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct WebSearchResultError { + #[serde(rename = "type")] + pub error_type: WebSearchResultErrorType, + pub error_code: String, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum WebSearchResultErrorType { + WebSearchToolResultError, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct WebSearchResult { + #[serde(rename = "type")] + pub result_type: WebSearchResultType, + pub url: String, + pub title: String, + pub encrypted_content: String, + pub page_age: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum WebSearchResultType { + WebSearchResult, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum WebFetchToolResultContent { + WebFetchToolResultError(ServerToolError), + WebFetchResult(WebFetchResult), +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct WebFetchResult { + pub url: String, + pub content: WebFetchDocument, + pub retrieved_at: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum WebFetchDocument { + Document(Box), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum CodeExecutionToolResultContent { + CodeExecutionToolResultError(ServerToolError), + CodeExecutionResult(CodeExecutionResult), + EncryptedCodeExecutionResult(EncryptedCodeExecutionResult), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum BashCodeExecutionToolResultContent { + BashCodeExecutionToolResultError(ServerToolError), + BashCodeExecutionResult(CodeExecutionResult), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct CodeExecutionResult { + pub stdout: String, + pub stderr: String, + pub return_code: i64, + pub content: Vec, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct EncryptedCodeExecutionResult { + pub encrypted_stdout: String, + pub stderr: String, + pub return_code: i64, + pub content: Vec, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum CodeExecutionOutput { + CodeExecutionOutput { + file_id: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum BashCodeExecutionOutput { + BashCodeExecutionOutput { + file_id: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum TextEditorCodeExecutionToolResultContent { + TextEditorCodeExecutionToolResultError(ServerToolError), + TextEditorCodeExecutionViewResult(TextEditorViewResult), + TextEditorCodeExecutionCreateResult(TextEditorCreateResult), + TextEditorCodeExecutionStrReplaceResult(TextEditorStrReplaceResult), +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct TextEditorViewResult { + pub content: String, + pub file_type: TextEditorFileType, + pub num_lines: Option, + pub start_line: Option, + pub total_lines: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum TextEditorFileType { + Text, + Image, + Pdf, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct TextEditorCreateResult { + pub is_file_update: bool, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct TextEditorStrReplaceResult { + pub lines: Option>, + pub new_lines: Option, + pub new_start: Option, + pub old_lines: Option, + pub old_start: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ToolSearchToolResultContent { + ToolSearchToolResultError(ServerToolError), + ToolSearchToolSearchResult(ToolSearchResult), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ToolSearchResult { + pub tool_references: Vec, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ToolSearchReference { + ToolReference(ToolReferenceBlock), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum AdvisorToolResultContent { + AdvisorToolResultError(ServerToolError), + AdvisorResult { + text: String, + #[serde(skip_serializing_if = "Option::is_none")] + stop_reason: Option, + #[serde(flatten)] + extra: Map, + }, + AdvisorRedactedResult { + encrypted_content: String, + #[serde(skip_serializing_if = "Option::is_none")] + stop_reason: Option, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ContainerUploadBlock { + pub file_id: String, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct McpToolUseBlock { + pub id: String, + pub name: String, + pub server_name: String, + pub input: Map, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct McpToolResultBlock { + pub tool_use_id: String, + pub content: Option, + pub is_error: Option, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(untagged)] +pub enum McpToolResultContent { + Text(String), + Blocks(Vec), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum McpToolResultText { + Text(TextBlock), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct McpToolListingBlock { + pub mcp_server_name: String, + pub tools: Vec, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct McpListedTool { + pub name: String, + pub description: Option, + pub input_schema: JsonSchema, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct CompactionBlock { + pub content: Option, + pub encrypted_content: Option, + pub signature: Option, + pub tool_changes: Option>, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ToolChange { + ToolAddition(ToolChangeBlock), + ToolRemoval(ToolChangeBlock), +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ToolChangeBlock { + pub tool: ToolChangeTarget, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ToolChangeTarget { + ToolReference { + name: String, + #[serde(flatten)] + extra: Map, + }, + McpToolReference { + server_name: String, + name: String, + #[serde(flatten)] + extra: Map, + }, + McpToolsetReference { + server_name: String, + #[serde(flatten)] + extra: Map, + }, + ToolDefinition { + definition: Box, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct FallbackBlock { + pub from: FallbackModel, + pub to: FallbackModel, + pub trigger: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct FallbackModel { + pub model: String, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum FallbackTrigger { + Refusal { + #[serde(skip_serializing_if = "Option::is_none")] + category: Option, + #[serde(flatten)] + extra: Map, + }, +} diff --git a/litellm-rust/crates/llms-types/src/formats/messages/metadata.rs b/litellm-rust/crates/llms-types/src/formats/messages/metadata.rs new file mode 100644 index 00000000000..6e5ebb2b98d --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/messages/metadata.rs @@ -0,0 +1,340 @@ +use serde_json::{Map, Value}; + +use crate::json_schema::JsonSchema; +use crate::recognized::Recognized; + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct MessagesMetadata { + pub user_id: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct OutputFormat { + #[serde(rename = "type")] + pub format_type: OutputFormatType, + pub schema: JsonSchema, + pub strict: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum OutputFormatType { + JsonSchema, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct MessagesCompaction { + #[serde(rename = "type")] + pub compaction_type: CompactionType, + pub instructions: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum CompactionType { + Summarize, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct MessagesContainer { + pub id: Option, + pub expires_at: Option, + pub skills: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ContainerSkill { + #[serde(rename = "type")] + pub skill_type: SkillType, + pub skill_id: String, + pub version: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum SkillType { + Anthropic, + Custom, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct McpServer { + #[serde(rename = "type")] + pub server_type: McpServerType, + pub url: String, + pub name: String, + pub authorization_token: Option, + pub tool_configuration: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum McpServerType { + Url, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct McpToolConfiguration { + pub allowed_tools: Option>, + pub enabled: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct StopDetails { + #[serde(rename = "type")] + pub detail_type: Recognized, + pub category: Option, + pub explanation: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum StopDetailsType { + Refusal, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ContextManagementResponse { + pub applied_edits: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct AppliedEdit { + #[serde(rename = "type")] + pub edit_type: Option, + pub cleared_input_tokens: Option, + pub cleared_tool_uses: Option, + pub cleared_thinking_turns: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct Safeguard { + #[serde(rename = "type")] + pub safeguard_type: String, + pub classifier_context: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[derive( + Clone, + Debug, + PartialEq, + Eq, + strum::Display, + strum::EnumString, + strum::AsRefStr, + serde_with::DeserializeFromStr, + serde_with::SerializeDisplay, +)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[cfg_attr(feature = "schema", schemars(from = "String", into = "String"))] +pub enum MessageRole { + #[strum(serialize = "user")] + User, + #[strum(serialize = "assistant")] + Assistant, + #[strum(serialize = "system")] + System, + #[strum(default, transparent)] + Other(String), +} + +impl MessageRole { + pub fn as_str(&self) -> &str { + self.as_ref() + } +} + +impl From for MessageRole { + fn from(value: String) -> Self { + value.parse().unwrap_or_else(|never| match never {}) + } +} + +impl From for String { + fn from(value: MessageRole) -> Self { + value.to_string() + } +} + +#[derive( + Clone, + Debug, + PartialEq, + Eq, + strum::Display, + strum::EnumString, + strum::AsRefStr, + serde_with::DeserializeFromStr, + serde_with::SerializeDisplay, +)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[cfg_attr(feature = "schema", schemars(from = "String", into = "String"))] +pub enum MessageType { + #[strum(serialize = "message")] + Message, + #[strum(default, transparent)] + Other(String), +} + +impl MessageType { + pub fn as_str(&self) -> &str { + self.as_ref() + } +} + +impl From for MessageType { + fn from(value: String) -> Self { + value.parse().unwrap_or_else(|never| match never {}) + } +} + +impl From for String { + fn from(value: MessageType) -> Self { + value.to_string() + } +} + +#[derive( + Clone, + Debug, + PartialEq, + Eq, + strum::Display, + strum::EnumString, + strum::AsRefStr, + serde_with::DeserializeFromStr, + serde_with::SerializeDisplay, +)] +#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] +#[cfg_attr(feature = "schema", schemars(from = "String", into = "String"))] +pub enum StopReason { + #[strum(serialize = "end_turn")] + EndTurn, + #[strum(serialize = "max_tokens")] + MaxTokens, + #[strum(serialize = "stop_sequence")] + StopSequence, + #[strum(serialize = "tool_use")] + ToolUse, + #[strum(serialize = "refusal")] + Refusal, + #[strum(serialize = "compaction")] + Compaction, + #[strum(serialize = "pause_turn")] + PauseTurn, + #[strum(serialize = "model_context_window_exceeded")] + ModelContextWindowExceeded, + #[strum(default, transparent)] + Other(String), +} + +impl StopReason { + pub fn as_str(&self) -> &str { + self.as_ref() + } +} + +impl From for StopReason { + fn from(value: String) -> Self { + value.parse().unwrap_or_else(|never| match never {}) + } +} + +impl From for String { + fn from(value: StopReason) -> Self { + value.to_string() + } +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(untagged)] +pub enum ContainerReference { + Id(String), + Parameters(Box), +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct MessagesDiagnosticsParam { + #[serde( + default, + skip_serializing_if = "Option::is_none", + with = "::serde_with::rust::double_option" + )] + #[cfg_attr(feature = "schema", schemars(with = "Option"))] + pub previous_message_id: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct MessagesDiagnostics { + pub cache_miss_reason: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum CacheMissReason { + ModelChanged(CacheMissedTokens), + SystemChanged(CacheMissedTokens), + ToolsChanged(CacheMissedTokens), + MessagesChanged(CacheMissedTokens), + PreviousMessageNotFound { + #[serde(flatten)] + extra: Map, + }, + Unavailable { + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct CacheMissedTokens { + pub cache_missed_input_tokens: u64, + #[serde(flatten)] + pub extra: Map, +} diff --git a/litellm-rust/crates/llms-types/src/formats/messages/mod.rs b/litellm-rust/crates/llms-types/src/formats/messages/mod.rs index 219e0ae63a0..08695ca7ecc 100644 --- a/litellm-rust/crates/llms-types/src/formats/messages/mod.rs +++ b/litellm-rust/crates/llms-types/src/formats/messages/mod.rs @@ -1,11 +1,54 @@ +mod content; +mod metadata; mod request; mod response; pub mod streaming; +mod tools; +mod usage; +pub use content::{ + AdvisorToolResultContent, BashCodeExecutionOutput, BashCodeExecutionToolResultContent, + BlockContent, BrowserStateBlock, BrowserStateChange, BrowserTab, CharCitation, Citation, + CitationsConfig, CodeExecutionOutput, CodeExecutionResult, CodeExecutionToolResultContent, + CompactionBlock, ContainerUploadBlock, ContentBlockCitation, ContentSource, DocumentBlock, + EncryptedCodeExecutionResult, FallbackBlock, FallbackModel, FallbackTrigger, ImageBlock, + McpListedTool, McpToolListingBlock, McpToolResultBlock, McpToolResultContent, + McpToolResultText, McpToolUseBlock, MessagesContentPart, PageCitation, RedactedThinkingBlock, + SearchResultBlock, SearchResultCitation, ServerToolError, ServerToolResultBlock, + ServerToolUseBlock, TextBlock, TextEditorCodeExecutionToolResultContent, + TextEditorCreateResult, TextEditorFileType, TextEditorStrReplaceResult, TextEditorViewResult, + ThinkingBlock, ToolCaller, ToolChange, ToolChangeBlock, ToolChangeTarget, ToolReferenceBlock, + ToolResultBlock, ToolSearchReference, ToolSearchResult, ToolSearchToolResultContent, + ToolUseBlock, WebFetchDocument, WebFetchResult, WebFetchToolResultContent, WebSearchCitation, + WebSearchResult, WebSearchResultError, WebSearchResultErrorType, WebSearchResultType, + WebSearchToolResultContent, +}; +pub use metadata::{ + AppliedEdit, CacheMissReason, CacheMissedTokens, CompactionType, ContainerReference, + ContainerSkill, ContextManagementResponse, McpServer, McpServerType, McpToolConfiguration, + MessageRole, MessageType, MessagesCompaction, MessagesContainer, MessagesDiagnostics, + MessagesDiagnosticsParam, MessagesMetadata, OutputFormat, OutputFormatType, Safeguard, + SkillType, StopDetails, StopDetailsType, StopReason, +}; pub use request::{ AdaptiveThinking, CacheControl, ContentBlock, ContentBlockType, ContextEdit, ContextManagement, - DisabledThinking, EffortLevel, EnabledThinking, Message, MessageContent, + ContextTrigger, DisabledThinking, EffortLevel, EnabledThinking, Message, MessageContent, MessagesOptionalParams, MessagesRequest, MessagesTool, OutputConfig, Speed, SystemPrompt, ThinkingConfig, ThinkingDisplay, }; pub use response::MessagesResponse; +pub use tools::{ + AdvisorTool, AdvisorToolName, AllowedCaller, BashToolName, BrowserToolsetConfigs, + BuiltinMessagesTool, ClientTool, CodeExecutionToolName, ComputerTool, ComputerTool20251124, + ComputerToolName, ComputerToolsetConfigs, CustomTool, CustomToolType, McpToolset, + MemoryToolName, MessagesToolParam, ResponseInclusion, ServerTool, StrReplaceBasedEditToolName, + StrReplaceEditorName, TextEditorTool20250728, ToolChoice, ToolChoiceType, ToolResultUrlSource, + ToolSearchBm25ToolName, ToolSearchRegexToolName, Toolset, ToolsetToolConfig, + UrlSourceToolReference, UserInputUrlSource, UserLocationType, WebFetchTool, + WebFetchTool20260309, WebFetchTool20260318, WebFetchToolName, WebFetchUrlSources, + WebSearchTool, WebSearchTool20260318, WebSearchToolName, WebSearchUserLocation, +}; +pub use usage::{ + CacheCreationUsage, MessagesOutputTokensDetails, MessagesUsage, ServerToolUsage, + UsageIteration, UsageIterationType, +}; diff --git a/litellm-rust/crates/llms-types/src/formats/messages/request.rs b/litellm-rust/crates/llms-types/src/formats/messages/request.rs index 16cc55e4f5a..d065f45d41c 100644 --- a/litellm-rust/crates/llms-types/src/formats/messages/request.rs +++ b/litellm-rust/crates/llms-types/src/formats/messages/request.rs @@ -4,14 +4,14 @@ use strum::IntoStaticStr; use crate::formats::chat_completions::ReasoningEffort; use crate::recognized::Recognized; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(untagged)] pub enum SystemPrompt { Text(String), Blocks(Vec), } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(untagged)] pub enum MessageContent { Text(String), @@ -30,16 +30,24 @@ pub enum MessageContent { )] #[cfg_attr(feature = "schema", derive(schemars::JsonSchema))] #[cfg_attr(feature = "schema", schemars(from = "String", into = "String"))] -#[strum(serialize_all = "snake_case")] pub enum ContentBlockType { + #[strum(serialize = "text")] Text, + #[strum(serialize = "thinking")] Thinking, + #[strum(serialize = "redacted_thinking")] RedactedThinking, + #[strum(serialize = "tool_use")] ToolUse, + #[strum(serialize = "server_tool_use")] ServerToolUse, + #[strum(serialize = "tool_result")] ToolResult, + #[strum(serialize = "compaction")] Compaction, + #[strum(serialize = "advisor_tool_result")] AdvisorToolResult, + #[strum(serialize = "web_search_tool_result")] WebSearchToolResult, #[strum(default, transparent)] Other(String), @@ -57,7 +65,7 @@ impl From for String { } } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct ContentBlock { #[serde(rename = "type", default, skip_serializing_if = "Option::is_none")] @@ -102,7 +110,7 @@ impl ContentBlock { } } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct CacheControl { #[serde(rename = "type", skip_serializing_if = "Option::is_none")] @@ -115,7 +123,7 @@ pub struct CacheControl { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct Message { pub role: String, pub content: MessageContent, @@ -123,24 +131,22 @@ pub struct Message { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] -#[derive(Copy, Hash, IntoStaticStr, Eq)] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Hash, IntoStaticStr, Eq, strum::VariantArray)] #[serde(rename_all = "lowercase")] -#[strum(serialize_all = "lowercase")] pub enum EffortLevel { + #[strum(serialize = "low")] Low, + #[strum(serialize = "medium")] Medium, + #[strum(serialize = "high")] High, + #[strum(serialize = "xhigh")] Xhigh, + #[strum(serialize = "max")] Max, } -impl EffortLevel { - pub fn as_str(self) -> &'static str { - self.into() - } -} - impl From for ReasoningEffort { fn from(level: EffortLevel) -> Self { match level { @@ -153,24 +159,19 @@ impl From for ReasoningEffort { } } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Copy, IntoStaticStr, Eq)] #[serde(rename_all = "lowercase")] -#[strum(serialize_all = "lowercase")] pub enum Speed { + #[strum(serialize = "fast")] Fast, + #[strum(serialize = "standard")] Standard, } -impl Speed { - pub fn as_str(self) -> &'static str { - self.into() - } -} - /// The tools whose presence changes how the request is sent. Every other tool, custom or /// server, deserializes as `Recognized::Unrecognized` and passes through verbatim. -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(tag = "type")] pub enum MessagesTool { #[serde(rename = "advisor_20260301")] @@ -190,7 +191,22 @@ pub enum MessagesTool { }, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ContextTrigger { + InputTokens { + value: u64, + #[serde(flatten)] + extra: Map, + }, + ToolUses { + value: u64, + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(tag = "type")] pub enum ContextEdit { #[serde(rename = "compact_20260112")] @@ -210,7 +226,7 @@ pub enum ContextEdit { }, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct ContextManagement { #[serde(default, skip_serializing_if = "Option::is_none")] @@ -219,7 +235,7 @@ pub struct ContextManagement { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct OutputConfig { #[serde(default, skip_serializing_if = "Option::is_none")] @@ -236,7 +252,7 @@ impl OutputConfig { } } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Copy, Eq)] #[serde(rename_all = "lowercase")] pub enum ThinkingDisplay { @@ -245,7 +261,7 @@ pub enum ThinkingDisplay { Updates, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct EnabledThinking { #[serde(default, skip_serializing_if = "Option::is_none")] @@ -256,7 +272,7 @@ pub struct EnabledThinking { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct AdaptiveThinking { #[serde(default, skip_serializing_if = "Option::is_none")] @@ -265,14 +281,14 @@ pub struct AdaptiveThinking { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct DisabledThinking { #[serde(flatten)] pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(tag = "type", rename_all = "lowercase")] pub enum ThinkingConfig { Enabled(EnabledThinking), @@ -296,7 +312,7 @@ impl ThinkingConfig { } } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct MessagesRequest { pub model: String, pub messages: Vec, @@ -304,7 +320,7 @@ pub struct MessagesRequest { pub params: MessagesOptionalParams, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct MessagesOptionalParams { #[serde(skip_serializing_if = "Option::is_none")] @@ -676,7 +692,10 @@ mod tests { #[rstest] fn speed_names_match_the_wire(#[values(Speed::Fast, Speed::Standard)] speed: Speed) { - assert_eq!(serde_json::to_value(speed).unwrap(), json!(speed.as_str())); + assert_eq!( + serde_json::to_value(speed).unwrap(), + json!(<&'static str>::from(speed)) + ); } #[rstest] @@ -690,10 +709,13 @@ mod tests { )] level: EffortLevel, ) { - assert_eq!(serde_json::to_value(level).unwrap(), json!(level.as_str())); + assert_eq!( + serde_json::to_value(level).unwrap(), + json!(<&'static str>::from(level)) + ); assert_eq!( serde_json::to_value(ReasoningEffort::from(level)).unwrap(), - json!(level.as_str()) + json!(<&'static str>::from(level)) ); } } diff --git a/litellm-rust/crates/llms-types/src/formats/messages/response.rs b/litellm-rust/crates/llms-types/src/formats/messages/response.rs index 2d8e1c054fa..84966aa2a0a 100644 --- a/litellm-rust/crates/llms-types/src/formats/messages/response.rs +++ b/litellm-rust/crates/llms-types/src/formats/messages/response.rs @@ -1,6 +1,6 @@ use serde_json::{Map, Value}; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct MessagesResponse { pub id: String, #[serde(rename = "type")] diff --git a/litellm-rust/crates/llms-types/src/formats/messages/streaming.rs b/litellm-rust/crates/llms-types/src/formats/messages/streaming.rs index abdcfa26a8c..88f05edd6ed 100644 --- a/litellm-rust/crates/llms-types/src/formats/messages/streaming.rs +++ b/litellm-rust/crates/llms-types/src/formats/messages/streaming.rs @@ -1,6 +1,6 @@ use serde_json::{Map, Value}; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct MessagesStreamUsage { #[serde(default, skip_serializing_if = "Option::is_none")] @@ -17,7 +17,7 @@ pub struct MessagesStreamUsage { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct MessagesStreamMessage { pub id: String, #[serde(rename = "type")] @@ -32,7 +32,7 @@ pub struct MessagesStreamMessage { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(tag = "type", rename_all = "snake_case")] pub enum MessagesContentBlockDelta { TextDelta { @@ -56,7 +56,7 @@ pub enum MessagesContentBlockDelta { }, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct MessagesContentBlock { #[serde(rename = "type")] pub block_type: String, @@ -82,7 +82,7 @@ pub struct MessagesContentBlock { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct MessagesDelta { #[serde(default, skip_serializing_if = "Option::is_none")] @@ -97,7 +97,7 @@ pub struct MessagesDelta { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct MessagesStreamError { #[serde(rename = "type")] pub error_type: String, @@ -108,7 +108,7 @@ pub struct MessagesStreamError { pub extra: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(tag = "type", rename_all = "snake_case")] pub enum MessagesStreamEvent { MessageStart { diff --git a/litellm-rust/crates/llms-types/src/formats/messages/tools.rs b/litellm-rust/crates/llms-types/src/formats/messages/tools.rs new file mode 100644 index 00000000000..e6c5d6dc345 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/messages/tools.rs @@ -0,0 +1,564 @@ +use indexmap::IndexMap; +use serde_json::{Map, Value}; + +use super::{CacheControl, CitationsConfig, McpListedTool}; +use crate::json_schema::JsonSchema; + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +pub enum AllowedCaller { + #[serde(rename = "direct")] + Direct, + #[serde(rename = "code_execution_20250825")] + CodeExecution20250825, + #[serde(rename = "code_execution_20260120")] + CodeExecution20260120, + #[serde(rename = "code_execution_20260521")] + CodeExecution20260521, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct WebSearchUserLocation { + #[serde(rename = "type")] + pub location_type: UserLocationType, + pub city: Option, + pub country: Option, + pub region: Option, + pub timezone: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum UserLocationType { + Approximate, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ToolChoice { + #[serde(rename = "type")] + pub choice_type: ToolChoiceType, + pub name: Option, + pub disable_parallel_tool_use: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum ToolChoiceType { + Auto, + Any, + Tool, + None, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct CustomTool { + #[serde(rename = "type")] + pub tool_type: Option, + pub name: String, + pub input_schema: JsonSchema, + pub description: Option, + pub strict: Option, + pub cache_control: Option, + pub defer_loading: Option, + pub allowed_callers: Option>, + pub input_examples: Option>>, + pub eager_input_streaming: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum CustomToolType { + Custom, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum BashToolName { + Bash, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum StrReplaceEditorName { + StrReplaceEditor, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum StrReplaceBasedEditToolName { + StrReplaceBasedEditTool, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum MemoryToolName { + Memory, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum ComputerToolName { + Computer, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum CodeExecutionToolName { + CodeExecution, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum ToolSearchRegexToolName { + ToolSearchToolRegex, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum ToolSearchBm25ToolName { + ToolSearchToolBm25, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum WebSearchToolName { + WebSearch, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum WebFetchToolName { + WebFetch, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum AdvisorToolName { + Advisor, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ServerTool { + pub name: N, + pub allowed_callers: Option>, + pub cache_control: Option, + pub defer_loading: Option, + pub strict: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ClientTool { + pub name: N, + pub allowed_callers: Option>, + pub cache_control: Option, + pub defer_loading: Option, + pub input_examples: Option>>, + pub strict: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct TextEditorTool20250728 { + pub name: StrReplaceBasedEditToolName, + pub allowed_callers: Option>, + pub cache_control: Option, + pub defer_loading: Option, + pub input_examples: Option>>, + pub max_characters: Option, + pub strict: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ComputerTool { + pub name: ComputerToolName, + pub display_width_px: u64, + pub display_height_px: u64, + pub display_number: Option, + pub allowed_callers: Option>, + pub cache_control: Option, + pub defer_loading: Option, + pub input_examples: Option>>, + pub strict: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ComputerTool20251124 { + pub name: ComputerToolName, + pub display_width_px: u64, + pub display_height_px: u64, + pub display_number: Option, + pub enable_zoom: Option, + pub allowed_callers: Option>, + pub cache_control: Option, + pub defer_loading: Option, + pub input_examples: Option>>, + pub strict: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Copy, Eq)] +#[serde(rename_all = "snake_case")] +pub enum ResponseInclusion { + Full, + Excluded, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct WebSearchTool { + pub name: WebSearchToolName, + pub allowed_callers: Option>, + pub allowed_domains: Option>, + pub blocked_domains: Option>, + pub cache_control: Option, + pub defer_loading: Option, + pub max_uses: Option, + pub strict: Option, + pub user_location: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct WebSearchTool20260318 { + pub name: WebSearchToolName, + pub allowed_callers: Option>, + pub allowed_domains: Option>, + pub blocked_domains: Option>, + pub cache_control: Option, + pub defer_loading: Option, + pub max_uses: Option, + pub response_inclusion: Option, + pub strict: Option, + pub user_location: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum UrlSourceToolReference { + ToolReference { + name: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ToolResultUrlSource { + All { + #[serde(flatten)] + extra: Map, + }, + None { + #[serde(flatten)] + extra: Map, + }, + Only { + tools: Vec, + #[serde(flatten)] + extra: Map, + }, + Except { + tools: Vec, + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum UserInputUrlSource { + All { + #[serde(flatten)] + extra: Map, + }, + None { + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct WebFetchUrlSources { + pub client_tool_results: Option, + pub server_tool_results: Option, + pub user_input: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct WebFetchTool { + pub name: WebFetchToolName, + pub allowed_callers: Option>, + pub allowed_domains: Option>, + pub blocked_domains: Option>, + pub cache_control: Option, + pub citations: Option, + pub defer_loading: Option, + pub max_content_tokens: Option, + pub max_uses: Option, + pub strict: Option, + pub url_sources: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct WebFetchTool20260309 { + pub name: WebFetchToolName, + pub allowed_callers: Option>, + pub allowed_domains: Option>, + pub blocked_domains: Option>, + pub cache_control: Option, + pub citations: Option, + pub defer_loading: Option, + pub max_content_tokens: Option, + pub max_uses: Option, + pub strict: Option, + pub url_sources: Option, + pub use_cache: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct WebFetchTool20260318 { + pub name: WebFetchToolName, + pub allowed_callers: Option>, + pub allowed_domains: Option>, + pub blocked_domains: Option>, + pub cache_control: Option, + pub citations: Option, + pub defer_loading: Option, + pub max_content_tokens: Option, + pub max_uses: Option, + pub response_inclusion: Option, + pub strict: Option, + pub url_sources: Option, + pub use_cache: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct AdvisorTool { + pub name: AdvisorToolName, + pub model: String, + pub allowed_callers: Option>, + pub cache_control: Option, + pub caching: Option, + pub defer_loading: Option, + pub max_tokens: Option, + pub max_uses: Option, + pub strict: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ToolsetToolConfig { + pub defer_loading: Option, + pub enabled: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct McpToolset { + pub mcp_server_name: String, + pub cache_control: Option, + pub configs: Option>, + pub default_config: Option, + pub tools: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct BrowserToolsetConfigs { + #[serde(rename = "type")] + pub type_text: Option, + pub close_tab: Option, + pub double_click: Option, + pub file_upload: Option, + pub find: Option, + pub form_input: Option, + pub get_page_text: Option, + pub hold_key: Option, + pub hover: Option, + pub javascript_exec: Option, + pub key: Option, + pub left_click: Option, + pub left_click_drag: Option, + pub left_mouse_down: Option, + pub left_mouse_up: Option, + pub list_tabs: Option, + pub middle_click: Option, + pub mouse_move: Option, + pub navigate: Option, + pub new_tab: Option, + pub read_console: Option, + pub read_network: Option, + pub read_page: Option, + pub right_click: Option, + pub screenshot: Option, + pub scroll: Option, + pub scroll_to: Option, + pub switch_tab: Option, + pub triple_click: Option, + pub wait: Option, + pub zoom: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ComputerToolsetConfigs { + #[serde(rename = "type")] + pub type_text: Option, + pub cursor_position: Option, + pub double_click: Option, + pub hold_key: Option, + pub key: Option, + pub left_click: Option, + pub left_click_drag: Option, + pub left_mouse_down: Option, + pub left_mouse_up: Option, + pub middle_click: Option, + pub mouse_move: Option, + pub right_click: Option, + pub screenshot: Option, + pub scroll: Option, + pub triple_click: Option, + pub wait: Option, + pub zoom: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct Toolset { + pub cache_control: Option, + pub configs: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type")] +pub enum BuiltinMessagesTool { + #[serde(rename = "bash_20241022")] + Bash20241022(ClientTool), + #[serde(rename = "bash_20250124")] + Bash20250124(ClientTool), + #[serde(rename = "text_editor_20241022")] + TextEditor20241022(ClientTool), + #[serde(rename = "text_editor_20250124")] + TextEditor20250124(ClientTool), + #[serde(rename = "text_editor_20250429")] + TextEditor20250429(ClientTool), + #[serde(rename = "text_editor_20250728")] + TextEditor20250728(TextEditorTool20250728), + #[serde(rename = "memory_20250818")] + Memory20250818(ClientTool), + #[serde(rename = "computer_20241022")] + Computer20241022(ComputerTool), + #[serde(rename = "computer_20250124")] + Computer20250124(ComputerTool), + #[serde(rename = "computer_20251124")] + Computer20251124(ComputerTool20251124), + #[serde(rename = "code_execution_20250522")] + CodeExecution20250522(ServerTool), + #[serde(rename = "code_execution_20250825")] + CodeExecution20250825(ServerTool), + #[serde(rename = "code_execution_20260120")] + CodeExecution20260120(ServerTool), + #[serde(rename = "code_execution_20260521")] + CodeExecution20260521(ServerTool), + #[serde(rename = "tool_search_tool_regex_20251119")] + ToolSearchRegex20251119(ServerTool), + #[serde(rename = "tool_search_tool_regex")] + ToolSearchRegex(ServerTool), + #[serde(rename = "tool_search_tool_bm25_20251119")] + ToolSearchBm2520251119(ServerTool), + #[serde(rename = "tool_search_tool_bm25")] + ToolSearchBm25(ServerTool), + #[serde(rename = "web_search_20250305")] + WebSearch20250305(WebSearchTool), + #[serde(rename = "web_search_20260209")] + WebSearch20260209(WebSearchTool), + #[serde(rename = "web_search_20260318")] + WebSearch20260318(WebSearchTool20260318), + #[serde(rename = "web_fetch_20250910")] + WebFetch20250910(WebFetchTool), + #[serde(rename = "web_fetch_20260209")] + WebFetch20260209(WebFetchTool), + #[serde(rename = "web_fetch_20260309")] + WebFetch20260309(WebFetchTool20260309), + #[serde(rename = "web_fetch_20260318")] + WebFetch20260318(WebFetchTool20260318), + #[serde(rename = "advisor_20260301")] + Advisor20260301(AdvisorTool), + #[serde(rename = "browser_toolset_20260801")] + BrowserToolset20260801(Toolset), + #[serde(rename = "computer_toolset_20260801")] + ComputerToolset20260801(Toolset), + #[serde(rename = "mcp_toolset")] + McpToolset(McpToolset), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(untagged)] +pub enum MessagesToolParam { + Builtin(Box), + Custom(Box), +} diff --git a/litellm-rust/crates/llms-types/src/formats/messages/usage.rs b/litellm-rust/crates/llms-types/src/formats/messages/usage.rs new file mode 100644 index 00000000000..4ce0ef1d5c4 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/messages/usage.rs @@ -0,0 +1,75 @@ +use serde_json::{Map, Value}; + +use crate::recognized::Recognized; + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ServerToolUsage { + pub web_search_requests: Option, + pub web_fetch_requests: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct UsageIteration { + #[serde(rename = "type")] + pub iteration_type: Recognized, + pub input_tokens: Option, + pub output_tokens: Option, + pub cache_creation_input_tokens: Option, + pub cache_read_input_tokens: Option, + pub cache_creation: Option, + pub model: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum UsageIterationType { + Compaction, + Message, + AdvisorMessage, + FallbackMessage, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct MessagesUsage { + pub input_tokens: Option, + pub output_tokens: Option, + pub cache_creation_input_tokens: Option, + pub cache_read_input_tokens: Option, + pub server_tool_use: Option, + pub cache_creation: Option, + pub output_tokens_details: Option, + pub service_tier: Option, + pub inference_geo: Option, + pub speed: Option, + pub iterations: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct CacheCreationUsage { + pub ephemeral_1h_input_tokens: Option, + pub ephemeral_5m_input_tokens: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct MessagesOutputTokensDetails { + pub thinking_tokens: Option, + #[serde(flatten)] + pub extra: Map, +} diff --git a/litellm-rust/crates/llms-types/src/formats/mod.rs b/litellm-rust/crates/llms-types/src/formats/mod.rs index 53f2577090b..0353077c3a9 100644 --- a/litellm-rust/crates/llms-types/src/formats/mod.rs +++ b/litellm-rust/crates/llms-types/src/formats/mod.rs @@ -1,3 +1,6 @@ +//! Wire types for each provider API format (OpenAI Chat Completions, +//! Anthropic Messages, Responses, ...), one submodule per format. + pub mod audio_transcription; pub mod batches; pub mod chat_completions; diff --git a/litellm-rust/crates/llms-types/src/formats/ocr.rs b/litellm-rust/crates/llms-types/src/formats/ocr.rs index b491f5d82f1..68c991a6e8b 100644 --- a/litellm-rust/crates/llms-types/src/formats/ocr.rs +++ b/litellm-rust/crates/llms-types/src/formats/ocr.rs @@ -5,7 +5,7 @@ use serde_with::serde_as; use crate::serde_compat::{FiniteF64, LaxI64}; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(tag = "type")] pub enum OcrDocument { #[serde(rename = "document_url")] @@ -49,7 +49,7 @@ impl OcrDocument { } } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Copy, Default, Eq)] #[serde(rename_all = "lowercase")] pub enum OcrResponseFormat { @@ -59,7 +59,7 @@ pub enum OcrResponseFormat { } #[serde_as] -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct OcrPageDimensions { #[serde_as(deserialize_as = "Option")] @@ -70,7 +70,7 @@ pub struct OcrPageDimensions { pub width: Option, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct OcrPageImage { pub image_base64: Option, @@ -80,7 +80,7 @@ pub struct OcrPageImage { } #[serde_as] -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct OcrPage { #[serde_as(deserialize_as = "LaxI64")] @@ -93,7 +93,7 @@ pub struct OcrPage { } #[serde_as] -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct OcrUsageInfo { #[serde_as(deserialize_as = "Option")] @@ -108,7 +108,7 @@ pub struct OcrUsageInfo { pub extra_fields: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct LiteLLMOcrResponse { pub pages: Vec, pub model: String, @@ -150,3 +150,15 @@ impl LiteLLMOcrResponse { fn ocr_object() -> String { "ocr".into() } + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct OcrBoundingBox { + pub top_left_x: Option, + pub top_left_y: Option, + pub bottom_right_x: Option, + pub bottom_right_y: Option, + #[serde(flatten)] + pub extra: Map, +} diff --git a/litellm-rust/crates/llms-types/src/formats/responses/mod.rs b/litellm-rust/crates/llms-types/src/formats/responses/mod.rs index 0aefd8a8698..a45a971bae6 100644 --- a/litellm-rust/crates/llms-types/src/formats/responses/mod.rs +++ b/litellm-rust/crates/llms-types/src/formats/responses/mod.rs @@ -1,4 +1,6 @@ +mod output; mod response; pub mod streaming_websocket; +pub use output::*; pub use response::ResponsesApiResponse; diff --git a/litellm-rust/crates/llms-types/src/formats/responses/output.rs b/litellm-rust/crates/llms-types/src/formats/responses/output.rs new file mode 100644 index 00000000000..ba4c011d458 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/formats/responses/output.rs @@ -0,0 +1,831 @@ +use crate::formats::chat_completions::PromptCacheBreakpoint; +use serde_json::{Map, Value}; +use std::collections::BTreeMap; + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesOutputItem { + Message(ResponsesMessage), + FunctionCall(ResponsesFunctionCall), + CustomToolCall(ResponsesCustomToolCall), + Reasoning(ResponsesReasoning), + WebSearchCall(ResponsesWebSearchCall), + FileSearchCall(ResponsesFileSearchCall), + ImageGenerationCall(ResponsesImageGenerationCall), + CodeInterpreterCall(ResponsesCodeInterpreterCall), + McpCall(ResponsesMcpCall), + McpListTools(ResponsesMcpListTools), + FunctionCallOutput(ResponsesFunctionCallOutput), + ComputerCall(ResponsesComputerCall), + ComputerCallOutput(ResponsesComputerCallOutput), + Program(ResponsesProgram), + ProgramOutput(ResponsesProgramOutput), + ToolSearchCall(ResponsesToolSearchCall), + ToolSearchOutput(ResponsesToolSearchOutput), + AdditionalTools(ResponsesAdditionalTools), + Compaction(ResponsesCompaction), + LocalShellCall(ResponsesLocalShellCall), + LocalShellCallOutput(ResponsesLocalShellCallOutput), + ShellCall(ResponsesShellCall), + ShellCallOutput(ResponsesShellCallOutput), + ApplyPatchCall(ResponsesApplyPatchCall), + ApplyPatchCallOutput(ResponsesApplyPatchCallOutput), + McpApprovalRequest(ResponsesMcpApprovalRequest), + McpApprovalResponse(ResponsesMcpApprovalResponse), + CustomToolCallOutput(ResponsesCustomToolCallOutput), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesContentPart { + OutputText { + text: String, + #[serde(skip_serializing_if = "Option::is_none")] + annotations: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + logprobs: Option>, + #[serde(flatten)] + extra: Map, + }, + Refusal { + refusal: String, + #[serde(flatten)] + extra: Map, + }, + SummaryText { + text: String, + #[serde(flatten)] + extra: Map, + }, + ReasoningText { + text: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesAnnotation { + UrlCitation(ResponsesUrlCitation), + FileCitation(ResponsesFileCitation), + FilePath(ResponsesFileCitation), + ContainerFileCitation(ResponsesFileCitation), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesCodeOutput { + Logs { + logs: String, + #[serde(flatten)] + extra: Map, + }, + Image { + url: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesMessage { + pub id: Option, + pub status: Option, + pub role: Option, + pub phase: Option, + pub content: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesFunctionCall { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub name: Option, + pub arguments: Option, + pub namespace: Option, + pub r#async: Option, + pub caller: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesCustomToolCall { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub name: Option, + pub input: Option, + pub namespace: Option, + pub r#async: Option, + pub caller: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesReasoning { + pub id: Option, + pub status: Option, + pub summary: Option>, + pub content: Option>, + pub encrypted_content: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesWebSearchCall { + pub id: Option, + pub status: Option, + pub action: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesFileSearchCall { + pub id: Option, + pub status: Option, + pub queries: Option>, + pub results: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesFileSearchResult { + pub file_id: Option, + pub filename: Option, + pub score: Option, + pub text: Option, + pub attributes: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesImageGenerationCall { + pub id: Option, + pub status: Option, + pub result: Option, + pub action: Option, + pub background: Option, + pub output_format: Option, + pub quality: Option, + pub revised_prompt: Option, + pub size: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesCodeInterpreterCall { + pub id: Option, + pub status: Option, + pub code: Option, + pub container_id: Option, + pub outputs: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesMcpCall { + pub id: Option, + pub status: Option, + pub name: Option, + pub server_label: Option, + pub arguments: Option, + pub output: Option, + pub error: Option, + pub approval_request_id: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesUrlCitation { + pub url: Option, + pub title: Option, + pub start_index: Option, + pub end_index: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesFileCitation { + pub file_id: Option, + pub filename: Option, + pub index: Option, + pub start_index: Option, + pub end_index: Option, + pub container_id: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesWebSearchAction { + Search { + #[serde(skip_serializing_if = "Option::is_none")] + query: Option, + #[serde(skip_serializing_if = "Option::is_none")] + queries: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + sources: Option>, + #[serde(flatten)] + extra: Map, + }, + OpenPage { + #[serde(skip_serializing_if = "Option::is_none")] + url: Option, + #[serde(flatten)] + extra: Map, + }, + #[serde(alias = "find")] + FindInPage { + url: String, + pattern: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesWebSearchSource { + #[serde(rename = "type")] + pub source_type: Option, + pub url: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesMcpListTools { + pub id: Option, + pub server_label: Option, + pub tools: Option>, + pub error: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesMcpTool { + pub name: Option, + pub description: Option, + pub input_schema: Option, + pub annotations: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(untagged)] +pub enum ResponsesMcpError { + Message(String), + Detail(ResponsesMcpErrorDetail), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesMcpErrorDetail { + McpProtocolError { + code: i64, + message: String, + #[serde(flatten)] + extra: Map, + }, + HttpError { + code: i64, + message: String, + #[serde(flatten)] + extra: Map, + }, + McpToolExecutionError { + content: Value, + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesToolCaller { + Direct { + #[serde(flatten)] + extra: Map, + }, + Program { + caller_id: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(untagged)] +pub enum ResponsesToolOutput { + Text(String), + Content(Vec), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesInputContent { + InputText { + text: String, + #[serde(skip_serializing_if = "Option::is_none")] + prompt_cache_breakpoint: Option, + #[serde(flatten)] + extra: Map, + }, + InputImage { + detail: String, + #[serde(skip_serializing_if = "Option::is_none")] + file_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + image_url: Option, + #[serde(skip_serializing_if = "Option::is_none")] + prompt_cache_breakpoint: Option, + #[serde(flatten)] + extra: Map, + }, + InputFile { + #[serde(skip_serializing_if = "Option::is_none")] + detail: Option, + #[serde(skip_serializing_if = "Option::is_none")] + file_data: Option, + #[serde(skip_serializing_if = "Option::is_none")] + file_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + file_url: Option, + #[serde(skip_serializing_if = "Option::is_none")] + filename: Option, + #[serde(skip_serializing_if = "Option::is_none")] + prompt_cache_breakpoint: Option, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesFunctionCallOutput { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub output: Option, + pub caller: Option, + pub created_by: Option, + pub name: Option, + pub namespace: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesCustomToolCallOutput { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub output: Option, + pub caller: Option, + pub created_by: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesSafetyCheck { + pub id: Option, + pub code: Option, + pub message: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesComputerCall { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub pending_safety_checks: Option>, + pub action: Option, + pub actions: Option>, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesComputerAction { + Click { + button: String, + x: i64, + y: i64, + #[serde(skip_serializing_if = "Option::is_none")] + keys: Option>, + #[serde(flatten)] + extra: Map, + }, + DoubleClick { + x: i64, + y: i64, + #[serde(skip_serializing_if = "Option::is_none")] + keys: Option>, + #[serde(flatten)] + extra: Map, + }, + Drag { + path: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + keys: Option>, + #[serde(flatten)] + extra: Map, + }, + Keypress { + keys: Vec, + #[serde(flatten)] + extra: Map, + }, + Move { + x: i64, + y: i64, + #[serde(skip_serializing_if = "Option::is_none")] + keys: Option>, + #[serde(flatten)] + extra: Map, + }, + Screenshot { + #[serde(flatten)] + extra: Map, + }, + Scroll { + scroll_x: i64, + scroll_y: i64, + x: i64, + y: i64, + #[serde(skip_serializing_if = "Option::is_none")] + keys: Option>, + #[serde(flatten)] + extra: Map, + }, + Type { + text: String, + #[serde(flatten)] + extra: Map, + }, + Wait { + #[serde(flatten)] + extra: Map, + }, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct ResponsesCoordinate { + pub x: i64, + pub y: i64, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesComputerCallOutput { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub output: Option, + pub acknowledged_safety_checks: Option>, + pub created_by: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesComputerOutput { + ComputerScreenshot { + #[serde(skip_serializing_if = "Option::is_none")] + file_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + image_url: Option, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesProgram { + pub id: Option, + pub call_id: Option, + pub code: Option, + pub fingerprint: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesProgramOutput { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub result: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesToolSearchCall { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub arguments: Option, + pub execution: Option, + pub created_by: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesToolSearchOutput { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub execution: Option, + pub tools: Option>>, + pub created_by: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesAdditionalTools { + pub id: Option, + pub role: Option, + pub tools: Option>>, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesCompaction { + pub id: Option, + pub encrypted_content: Option, + pub created_by: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesLocalShellCall { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub action: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesLocalShellAction { + Exec { + command: Vec, + env: BTreeMap, + #[serde(skip_serializing_if = "Option::is_none")] + timeout_ms: Option, + #[serde(skip_serializing_if = "Option::is_none")] + user: Option, + #[serde(skip_serializing_if = "Option::is_none")] + working_directory: Option, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesLocalShellCallOutput { + pub id: Option, + pub status: Option, + pub output: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesShellCall { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub action: Option, + pub environment: Option, + pub caller: Option, + pub created_by: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesShellAction { + pub commands: Option>, + pub max_output_length: Option, + pub timeout_ms: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesShellEnvironment { + Local { + #[serde(flatten)] + extra: Map, + }, + ContainerReference { + container_id: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesShellCallOutput { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub max_output_length: Option, + pub output: Option>, + pub caller: Option, + pub created_by: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesShellOutputChunk { + pub outcome: Option, + pub stdout: Option, + pub stderr: Option, + pub created_by: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesShellOutcome { + Timeout { + #[serde(flatten)] + extra: Map, + }, + Exit { + exit_code: i64, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesApplyPatchCall { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub operation: Option, + pub caller: Option, + pub created_by: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponsesApplyPatchOperation { + CreateFile { + path: String, + diff: String, + #[serde(flatten)] + extra: Map, + }, + DeleteFile { + path: String, + #[serde(flatten)] + extra: Map, + }, + UpdateFile { + path: String, + diff: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesApplyPatchCallOutput { + pub id: Option, + pub status: Option, + pub call_id: Option, + pub output: Option, + pub caller: Option, + pub created_by: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesMcpApprovalRequest { + pub id: Option, + pub server_label: Option, + pub name: Option, + pub arguments: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesMcpApprovalResponse { + pub id: Option, + pub approval_request_id: Option, + pub approve: Option, + pub reason: Option, + #[serde(flatten)] + pub extra: Map, +} diff --git a/litellm-rust/crates/llms-types/src/formats/responses/response.rs b/litellm-rust/crates/llms-types/src/formats/responses/response.rs index 7017d0fa4e4..813294b0f1f 100644 --- a/litellm-rust/crates/llms-types/src/formats/responses/response.rs +++ b/litellm-rust/crates/llms-types/src/formats/responses/response.rs @@ -1,6 +1,6 @@ use serde_json::{Map, Value}; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct ResponsesApiResponse { pub id: String, pub model: String, diff --git a/litellm-rust/crates/llms-types/src/formats/responses/streaming_websocket.rs b/litellm-rust/crates/llms-types/src/formats/responses/streaming_websocket.rs index 6b648a21bd1..6053eea3e14 100644 --- a/litellm-rust/crates/llms-types/src/formats/responses/streaming_websocket.rs +++ b/litellm-rust/crates/llms-types/src/formats/responses/streaming_websocket.rs @@ -1,3 +1,5 @@ +use super::ResponsesOutputItem; +use crate::recognized::Recognized; use serde_json::{Map, Value}; #[derive( @@ -36,7 +38,7 @@ impl ResponsesWsEventType { } } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] pub struct ResponsesWsEvent { #[serde(rename = "type")] pub event_type: ResponsesWsEventType, @@ -62,7 +64,7 @@ impl ResponsesWsEvent { } } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Eq)] pub struct ResponsesErrorFrame { #[serde(rename = "type")] @@ -82,7 +84,7 @@ impl ResponsesErrorFrame { } } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Eq)] pub struct ResponsesErrorBody { #[serde(rename = "type")] @@ -90,6 +92,18 @@ pub struct ResponsesErrorBody { pub message: String, } +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct ResponsesEventResponse { + pub id: Option, + pub model: Option, + pub status: Option, + pub output: Option>>, + #[serde(flatten)] + pub extra: Map, +} + #[cfg(test)] mod tests { use rstest::rstest; diff --git a/litellm-rust/crates/llms-types/src/headers.rs b/litellm-rust/crates/llms-types/src/headers.rs index bf4f42b493d..eefb76280f7 100644 --- a/litellm-rust/crates/llms-types/src/headers.rs +++ b/litellm-rust/crates/llms-types/src/headers.rs @@ -1,6 +1,6 @@ use serde_json::{Map, Value}; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Default)] pub struct ProviderSpecificHeader { #[serde(default)] @@ -9,7 +9,7 @@ pub struct ProviderSpecificHeader { pub extra_headers: Map, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(untagged)] pub enum ProviderSpecificHeaders { One(ProviderSpecificHeader), diff --git a/litellm-rust/crates/llms-types/src/json_schema.rs b/litellm-rust/crates/llms-types/src/json_schema.rs new file mode 100644 index 00000000000..50a96ebef09 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/json_schema.rs @@ -0,0 +1,51 @@ +use indexmap::IndexMap; +use serde_json::{Map, Value}; + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(untagged)] +pub enum JsonSchema { + Boolean(bool), + Object(Box), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(untagged)] +pub enum JsonSchemaType { + Name(String), + Names(Vec), +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(untagged)] +pub enum JsonSchemaItems { + Schema(Box), + Tuple(Vec), +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct JsonSchemaObject { + #[serde(rename = "type")] + pub schema_type: Option, + pub properties: Option>, + pub required: Option>, + #[serde(rename = "additionalProperties")] + pub additional_properties: Option>, + pub items: Option, + #[serde(rename = "prefixItems")] + pub prefix_items: Option>, + #[serde(rename = "$defs")] + pub defs: Option>, + #[serde(rename = "$ref")] + pub reference: Option, + #[serde(rename = "anyOf")] + pub any_of: Option>, + #[serde(rename = "allOf")] + pub all_of: Option>, + #[serde(rename = "oneOf")] + pub one_of: Option>, + pub strict: Option, + #[serde(flatten)] + pub extra: Map, +} diff --git a/litellm-rust/crates/llms-types/src/lib.rs b/litellm-rust/crates/llms-types/src/lib.rs index 116c11c0f88..814122f2ab9 100644 --- a/litellm-rust/crates/llms-types/src/lib.rs +++ b/litellm-rust/crates/llms-types/src/lib.rs @@ -6,6 +6,7 @@ macro_rules_attribute::attribute_alias! { pub mod formats; pub mod headers; +pub mod json_schema; pub mod providers; pub mod recognized; pub mod serde_compat; diff --git a/litellm-rust/crates/llms-types/src/providers/AGENTS.md b/litellm-rust/crates/llms-types/src/providers/AGENTS.md new file mode 100644 index 00000000000..128d29652d1 --- /dev/null +++ b/litellm-rust/crates/llms-types/src/providers/AGENTS.md @@ -0,0 +1,18 @@ +# references + +These links describe the provider-specific wire types in this directory. Shared Messages payload references live in `../formats/messages/AGENTS.md` + +## anthropic.rs + +`AnthropicBeta`, `BetaSet` and `BetaProvider` represent the `anthropic-beta` header values, their +wire spelling and per-host support + +- https://platform.claude.com/docs/en/api/beta-headers.md + +## minimax.rs + +MiniMax's Messages reference documents its image, video, and mid-conversation system content extensions. Its cache reference documents `cache_control` on those blocks + +- https://platform.minimax.io/docs/api-reference/text-chat-anthropic +- https://platform.minimax.io/docs/api-reference/text-chat-anthropic.md +- https://platform.minimax.io/docs/api-reference/anthropic-api-compatible-cache.md diff --git a/litellm-rust/crates/llms-types/src/providers/anthropic.rs b/litellm-rust/crates/llms-types/src/providers/anthropic.rs index 3e0c4369b11..97bb26a4088 100644 --- a/litellm-rust/crates/llms-types/src/providers/anthropic.rs +++ b/litellm-rust/crates/llms-types/src/providers/anthropic.rs @@ -7,6 +7,39 @@ use std::{ str::FromStr, }; +use strum::VariantArray; + +/// A provider column of `litellm/anthropic_beta_headers_config.json`: which betas a host accepts +/// and under which name. +#[derive( + Clone, + Copy, + Debug, + PartialEq, + Eq, + Hash, + strum::AsRefStr, + strum::Display, + strum::EnumString, + strum::VariantArray, +)] +pub enum BetaProvider { + #[strum(serialize = "anthropic")] + Anthropic, + #[strum(serialize = "azure_ai")] + AzureAi, + #[strum(serialize = "bedrock_converse")] + BedrockConverse, + #[strum(serialize = "bedrock")] + Bedrock, + #[strum(serialize = "bedrock_mantle")] + BedrockMantle, + #[strum(serialize = "vertex_ai")] + VertexAi, + #[strum(serialize = "databricks")] + Databricks, +} + /// One value of the `anthropic-beta` header. Equality, ordering and hashing follow the wire /// string, so a value parsed from a caller's header never disagrees with the matching variant. #[derive(Clone, Debug, strum::AsRefStr, strum::Display, strum::EnumString)] @@ -35,12 +68,68 @@ pub enum AnthropicBeta { PerTurnControl20260701, #[strum(serialize = "dangerous-tool-use-2026-09-03")] DangerousToolUse20260903, + #[strum(serialize = "bash_20241022")] + Bash20241022, + #[strum(serialize = "bash_20250124")] + Bash20250124, + #[strum(serialize = "claude-code-20250219")] + ClaudeCode20250219, + #[strum(serialize = "code-execution-2025-08-25")] + CodeExecution20250825, + #[strum(serialize = "computer-use-2025-01-24")] + ComputerUse20250124, + #[strum(serialize = "computer-use-2025-11-24")] + ComputerUse20251124, + #[strum(serialize = "context-1m-2025-08-07")] + Context1m20250807, + #[strum(serialize = "effort-2025-11-24")] + Effort20251124, + #[strum(serialize = "files-api-2025-04-14")] + FilesApi20250414, + #[strum(serialize = "fine-grained-tool-streaming-2025-05-14")] + FineGrainedToolStreaming20250514, + #[strum(serialize = "interleaved-thinking-2025-05-14")] + InterleavedThinking20250514, + #[strum(serialize = "mcp-client-2025-04-04")] + McpClient20250404, + #[strum(serialize = "mcp-client-2025-11-20")] + McpClient20251120, + #[strum(serialize = "mcp-servers-2025-12-04")] + McpServers20251204, + #[strum(serialize = "mid-conversation-output-config-2026-07-01")] + MidConversationOutputConfig20260701, + #[strum(serialize = "mid-conversation-tool-changes-2026-07-01")] + MidConversationToolChanges20260701, + #[strum(serialize = "output-128k-2025-02-19")] + Output128k20250219, + #[strum(serialize = "prompt-caching-scope-2026-01-05")] + PromptCachingScope20260105, + #[strum(serialize = "skills-2025-10-02")] + Skills20251002, + #[strum(serialize = "structured-output-2024-03-01")] + StructuredOutput20240301, + #[strum(serialize = "text_editor_20241022")] + TextEditor20241022, + #[strum(serialize = "text_editor_20250124")] + TextEditor20250124, + #[strum(serialize = "thinking-binding-controls-2026-08-01")] + ThinkingBindingControls20260801, + #[strum(serialize = "thinking-display-updates-2026-08-18")] + ThinkingDisplayUpdates20260818, + #[strum(serialize = "token-efficient-tools-2025-02-19")] + TokenEfficientTools20250219, + #[strum(serialize = "tool-examples-2025-10-29")] + ToolExamples20251029, + #[strum(serialize = "tool-search-tool-2025-10-19")] + ToolSearchTool20251019, + #[strum(serialize = "inline-tools-2026-09-15")] + InlineTools20260915, #[strum(default, transparent)] Other(String), } impl AnthropicBeta { - pub const KNOWN: [Self; 12] = [ + pub const KNOWN: [Self; 40] = [ Self::Oauth20250420, Self::WebFetch20250910, Self::WebSearch20250305, @@ -53,11 +142,153 @@ impl AnthropicBeta { Self::AdvisorTool20260301, Self::PerTurnControl20260701, Self::DangerousToolUse20260903, + Self::Bash20241022, + Self::Bash20250124, + Self::ClaudeCode20250219, + Self::CodeExecution20250825, + Self::ComputerUse20250124, + Self::ComputerUse20251124, + Self::Context1m20250807, + Self::Effort20251124, + Self::FilesApi20250414, + Self::FineGrainedToolStreaming20250514, + Self::InterleavedThinking20250514, + Self::McpClient20250404, + Self::McpClient20251120, + Self::McpServers20251204, + Self::MidConversationOutputConfig20260701, + Self::MidConversationToolChanges20260701, + Self::Output128k20250219, + Self::PromptCachingScope20260105, + Self::Skills20251002, + Self::StructuredOutput20240301, + Self::TextEditor20241022, + Self::TextEditor20250124, + Self::ThinkingBindingControls20260801, + Self::ThinkingDisplayUpdates20260818, + Self::TokenEfficientTools20250219, + Self::ToolExamples20251029, + Self::ToolSearchTool20251019, + Self::InlineTools20260915, ]; pub fn as_str(&self) -> &str { self.as_ref() } + + pub fn on(&self, provider: BetaProvider) -> Option { + let resolved = match self { + Self::Other(raw) => raw.parse().unwrap(), + _ => self.clone(), + }; + if resolved.rejected_by().contains(&provider) { + return None; + } + Some(match (resolved, provider) { + ( + Self::AdvancedToolUse20251120, + BetaProvider::Bedrock | BetaProvider::BedrockMantle | BetaProvider::VertexAi, + ) => Self::ToolSearchTool20251019, + (beta, _) => beta, + }) + } + + fn rejected_by(&self) -> &'static [BetaProvider] { + match self { + Self::AdvancedToolUse20251120 | Self::ContextManagement20250627 => { + &[BetaProvider::BedrockConverse] + } + Self::AdvisorTool20260301 | Self::Compact20260904 => &[ + BetaProvider::AzureAi, + BetaProvider::BedrockConverse, + BetaProvider::Bedrock, + BetaProvider::BedrockMantle, + BetaProvider::VertexAi, + BetaProvider::Databricks, + ], + Self::Bash20241022 + | Self::Bash20250124 + | Self::McpServers20251204 + | Self::StructuredOutput20240301 + | Self::TextEditor20241022 + | Self::TextEditor20250124 => BetaProvider::VARIANTS, + Self::ClaudeCode20250219 | Self::ToolExamples20251029 => &[ + BetaProvider::Anthropic, + BetaProvider::AzureAi, + BetaProvider::BedrockConverse, + BetaProvider::Bedrock, + BetaProvider::VertexAi, + BetaProvider::Databricks, + ], + Self::CodeExecution20250825 + | Self::FilesApi20250414 + | Self::McpClient20250404 + | Self::McpClient20251120 + | Self::PromptCachingScope20260105 + | Self::Skills20251002 + | Self::WebFetch20250910 => &[ + BetaProvider::BedrockConverse, + BetaProvider::Bedrock, + BetaProvider::BedrockMantle, + BetaProvider::VertexAi, + ], + Self::Compact20260112 => &[BetaProvider::AzureAi, BetaProvider::BedrockConverse], + Self::ComputerUse20250124 | Self::ComputerUse20251124 | Self::Context1m20250807 => &[], + Self::DangerousToolUse20260903 => { + &[BetaProvider::BedrockConverse, BetaProvider::Databricks] + } + Self::Effort20251124 => &[BetaProvider::VertexAi], + Self::FastMode20260201 | Self::Oauth20250420 => &[ + BetaProvider::AzureAi, + BetaProvider::BedrockConverse, + BetaProvider::Bedrock, + BetaProvider::BedrockMantle, + BetaProvider::VertexAi, + ], + Self::FineGrainedToolStreaming20250514 => { + &[BetaProvider::AzureAi, BetaProvider::VertexAi] + } + Self::InterleavedThinking20250514 | Self::WebSearch20250305 => { + &[BetaProvider::BedrockConverse, BetaProvider::Bedrock] + } + Self::MidConversationOutputConfig20260701 + | Self::MidConversationToolChanges20260701 => &[ + BetaProvider::AzureAi, + BetaProvider::BedrockConverse, + BetaProvider::BedrockMantle, + BetaProvider::VertexAi, + BetaProvider::Databricks, + ], + Self::Output128k20250219 | Self::TokenEfficientTools20250219 => &[ + BetaProvider::AzureAi, + BetaProvider::BedrockConverse, + BetaProvider::Bedrock, + BetaProvider::VertexAi, + ], + Self::PerTurnControl20260701 => &[ + BetaProvider::BedrockConverse, + BetaProvider::Bedrock, + BetaProvider::Databricks, + ], + Self::StructuredOutputs20251113 => &[BetaProvider::Bedrock, BetaProvider::VertexAi], + Self::ThinkingBindingControls20260801 => &[BetaProvider::AzureAi], + Self::ThinkingDisplayUpdates20260818 => &[BetaProvider::Databricks], + Self::ToolSearchTool20251019 => &[ + BetaProvider::Anthropic, + BetaProvider::AzureAi, + BetaProvider::BedrockConverse, + BetaProvider::Databricks, + ], + Self::InlineTools20260915 => &[ + BetaProvider::AzureAi, + BetaProvider::BedrockConverse, + BetaProvider::Bedrock, + BetaProvider::BedrockMantle, + BetaProvider::Databricks, + ], + Self::Other(_) => BetaProvider::VARIANTS, + } + } } impl PartialEq for AnthropicBeta { @@ -149,10 +380,20 @@ impl fmt::Display for BetaSet { #[cfg(test)] mod tests { + use std::collections::BTreeSet; + + use indexmap::IndexMap; use rstest::rstest; use super::*; + fn beta_headers_config() -> IndexMap { + serde_json::from_str(include_str!( + "../../../../../litellm/anthropic_beta_headers_config.json" + )) + .unwrap() + } + fn set(header: &str) -> BetaSet { header.parse().unwrap_or_else(|never| match never {}) } @@ -171,7 +412,35 @@ mod tests { AnthropicBeta::FastMode20260201, AnthropicBeta::AdvisorTool20260301, AnthropicBeta::PerTurnControl20260701, - AnthropicBeta::DangerousToolUse20260903 + AnthropicBeta::DangerousToolUse20260903, + AnthropicBeta::Bash20241022, + AnthropicBeta::Bash20250124, + AnthropicBeta::ClaudeCode20250219, + AnthropicBeta::CodeExecution20250825, + AnthropicBeta::ComputerUse20250124, + AnthropicBeta::ComputerUse20251124, + AnthropicBeta::Context1m20250807, + AnthropicBeta::Effort20251124, + AnthropicBeta::FilesApi20250414, + AnthropicBeta::FineGrainedToolStreaming20250514, + AnthropicBeta::InterleavedThinking20250514, + AnthropicBeta::McpClient20250404, + AnthropicBeta::McpClient20251120, + AnthropicBeta::McpServers20251204, + AnthropicBeta::MidConversationOutputConfig20260701, + AnthropicBeta::MidConversationToolChanges20260701, + AnthropicBeta::Output128k20250219, + AnthropicBeta::PromptCachingScope20260105, + AnthropicBeta::Skills20251002, + AnthropicBeta::StructuredOutput20240301, + AnthropicBeta::TextEditor20241022, + AnthropicBeta::TextEditor20250124, + AnthropicBeta::ThinkingBindingControls20260801, + AnthropicBeta::ThinkingDisplayUpdates20260818, + AnthropicBeta::TokenEfficientTools20250219, + AnthropicBeta::ToolExamples20251029, + AnthropicBeta::ToolSearchTool20251019, + AnthropicBeta::InlineTools20260915 )] beta: AnthropicBeta, ) { @@ -181,6 +450,119 @@ mod tests { assert!(AnthropicBeta::KNOWN.contains(&beta)); } + #[rstest] + fn known_betas_are_exactly_the_config_keys() { + let config = beta_headers_config(); + let config_keys: BTreeSet = config + .iter() + .filter(|(provider, _)| provider.as_str() != "description") + .flat_map(|(_, column)| { + column + .as_object() + .into_iter() + .flat_map(|values| values.keys().cloned()) + }) + .collect(); + let known_keys: BTreeSet = AnthropicBeta::KNOWN + .iter() + .map(|beta| beta.as_str().to_string()) + .collect(); + assert_eq!(known_keys, config_keys); + } + + #[rstest] + #[case::anthropic(BetaProvider::Anthropic)] + #[case::azure_ai(BetaProvider::AzureAi)] + #[case::bedrock_converse(BetaProvider::BedrockConverse)] + #[case::bedrock(BetaProvider::Bedrock)] + #[case::bedrock_mantle(BetaProvider::BedrockMantle)] + #[case::vertex_ai(BetaProvider::VertexAi)] + #[case::databricks(BetaProvider::Databricks)] + fn provider_columns_match_beta_provider_variants(#[case] provider: BetaProvider) { + let config = beta_headers_config(); + let config_columns: Vec = config + .keys() + .filter(|column| column.as_str() != "description") + .cloned() + .collect(); + let beta_provider_columns: Vec = BetaProvider::VARIANTS + .iter() + .map(ToString::to_string) + .collect(); + assert_eq!(config_columns, beta_provider_columns); + assert_eq!( + provider.to_string().parse::().unwrap(), + provider + ); + } + + #[rstest] + fn on_matches_every_config_cell() { + let config = beta_headers_config(); + for provider in BetaProvider::VARIANTS { + for beta in &AnthropicBeta::KNOWN { + let expected = config + .get(provider.as_ref()) + .and_then(serde_json::Value::as_object) + .and_then(|column| column.get(beta.as_str())) + .and_then(serde_json::Value::as_str) + .map(str::to_string); + let actual = beta.on(*provider).map(|name| name.to_string()); + assert_eq!(actual, expected, "provider {provider}, beta {beta}"); + if matches!(beta, AnthropicBeta::AdvancedToolUse20251120) + && matches!( + provider, + BetaProvider::Bedrock + | BetaProvider::BedrockMantle + | BetaProvider::VertexAi + ) + { + let parsed: AnthropicBeta = actual.as_deref().unwrap().parse().unwrap(); + assert!(AnthropicBeta::KNOWN.contains(&parsed)); + assert!(!matches!(parsed, AnthropicBeta::Other(_))); + } + } + } + } + + #[rstest] + #[case::renamed( + "advanced-tool-use-2025-11-20", + BetaProvider::Bedrock, + Some(AnthropicBeta::ToolSearchTool20251019) + )] + #[case::kept( + "oauth-2025-04-20", + BetaProvider::Anthropic, + Some(AnthropicBeta::Oauth20250420) + )] + #[case::rejected("effort-2025-11-24", BetaProvider::VertexAi, None)] + #[case::unknown("example-beta-2099-01-01", BetaProvider::Anthropic, None)] + fn on_resolves_known_spellings_held_as_other( + #[case] raw: &str, + #[case] provider: BetaProvider, + #[case] expected: Option, + ) { + let actual = AnthropicBeta::Other(raw.to_string()).on(provider); + assert_eq!(actual, expected); + assert!(!matches!(actual, Some(AnthropicBeta::Other(_)))); + } + + #[rstest] + #[case::anthropic(BetaProvider::Anthropic)] + #[case::azure_ai(BetaProvider::AzureAi)] + #[case::bedrock_converse(BetaProvider::BedrockConverse)] + #[case::bedrock(BetaProvider::Bedrock)] + #[case::bedrock_mantle(BetaProvider::BedrockMantle)] + #[case::vertex_ai(BetaProvider::VertexAi)] + #[case::databricks(BetaProvider::Databricks)] + fn unknown_beta_is_rejected_by_every_provider(#[case] provider: BetaProvider) { + assert_eq!( + AnthropicBeta::Other("example-beta-2099-01-01".into()).on(provider), + None + ); + } + #[test] fn unknown_values_are_kept_verbatim() { let parsed: AnthropicBeta = "claude-code-20250219".parse().unwrap(); diff --git a/litellm-rust/crates/llms-types/src/providers/minimax.rs b/litellm-rust/crates/llms-types/src/providers/minimax.rs new file mode 100644 index 00000000000..7ddbf32af2c --- /dev/null +++ b/litellm-rust/crates/llms-types/src/providers/minimax.rs @@ -0,0 +1,59 @@ +use serde_json::{Map, Value}; + +use crate::formats::messages::CacheControl; + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum MinimaxMessagesContentBlock { + Image(MinimaxMediaBlock), + Video(MinimaxMediaBlock), + MidConvSystem { + text: String, + #[serde(flatten)] + extra: Map, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +pub struct MinimaxMediaBlock { + pub source: MinimaxMediaSource, + pub cache_control: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum MinimaxMediaSource { + Base64 { + media_type: String, + data: String, + #[serde(flatten)] + options: MinimaxMediaOptions, + }, + Url { + url: String, + #[serde(flatten)] + options: MinimaxMediaOptions, + }, +} + +#[serde_with::skip_serializing_none] +#[macro_rules_attribute::apply(crate::wire_type)] +#[derive(Default)] +pub struct MinimaxMediaOptions { + pub detail: Option, + pub fps: Option, + pub max_long_side_pixel: Option, + #[serde(flatten)] + pub extra: Map, +} + +#[macro_rules_attribute::apply(crate::wire_type)] +#[serde(rename_all = "snake_case")] +pub enum MinimaxMediaDetail { + Low, + Default, + High, +} diff --git a/litellm-rust/crates/llms-types/src/providers/mod.rs b/litellm-rust/crates/llms-types/src/providers/mod.rs index e529997219e..c9eafe964ef 100644 --- a/litellm-rust/crates/llms-types/src/providers/mod.rs +++ b/litellm-rust/crates/llms-types/src/providers/mod.rs @@ -1 +1,2 @@ pub mod anthropic; +pub mod minimax; diff --git a/litellm-rust/crates/llms-types/src/recognized.rs b/litellm-rust/crates/llms-types/src/recognized.rs index 148d65381a5..4fe5f66b31b 100644 --- a/litellm-rust/crates/llms-types/src/recognized.rs +++ b/litellm-rust/crates/llms-types/src/recognized.rs @@ -1,6 +1,6 @@ use serde_json::Value; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[serde(untagged)] pub enum Recognized { Known(T), diff --git a/litellm-rust/crates/llms-types/tests/chat.rs b/litellm-rust/crates/llms-types/tests/chat.rs new file mode 100644 index 00000000000..76aa791b510 --- /dev/null +++ b/litellm-rust/crates/llms-types/tests/chat.rs @@ -0,0 +1,163 @@ +use litellm_llms_types::formats::chat_completions::{ + ChatContentPart, ChatLogprobs, ChatMediaUrl, PromptCacheMode, +}; +use litellm_llms_types::formats::messages::ContentSource; +use rstest::rstest; +use serde_json::{Value, json}; + +#[rstest] +#[case::parts(json!([ + {"type":"text","text":"hello","cache_control":{"type":"ephemeral"},"prompt_cache_breakpoint":{"mode":"explicit"}}, + {"type":"image_url","image_url":{"url":"https://example.test/image","detail":"high","format":"image/png"}}, + {"type":"video_url","video_url":"https://example.test/video"}, + {"type":"input_audio","input_audio":{"data":"AA==","format":"wav"}}, + {"type":"file","file":{"file_id":"file_1","file_data":"JVBE","filename":"a.mp4","format":"video/mp4","detail":"low","video_metadata":{"fps":1,"start_offset":"1s","end_offset":"2s"}}}, + {"type":"document","source":{"type":"text","media_type":"text/plain","data":"doc"},"title":"T","context":"C","citations":{"enabled":true}}, + {"type":"refusal","refusal":"refused"} +]))] +fn content_parts_round_trip(#[case] wire: Value) { + let parts: Vec = serde_json::from_value(wire.clone()).unwrap(); + let [ + ChatContentPart::Text { + text, + cache_control, + prompt_cache_breakpoint: Some(breakpoint), + extra: text_extra, + }, + ChatContentPart::ImageUrl { + image_url: ChatMediaUrl::Parameters(image), + prompt_cache_breakpoint: None, + .. + }, + ChatContentPart::VideoUrl { + video_url: ChatMediaUrl::Url(video), + .. + }, + ChatContentPart::InputAudio { input_audio, .. }, + ChatContentPart::File { file, .. }, + ChatContentPart::Document { + source, + title: Some(title), + context: Some(context), + citations: Some(citations), + .. + }, + ChatContentPart::Refusal { refusal, .. }, + ] = parts.as_slice() + else { + panic!("expected typed content parts"); + }; + assert_eq!(text, "hello"); + assert_eq!( + cache_control.as_ref().unwrap().cache_type.as_deref(), + Some("ephemeral") + ); + assert_eq!(breakpoint.mode, PromptCacheMode::Explicit); + assert!(breakpoint.extra.is_empty()); + assert!(text_extra.is_empty()); + assert_eq!(image.url, "https://example.test/image"); + assert_eq!(image.detail.as_deref(), Some("high")); + assert_eq!(image.format.as_deref(), Some("image/png")); + assert!(image.extra.is_empty()); + assert_eq!(video, "https://example.test/video"); + assert_eq!(input_audio.data, "AA=="); + assert_eq!(input_audio.format, "wav"); + assert_eq!(file.file_id.as_deref(), Some("file_1")); + assert_eq!(file.file_data.as_deref(), Some("JVBE")); + assert_eq!(file.filename.as_deref(), Some("a.mp4")); + assert_eq!(file.format.as_deref(), Some("video/mp4")); + assert_eq!(file.detail.as_deref(), Some("low")); + assert!(file.extra.is_empty()); + let metadata = file.video_metadata.as_ref().unwrap(); + assert_eq!(metadata.fps, Some(1.into())); + assert_eq!(metadata.start_offset.as_deref(), Some("1s")); + assert_eq!(metadata.end_offset.as_deref(), Some("2s")); + assert!(metadata.extra.is_empty()); + let ContentSource::Text { + data, media_type, .. + } = source.as_ref() + else { + panic!("expected document text source"); + }; + assert_eq!(data, "doc"); + assert_eq!(media_type, "text/plain"); + assert_eq!(title, "T"); + assert_eq!(context, "C"); + assert_eq!(citations.enabled, Some(true)); + assert_eq!(refusal, "refused"); + assert_eq!(serde_json::to_value(parts).unwrap(), wire); +} + +#[rstest] +fn logprobs_round_trip_with_tokens_and_alternatives() { + let wire = json!({ + "content":[{"token":"hi","logprob":-1,"bytes":[104,105],"top_logprobs":[{"token":"hey","logprob":-2.5,"bytes":[104]}]}], + "refusal":[{"token":"refused","logprob":-3}] + }); + let logprobs: ChatLogprobs = serde_json::from_value(wire.clone()).unwrap(); + let token = &logprobs.content.as_ref().unwrap()[0]; + assert_eq!(token.token, "hi"); + assert_eq!(token.logprob, (-1).into()); + assert_eq!(token.bytes.as_deref(), Some([104, 105].as_slice())); + let alternative = &token.top_logprobs.as_ref().unwrap()[0]; + assert_eq!(alternative.token, "hey"); + assert_eq!(alternative.bytes.as_deref(), Some([104].as_slice())); + assert_eq!(logprobs.refusal.as_ref().unwrap()[0].token, "refused"); + assert!(logprobs.refusal.as_ref().unwrap()[0].top_logprobs.is_none()); + assert_eq!(serde_json::to_value(logprobs).unwrap(), wire); +} + +#[rstest] +#[case::text(json!({"type":"text","text":false}))] +#[case::missing_image(json!({"type":"image_url"}))] +#[case::audio_shape(json!({"type":"input_audio","input_audio":{"data":7,"format":"wav"}}))] +#[case::audio_missing_format(json!({"type":"input_audio","input_audio":{"data":"AA=="}}))] +#[case::image_missing_url(json!({"type":"image_url","image_url":{"detail":"high"}}))] +#[case::file_metadata(json!({"type":"file","file":{"video_metadata":{"fps":"fast"}}}))] +#[case::document_source(json!({"type":"document","source":{"type":"url","url":false}}))] +#[case::refusal_missing_text(json!({"type":"refusal"}))] +#[case::file_missing_file(json!({"type":"file"}))] +#[case::document_missing_source(json!({"type":"document"}))] +#[case::video_missing_url(json!({"type":"video_url"}))] +#[case::cache_control_shape(json!({"type":"text","text":"t","cache_control":"ephemeral"}))] +#[case::breakpoint_missing_mode(json!({"type":"text","text":"t","prompt_cache_breakpoint":{}}))] +#[case::breakpoint_unknown_mode(json!({"type":"text","text":"t","prompt_cache_breakpoint":{"mode":"auto"}}))] +#[case::unknown_tag(json!({"type":"future"}))] +fn content_parts_reject_malformed_typed_fields(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} + +#[rstest] +fn partial_file_preserves_extensions_and_omits_null_optionals() { + let wire = json!({"type":"file","file":{"file_id":null,"filename":"a.pdf","extension":[1,null]},"future":true}); + let part: ChatContentPart = serde_json::from_value(wire).unwrap(); + let ChatContentPart::File { + file, + prompt_cache_breakpoint: None, + extra, + } = &part + else { + panic!("expected file") + }; + assert!(file.file_id.is_none()); + assert!(file.file_data.is_none()); + assert_eq!(file.filename.as_deref(), Some("a.pdf")); + assert_eq!(extra["future"], json!(true)); + assert_eq!( + serde_json::to_value(part).unwrap(), + json!({"type":"file","file":{"filename":"a.pdf","extension":[1,null]},"future":true}) + ); +} + +#[rstest] +#[case::missing_token(json!({"logprob":-1}))] +#[case::missing_logprob(json!({"token":"hi"}))] +#[case::wrong_bytes(json!({"token":"hi","logprob":-1,"bytes":[256]}))] +fn token_logprobs_reject_malformed_fields(#[case] wire: Value) { + assert!( + serde_json::from_value::( + wire + ) + .is_err() + ); +} diff --git a/litellm-rust/crates/llms-types/tests/json_schema.rs b/litellm-rust/crates/llms-types/tests/json_schema.rs new file mode 100644 index 00000000000..a626304910c --- /dev/null +++ b/litellm-rust/crates/llms-types/tests/json_schema.rs @@ -0,0 +1,148 @@ +use litellm_llms_types::json_schema::{ + JsonSchema, JsonSchemaItems, JsonSchemaObject, JsonSchemaType, +}; +use rstest::rstest; +use serde::{Serialize, de::DeserializeOwned}; +use serde_json::{Value, json}; + +fn round_trip(wire: Value) -> T +where + T: DeserializeOwned + Serialize, +{ + let parsed: T = serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(serde_json::to_value(&parsed).unwrap(), wire); + parsed +} + +#[rstest] +#[case::boolean(json!(false))] +#[case::nested(json!({ + "type":["object","null"], + "properties":{"nested":{"type":"array","items":{"$ref":"#/$defs/item"}},"free":true}, + "additionalProperties":{"type":"string"}, + "$defs":{"item":{"anyOf":[{"type":"string"},false]}}, + "required":["nested"], + "enum":[{"custom":[1,null]}], + "future":null +}))] +fn recursive_schema_round_trips(#[case] wire: Value) { + let schema = round_trip::(wire); + match schema { + JsonSchema::Boolean(allowed) => assert!(!allowed), + JsonSchema::Object(schema) => { + assert_eq!( + schema.schema_type, + Some(JsonSchemaType::Names(vec!["object".into(), "null".into()])) + ); + let properties = schema.properties.as_ref().unwrap(); + let JsonSchema::Object(nested) = &properties["nested"] else { + panic!("expected nested schema") + }; + let Some(JsonSchemaItems::Schema(items)) = &nested.items else { + panic!("expected single items schema") + }; + let JsonSchema::Object(items) = items.as_ref() else { + panic!("expected items schema") + }; + assert_eq!(items.reference.as_deref(), Some("#/$defs/item")); + assert_eq!(properties["free"], JsonSchema::Boolean(true)); + let JsonSchema::Object(definition) = &schema.defs.as_ref().unwrap()["item"] else { + panic!("expected schema definition") + }; + assert_eq!( + definition.any_of.as_ref().unwrap()[1], + JsonSchema::Boolean(false) + ); + assert_eq!(schema.extra["enum"], json!([{"custom":[1,null]}])); + } + } +} + +#[rstest] +fn schema_maps_keep_wire_key_order() { + let wire = r#"{"type":"object","properties":{"zeta":{"type":"string"},"alpha":{"type":"integer"}},"$defs":{"y":true,"b":false}}"#; + let schema: JsonSchemaObject = serde_json::from_str(wire).unwrap(); + let names: Vec<&str> = schema + .properties + .as_ref() + .unwrap() + .keys() + .map(String::as_str) + .collect(); + assert_eq!(names, ["zeta", "alpha"]); + assert_eq!(serde_json::to_string(&schema).unwrap(), wire); +} + +#[rstest] +fn schema_object_round_trips_supported_keywords() { + let schema = round_trip::(json!({ + "type":"object", + "properties":{"name":{"type":"string"}}, + "required":["name"], + "additionalProperties":false, + "$ref":"#/$defs/value", + "strict":true, + "extension":{"nested":[1,null]} + })); + assert_eq!( + schema.schema_type, + Some(JsonSchemaType::Name("object".into())) + ); + assert_eq!(schema.required.as_deref(), Some(["name".into()].as_slice())); + assert_eq!( + schema.additional_properties.as_deref(), + Some(&JsonSchema::Boolean(false)) + ); + assert_eq!(schema.reference.as_deref(), Some("#/$defs/value")); + assert_eq!(schema.strict, Some(true)); +} + +#[rstest] +#[case::scalar(json!(7))] +#[case::properties_shape(json!({"properties":[]}))] +#[case::nested_schema(json!({"properties":{"name":7}}))] +#[case::schema_type(json!({"type":["string",7]}))] +#[case::items_shape(json!({"items":7}))] +#[case::tuple_member(json!({"items":[{"type":"string"},7]}))] +#[case::prefix_items_shape(json!({"prefixItems":{"type":"string"}}))] +#[case::composite_shape(json!({"anyOf":["string"]}))] +fn schemas_reject_malformed_known_keywords(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} + +#[rstest] +#[case::draft_07_tuple(json!({"type":"array","items":[{"type":"string"},true]}))] +#[case::draft_2020_12_tuple(json!({"type":"array","prefixItems":[{"type":"string"},true],"items":false}))] +fn tuple_schemas_expose_positional_members(#[case] wire: Value) { + let schema = round_trip::(wire); + let members = match (&schema.items, &schema.prefix_items) { + (Some(JsonSchemaItems::Tuple(members)), None) => members, + (Some(JsonSchemaItems::Schema(rest)), Some(members)) => { + assert_eq!(rest.as_ref(), &JsonSchema::Boolean(false)); + members + } + _ => panic!("expected positional members"), + }; + let [JsonSchema::Object(first), JsonSchema::Boolean(true)] = members.as_slice() else { + panic!("expected string schema then true"); + }; + assert_eq!( + first.schema_type, + Some(JsonSchemaType::Name("string".into())) + ); + assert!(schema.extra.is_empty()); +} + +#[rstest] +fn empty_schema_omits_null_optionals_and_preserves_arbitrary_keywords() { + let schema: JsonSchemaObject = serde_json::from_value( + json!({"type":null,"properties":null,"const":{"arbitrary":[1,null]},"future":null}), + ) + .unwrap(); + assert!(schema.schema_type.is_none()); + assert!(schema.properties.is_none()); + assert_eq!( + serde_json::to_value(schema).unwrap(), + json!({"const":{"arbitrary":[1,null]},"future":null}) + ); +} diff --git a/litellm-rust/crates/llms-types/tests/messages_request.rs b/litellm-rust/crates/llms-types/tests/messages_request.rs index a0515c21459..4825932bc02 100644 --- a/litellm-rust/crates/llms-types/tests/messages_request.rs +++ b/litellm-rust/crates/llms-types/tests/messages_request.rs @@ -33,6 +33,15 @@ fn content_block_type_schema_remains_a_string() { #[case::compaction("compaction", ContentBlockType::Compaction)] #[case::advisor_result("advisor_tool_result", ContentBlockType::AdvisorToolResult)] #[case::web_search_result("web_search_tool_result", ContentBlockType::WebSearchToolResult)] +#[case::image("image", ContentBlockType::Other("image".into()))] +#[case::document("document", ContentBlockType::Other("document".into()))] +#[case::tool_addition("tool_addition", ContentBlockType::Other("tool_addition".into()))] +#[case::tool_removal("tool_removal", ContentBlockType::Other("tool_removal".into()))] +#[case::advisor("advisor_result", ContentBlockType::Other("advisor_result".into()))] +#[case::web_search_error( + "web_search_tool_result_error", + ContentBlockType::Other("web_search_tool_result_error".into()) +)] #[case::future_block("future_block", ContentBlockType::Other("future_block".into()))] #[case::case_sensitive("Tool_Use", ContentBlockType::Other("Tool_Use".into()))] #[case::empty("", ContentBlockType::Other(String::new()))] diff --git a/litellm-rust/crates/llms-types/tests/messages_types.rs b/litellm-rust/crates/llms-types/tests/messages_types.rs new file mode 100644 index 00000000000..51785a54c98 --- /dev/null +++ b/litellm-rust/crates/llms-types/tests/messages_types.rs @@ -0,0 +1,1512 @@ +use indexmap::IndexMap; +use litellm_llms_types::formats::messages::{ + AdvisorTool, AdvisorToolName, AdvisorToolResultContent, AllowedCaller, BashCodeExecutionOutput, + BashCodeExecutionToolResultContent, BashToolName, BlockContent, BrowserStateChange, + BrowserToolsetConfigs, BuiltinMessagesTool, CacheControl, CacheMissReason, Citation, + CitationsConfig, ClientTool, CodeExecutionOutput, CodeExecutionToolName, + CodeExecutionToolResultContent, ComputerTool, ComputerTool20251124, ComputerToolName, + ComputerToolsetConfigs, ContainerReference, ContentSource, ContextManagementResponse, + ContextTrigger, CustomTool, CustomToolType, FallbackTrigger, McpListedTool, McpServer, + McpToolResultContent, McpToolResultText, McpToolset, MemoryToolName, MessageRole, MessageType, + MessagesCompaction, MessagesContainer, MessagesContentPart, MessagesDiagnostics, + MessagesDiagnosticsParam, MessagesMetadata, MessagesToolParam, MessagesUsage, OutputFormat, + ResponseInclusion, Safeguard, ServerTool, SkillType, StopDetails, StopDetailsType, StopReason, + StrReplaceBasedEditToolName, StrReplaceEditorName, TextEditorCodeExecutionToolResultContent, + TextEditorFileType, TextEditorTool20250728, ToolCaller, ToolChange, ToolChangeTarget, + ToolChoice, ToolChoiceType, ToolResultUrlSource, ToolSearchBm25ToolName, ToolSearchReference, + ToolSearchRegexToolName, ToolSearchToolResultContent, Toolset, ToolsetToolConfig, + UrlSourceToolReference, UsageIterationType, UserInputUrlSource, UserLocationType, + WebFetchDocument, WebFetchTool, WebFetchTool20260309, WebFetchTool20260318, WebFetchToolName, + WebFetchToolResultContent, WebFetchUrlSources, WebSearchTool, WebSearchTool20260318, + WebSearchToolName, WebSearchToolResultContent, WebSearchUserLocation, +}; +use litellm_llms_types::json_schema::{JsonSchema, JsonSchemaObject, JsonSchemaType}; +use litellm_llms_types::recognized::Recognized; +use rstest::rstest; +use serde::{Serialize, de::DeserializeOwned}; +use serde_json::{Map, Value, json}; + +fn round_trip(wire: Value) -> T +where + T: DeserializeOwned + Serialize, +{ + let parsed: T = serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(serde_json::to_value(&parsed).unwrap(), wire); + parsed +} + +#[rstest] +fn metadata_contracts_round_trip() { + let metadata = round_trip::(json!({"user_id":"user_1","future":true})); + assert_eq!(metadata.user_id.as_deref(), Some("user_1")); + assert_eq!(metadata.extra.get("future"), Some(&json!(true))); + let output_format = round_trip::(json!({ + "type":"json_schema", + "schema":{"type":"object","properties":{"name":{"type":"string"}}}, + "strict":true + })); + assert_eq!(output_format.strict, Some(true)); + assert!(output_format.extra.is_empty()); + let JsonSchema::Object(schema) = &output_format.schema else { + panic!("expected output object schema"); + }; + assert!(schema.properties.as_ref().unwrap().contains_key("name")); + let compaction = + round_trip::(json!({"type":"summarize","instructions":"briefly"})); + assert_eq!(compaction.instructions.as_deref(), Some("briefly")); + assert!(compaction.extra.is_empty()); + let container = round_trip::(json!({ + "id":"container_1", + "expires_at":"2026-01-01T00:00:00Z", + "skills":[{"type":"custom","skill_id":"skill_1","version":"1"}] + })); + assert_eq!(container.id.as_deref(), Some("container_1")); + assert!(container.extra.is_empty()); + let [skill] = container.skills.as_ref().unwrap().as_slice() else { + panic!("expected container skill"); + }; + assert_eq!(skill.skill_type, SkillType::Custom); + assert_eq!(skill.skill_id, "skill_1"); + assert_eq!(skill.version.as_deref(), Some("1")); + assert!(skill.extra.is_empty()); + let reference = round_trip::(json!({"id":"container_1"})); + let ContainerReference::Parameters(parameters) = reference else { + panic!("expected container parameters"); + }; + assert_eq!(parameters.id.as_deref(), container.id.as_deref()); + assert!(parameters.skills.is_none()); + round_trip::( + json!({"id":"container_1","skills":[{"type":"anthropic","skill_id":"pptx"}]}), + ); + let server = round_trip::(json!({ + "type":"url", + "url":"https://example.test/mcp", + "name":"search", + "authorization_token":"token", + "tool_configuration":{"allowed_tools":["search"],"enabled":true} + })); + assert_eq!(server.url, "https://example.test/mcp"); + assert_eq!(server.name, "search"); + assert_eq!(server.authorization_token.as_deref(), Some("token")); + assert!(server.extra.is_empty()); + let configuration = server.tool_configuration.as_ref().unwrap(); + assert_eq!(configuration.enabled, Some(true)); + assert_eq!( + configuration.allowed_tools.as_deref(), + Some([String::from("search")].as_slice()) + ); + assert!(configuration.extra.is_empty()); + let context = round_trip::(json!({ + "applied_edits":[ + {"type":"clear_tool_uses_20250919","cleared_input_tokens":7,"cleared_tool_uses":2}, + {"type":"clear_thinking_20251015","cleared_input_tokens":3,"cleared_thinking_turns":1} + ] + })); + let [tool_uses, thinking] = context.applied_edits.as_ref().unwrap().as_slice() else { + panic!("expected two applied edits"); + }; + assert_eq!( + tool_uses.edit_type.as_deref(), + Some("clear_tool_uses_20250919") + ); + assert_eq!( + (tool_uses.cleared_input_tokens, tool_uses.cleared_tool_uses), + (Some(7), Some(2)) + ); + assert!(tool_uses.cleared_thinking_turns.is_none()); + assert_eq!( + ( + thinking.cleared_input_tokens, + thinking.cleared_thinking_turns + ), + (Some(3), Some(1)) + ); + assert!(tool_uses.extra.is_empty() && thinking.extra.is_empty()); + let safeguard = round_trip::( + json!({"type":"classifier","classifier_context":{"source":"test"}}), + ); + assert_eq!(safeguard.safeguard_type, "classifier"); + assert_eq!( + safeguard.classifier_context.as_ref().unwrap().get("source"), + Some(&json!("test")) + ); + assert!(safeguard.extra.is_empty()); +} + +#[rstest] +fn stop_details_expose_refusal_and_keep_unknown_types() { + let refusal = round_trip::( + json!({"type":"refusal","category":"cyber","explanation":"blocked"}), + ); + assert_eq!( + refusal.detail_type, + Recognized::Known(StopDetailsType::Refusal) + ); + assert_eq!(refusal.category.as_deref(), Some("cyber")); + assert_eq!(refusal.explanation.as_deref(), Some("blocked")); + assert!(refusal.extra.is_empty()); + let safeguard = + round_trip::(json!({"type":"safeguard","safeguard_types":["classifier"]})); + assert_eq!( + safeguard.detail_type, + Recognized::Unrecognized(json!("safeguard")) + ); + assert_eq!(safeguard.extra["safeguard_types"], json!(["classifier"])); +} + +#[rstest] +fn container_rejects_skill_without_id() { + assert!( + serde_json::from_value::( + json!({"skills":[{"type":"custom","version":"1"}]}) + ) + .is_err() + ); +} + +#[rstest] +#[case::without_url(json!({"type":"url","name":"search"}))] +#[case::without_name(json!({"type":"url","url":"https://example.test/mcp"}))] +fn mcp_server_rejects_missing_required_fields(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} + +#[rstest] +fn output_format_rejects_missing_schema() { + assert!(serde_json::from_value::(json!({"type":"json_schema"})).is_err()); +} + +#[rstest] +#[case::user(json!("user"))] +#[case::assistant(json!("assistant"))] +#[case::system(json!("system"))] +#[case::future(json!("future_role"))] +fn message_roles_round_trip(#[case] wire: Value) { + round_trip::(wire); +} + +#[rstest] +#[case::message(json!("message"))] +#[case::future(json!("future_type"))] +fn message_types_round_trip(#[case] wire: Value) { + round_trip::(wire); +} + +#[rstest] +#[case::end_turn(json!("end_turn"))] +#[case::refusal(json!("refusal"))] +#[case::compaction(json!("compaction"))] +#[case::future(json!("future_reason"))] +fn stop_reasons_round_trip(#[case] wire: Value) { + round_trip::(wire); +} + +#[rstest] +fn tool_choice_round_trips() { + let choice = round_trip::( + json!({"type":"tool","name":"search","disable_parallel_tool_use":true}), + ); + assert_eq!(choice.choice_type, ToolChoiceType::Tool); + assert_eq!(choice.name.as_deref(), Some("search")); + assert_eq!(choice.disable_parallel_tool_use, Some(true)); + assert!(choice.extra.is_empty()); +} + +#[rstest] +fn usage_contracts_round_trip() { + let usage = round_trip::(json!({ + "input_tokens":10, + "output_tokens":4, + "server_tool_use":{"web_search_requests":2,"web_fetch_requests":1}, + "cache_creation":{"ephemeral_1h_input_tokens":3,"ephemeral_5m_input_tokens":1}, + "output_tokens_details":{"thinking_tokens":2}, + "iterations":[ + {"type":"compaction","input_tokens":7,"output_tokens":1}, + {"type":"message","input_tokens":3,"output_tokens":3} + ], + "service_tier":"priority", + "inference_geo":"global", + "speed":"fast" + })); + assert_eq!(usage.inference_geo.as_deref(), Some("global")); + assert_eq!(usage.input_tokens, Some(10)); + assert_eq!(usage.output_tokens, Some(4)); + assert!(usage.extra.is_empty()); + let server = usage.server_tool_use.as_ref().unwrap(); + assert_eq!(server.web_search_requests, Some(2)); + assert_eq!(server.web_fetch_requests, Some(1)); + let cache = usage.cache_creation.as_ref().unwrap(); + assert_eq!(cache.ephemeral_1h_input_tokens, Some(3)); + assert_eq!(cache.ephemeral_5m_input_tokens, Some(1)); + assert_eq!( + usage + .output_tokens_details + .as_ref() + .unwrap() + .thinking_tokens, + Some(2) + ); + let [compaction, message] = usage.iterations.as_ref().unwrap().as_slice() else { + panic!("expected usage iterations"); + }; + assert_eq!( + compaction.iteration_type, + Recognized::Known(UsageIterationType::Compaction) + ); + assert_eq!( + message.iteration_type, + Recognized::Known(UsageIterationType::Message) + ); + assert_eq!( + compaction.input_tokens.unwrap() + message.input_tokens.unwrap(), + usage.input_tokens.unwrap() + ); + assert_eq!( + compaction.output_tokens.unwrap() + message.output_tokens.unwrap(), + usage.output_tokens.unwrap() + ); +} + +#[rstest] +fn usage_iterations_expose_advisor_and_fallback_models() { + let usage = round_trip::(json!({ + "iterations":[ + {"type":"advisor_message","model":"claude-opus-5-5","input_tokens":5,"output_tokens":2, + "cache_creation_input_tokens":4,"cache_read_input_tokens":1, + "cache_creation":{"ephemeral_1h_input_tokens":0,"ephemeral_5m_input_tokens":4}}, + {"type":"fallback_message","model":"claude-sonnet-5-5","input_tokens":6,"output_tokens":3, + "cache_creation_input_tokens":0,"cache_read_input_tokens":2}, + {"type":"future_iteration","input_tokens":1} + ] + })); + let [advisor, fallback, future] = usage.iterations.as_deref().unwrap() else { + panic!("expected three iterations"); + }; + assert_eq!( + advisor.iteration_type, + Recognized::Known(UsageIterationType::AdvisorMessage) + ); + assert_eq!(advisor.model.as_deref(), Some("claude-opus-5-5")); + assert_eq!( + ( + advisor.cache_creation_input_tokens, + advisor.cache_read_input_tokens + ), + (Some(4), Some(1)) + ); + assert_eq!( + advisor + .cache_creation + .as_ref() + .unwrap() + .ephemeral_5m_input_tokens, + advisor.cache_creation_input_tokens + ); + assert!(advisor.extra.is_empty()); + assert_eq!( + fallback.iteration_type, + Recognized::Known(UsageIterationType::FallbackMessage) + ); + assert_eq!(fallback.model.as_deref(), Some("claude-sonnet-5-5")); + assert_eq!(fallback.cache_read_input_tokens, Some(2)); + assert!(fallback.extra.is_empty()); + assert_eq!( + future.iteration_type, + Recognized::Unrecognized(json!("future_iteration")) + ); + assert_eq!(future.input_tokens, Some(1)); +} + +#[rstest] +#[case::base64(json!({"type":"base64","media_type":"image/png","data":"AA=="}))] +#[case::url(json!({"type":"url","url":"https://example.test/image"}))] +#[case::file(json!({"type":"file","file_id":"file_1"}))] +#[case::text(json!({"type":"text","media_type":"text/plain","data":"document"}))] +#[case::content(json!({"type":"content","content":"nested"}))] +fn content_sources_round_trip(#[case] wire: Value) { + round_trip::(wire); +} + +fn client_tool(name: N) -> ClientTool { + ClientTool { + name, + allowed_callers: None, + cache_control: None, + defer_loading: None, + input_examples: None, + strict: None, + extra: Map::new(), + } +} + +fn server_tool(name: N) -> ServerTool { + ServerTool { + name, + allowed_callers: None, + cache_control: None, + defer_loading: None, + strict: None, + extra: Map::new(), + } +} + +fn computer_tool(display_width_px: u64, display_height_px: u64) -> ComputerTool { + ComputerTool { + name: ComputerToolName::Computer, + display_width_px, + display_height_px, + display_number: None, + allowed_callers: None, + cache_control: None, + defer_loading: None, + input_examples: None, + strict: None, + extra: Map::new(), + } +} + +fn web_search_tool() -> WebSearchTool { + WebSearchTool { + name: WebSearchToolName::WebSearch, + allowed_callers: None, + allowed_domains: None, + blocked_domains: None, + cache_control: None, + defer_loading: None, + max_uses: None, + strict: None, + user_location: None, + extra: Map::new(), + } +} + +fn web_fetch_tool() -> WebFetchTool { + WebFetchTool { + name: WebFetchToolName::WebFetch, + allowed_callers: None, + allowed_domains: None, + blocked_domains: None, + cache_control: None, + citations: None, + defer_loading: None, + max_content_tokens: None, + max_uses: None, + strict: None, + url_sources: None, + extra: Map::new(), + } +} + +fn web_fetch_url_sources() -> WebFetchUrlSources { + WebFetchUrlSources { + client_tool_results: None, + server_tool_results: None, + user_input: None, + extra: Map::new(), + } +} + +fn tool_reference(name: &str) -> UrlSourceToolReference { + UrlSourceToolReference::ToolReference { + name: String::from(name), + extra: Map::new(), + } +} + +fn toolset_config(enabled: Option, defer_loading: Option) -> ToolsetToolConfig { + ToolsetToolConfig { + defer_loading, + enabled, + extra: Map::new(), + } +} + +fn cache_control(cache_type: &str, ttl: &str) -> CacheControl { + CacheControl { + cache_type: Some(String::from(cache_type)), + ttl: Some(String::from(ttl)), + ..CacheControl::default() + } +} + +fn browser_toolset_configs() -> BrowserToolsetConfigs { + BrowserToolsetConfigs { + type_text: None, + close_tab: None, + double_click: None, + file_upload: None, + find: None, + form_input: None, + get_page_text: None, + hold_key: None, + hover: None, + javascript_exec: None, + key: None, + left_click: None, + left_click_drag: None, + left_mouse_down: None, + left_mouse_up: None, + list_tabs: None, + middle_click: None, + mouse_move: None, + navigate: None, + new_tab: None, + read_console: None, + read_network: None, + read_page: None, + right_click: None, + screenshot: None, + scroll: None, + scroll_to: None, + switch_tab: None, + triple_click: None, + wait: None, + zoom: None, + extra: Map::new(), + } +} + +fn computer_toolset_configs() -> ComputerToolsetConfigs { + ComputerToolsetConfigs { + type_text: None, + cursor_position: None, + double_click: None, + hold_key: None, + key: None, + left_click: None, + left_click_drag: None, + left_mouse_down: None, + left_mouse_up: None, + middle_click: None, + mouse_move: None, + right_click: None, + screenshot: None, + scroll: None, + triple_click: None, + wait: None, + zoom: None, + extra: Map::new(), + } +} + +#[rstest] +#[case::bash_20241022( + json!({"type":"bash_20241022","name":"bash","input_examples":[{"command":"ls"}]}), + BuiltinMessagesTool::Bash20241022(ClientTool { + input_examples: Some(vec![Map::from_iter([(String::from("command"), json!("ls"))])]), + ..client_tool(BashToolName::Bash) + }) +)] +#[case::bash_20250124( + json!({"type":"bash_20250124","name":"bash","allowed_callers":["direct","code_execution_20260521"]}), + BuiltinMessagesTool::Bash20250124(ClientTool { + allowed_callers: Some(vec![AllowedCaller::Direct, AllowedCaller::CodeExecution20260521]), + ..client_tool(BashToolName::Bash) + }) +)] +#[case::text_editor_20241022( + json!({"type":"text_editor_20241022","name":"str_replace_editor","defer_loading":true}), + BuiltinMessagesTool::TextEditor20241022(ClientTool { + defer_loading: Some(true), + ..client_tool(StrReplaceEditorName::StrReplaceEditor) + }) +)] +#[case::text_editor_20250124( + json!({"type":"text_editor_20250124","name":"str_replace_editor","strict":true}), + BuiltinMessagesTool::TextEditor20250124(ClientTool { + strict: Some(true), + ..client_tool(StrReplaceEditorName::StrReplaceEditor) + }) +)] +#[case::text_editor_20250429( + json!({"type":"text_editor_20250429","name":"str_replace_based_edit_tool","cache_control":{"type":"ephemeral","ttl":"1h"}}), + BuiltinMessagesTool::TextEditor20250429(ClientTool { + cache_control: Some(cache_control("ephemeral", "1h")), + ..client_tool(StrReplaceBasedEditToolName::StrReplaceBasedEditTool) + }) +)] +#[case::text_editor_20250728( + json!({"type":"text_editor_20250728","name":"str_replace_based_edit_tool","max_characters":10000}), + BuiltinMessagesTool::TextEditor20250728(TextEditorTool20250728 { + name: StrReplaceBasedEditToolName::StrReplaceBasedEditTool, + allowed_callers: None, + cache_control: None, + defer_loading: None, + input_examples: None, + max_characters: Some(10000), + strict: None, + extra: Map::new(), + }) +)] +#[case::memory_20250818( + json!({"type":"memory_20250818","name":"memory","allowed_callers":["code_execution_20250825"]}), + BuiltinMessagesTool::Memory20250818(ClientTool { + allowed_callers: Some(vec![AllowedCaller::CodeExecution20250825]), + ..client_tool(MemoryToolName::Memory) + }) +)] +#[case::computer_20241022( + json!({"type":"computer_20241022","name":"computer","display_width_px":1024,"display_height_px":768,"display_number":1}), + BuiltinMessagesTool::Computer20241022(ComputerTool { + display_number: Some(1), + ..computer_tool(1024, 768) + }) +)] +#[case::computer_20250124( + json!({"type":"computer_20250124","name":"computer","display_width_px":1280,"display_height_px":800}), + BuiltinMessagesTool::Computer20250124(computer_tool(1280, 800)) +)] +#[case::computer_20251124( + json!({"type":"computer_20251124","name":"computer","display_width_px":1920,"display_height_px":1080,"enable_zoom":true}), + BuiltinMessagesTool::Computer20251124(ComputerTool20251124 { + name: ComputerToolName::Computer, + display_width_px: 1920, + display_height_px: 1080, + display_number: None, + enable_zoom: Some(true), + allowed_callers: None, + cache_control: None, + defer_loading: None, + input_examples: None, + strict: None, + extra: Map::new(), + }) +)] +#[case::code_execution_20250522( + json!({"type":"code_execution_20250522","name":"code_execution","strict":false}), + BuiltinMessagesTool::CodeExecution20250522(ServerTool { + strict: Some(false), + ..server_tool(CodeExecutionToolName::CodeExecution) + }) +)] +#[case::code_execution_20250825( + json!({"type":"code_execution_20250825","name":"code_execution","defer_loading":false}), + BuiltinMessagesTool::CodeExecution20250825(ServerTool { + defer_loading: Some(false), + ..server_tool(CodeExecutionToolName::CodeExecution) + }) +)] +#[case::code_execution_20260120( + json!({"type":"code_execution_20260120","name":"code_execution","allowed_callers":["direct"]}), + BuiltinMessagesTool::CodeExecution20260120(ServerTool { + allowed_callers: Some(vec![AllowedCaller::Direct]), + ..server_tool(CodeExecutionToolName::CodeExecution) + }) +)] +#[case::code_execution_20260521( + json!({"type":"code_execution_20260521","name":"code_execution","allowed_callers":["code_execution_20260120"]}), + BuiltinMessagesTool::CodeExecution20260521(ServerTool { + allowed_callers: Some(vec![AllowedCaller::CodeExecution20260120]), + ..server_tool(CodeExecutionToolName::CodeExecution) + }) +)] +#[case::tool_search_regex_20251119( + json!({"type":"tool_search_tool_regex_20251119","name":"tool_search_tool_regex","defer_loading":true}), + BuiltinMessagesTool::ToolSearchRegex20251119(ServerTool { + defer_loading: Some(true), + ..server_tool(ToolSearchRegexToolName::ToolSearchToolRegex) + }) +)] +#[case::tool_search_regex( + json!({"type":"tool_search_tool_regex","name":"tool_search_tool_regex","strict":true}), + BuiltinMessagesTool::ToolSearchRegex(ServerTool { + strict: Some(true), + ..server_tool(ToolSearchRegexToolName::ToolSearchToolRegex) + }) +)] +#[case::tool_search_bm25_20251119( + json!({"type":"tool_search_tool_bm25_20251119","name":"tool_search_tool_bm25","defer_loading":false}), + BuiltinMessagesTool::ToolSearchBm2520251119(ServerTool { + defer_loading: Some(false), + ..server_tool(ToolSearchBm25ToolName::ToolSearchToolBm25) + }) +)] +#[case::tool_search_bm25( + json!({"type":"tool_search_tool_bm25","name":"tool_search_tool_bm25","strict":false}), + BuiltinMessagesTool::ToolSearchBm25(ServerTool { + strict: Some(false), + ..server_tool(ToolSearchBm25ToolName::ToolSearchToolBm25) + }) +)] +#[case::web_search_20250305( + json!({"type":"web_search_20250305","name":"web_search","max_uses":3,"allowed_domains":["example.test"], + "user_location":{"type":"approximate","city":"San Francisco","country":"US","timezone":"America/Los_Angeles"}}), + BuiltinMessagesTool::WebSearch20250305(WebSearchTool { + max_uses: Some(3), + allowed_domains: Some(vec![String::from("example.test")]), + user_location: Some(WebSearchUserLocation { + location_type: UserLocationType::Approximate, + city: Some(String::from("San Francisco")), + country: Some(String::from("US")), + region: None, + timezone: Some(String::from("America/Los_Angeles")), + extra: Map::new(), + }), + ..web_search_tool() + }) +)] +#[case::web_search_20260209( + json!({"type":"web_search_20260209","name":"web_search","blocked_domains":["blocked.test"]}), + BuiltinMessagesTool::WebSearch20260209(WebSearchTool { + blocked_domains: Some(vec![String::from("blocked.test")]), + ..web_search_tool() + }) +)] +#[case::web_search_20260318( + json!({"type":"web_search_20260318","name":"web_search","response_inclusion":"excluded"}), + BuiltinMessagesTool::WebSearch20260318(WebSearchTool20260318 { + name: WebSearchToolName::WebSearch, + allowed_callers: None, + allowed_domains: None, + blocked_domains: None, + cache_control: None, + defer_loading: None, + max_uses: None, + response_inclusion: Some(ResponseInclusion::Excluded), + strict: None, + user_location: None, + extra: Map::new(), + }) +)] +#[case::web_fetch_20250910( + json!({"type":"web_fetch_20250910","name":"web_fetch","max_content_tokens":5000,"citations":{"enabled":true}, + "url_sources":{"client_tool_results":{"type":"only","tools":[{"type":"tool_reference","name":"lookup"}]}, + "server_tool_results":{"type":"all"},"user_input":{"type":"none"}}}), + BuiltinMessagesTool::WebFetch20250910(WebFetchTool { + max_content_tokens: Some(5000), + citations: Some(CitationsConfig { + enabled: Some(true), + extra: Map::new(), + }), + url_sources: Some(WebFetchUrlSources { + client_tool_results: Some(ToolResultUrlSource::Only { + tools: vec![tool_reference("lookup")], + extra: Map::new(), + }), + server_tool_results: Some(ToolResultUrlSource::All { extra: Map::new() }), + user_input: Some(UserInputUrlSource::None { extra: Map::new() }), + extra: Map::new(), + }), + ..web_fetch_tool() + }) +)] +#[case::web_fetch_20260209( + json!({"type":"web_fetch_20260209","name":"web_fetch","url_sources":{"server_tool_results":{"type":"except","tools":[{"type":"tool_reference","name":"web_search"}]}}}), + BuiltinMessagesTool::WebFetch20260209(WebFetchTool { + url_sources: Some(WebFetchUrlSources { + server_tool_results: Some(ToolResultUrlSource::Except { + tools: vec![tool_reference("web_search")], + extra: Map::new(), + }), + ..web_fetch_url_sources() + }), + ..web_fetch_tool() + }) +)] +#[case::web_fetch_20260309( + json!({"type":"web_fetch_20260309","name":"web_fetch","use_cache":false}), + BuiltinMessagesTool::WebFetch20260309(WebFetchTool20260309 { + name: WebFetchToolName::WebFetch, + allowed_callers: None, + allowed_domains: None, + blocked_domains: None, + cache_control: None, + citations: None, + defer_loading: None, + max_content_tokens: None, + max_uses: None, + strict: None, + url_sources: None, + use_cache: Some(false), + extra: Map::new(), + }) +)] +#[case::web_fetch_20260318( + json!({"type":"web_fetch_20260318","name":"web_fetch","use_cache":true,"response_inclusion":"full"}), + BuiltinMessagesTool::WebFetch20260318(WebFetchTool20260318 { + name: WebFetchToolName::WebFetch, + allowed_callers: None, + allowed_domains: None, + blocked_domains: None, + cache_control: None, + citations: None, + defer_loading: None, + max_content_tokens: None, + max_uses: None, + response_inclusion: Some(ResponseInclusion::Full), + strict: None, + url_sources: None, + use_cache: Some(true), + extra: Map::new(), + }) +)] +#[case::advisor_20260301( + json!({"type":"advisor_20260301","name":"advisor","model":"claude-opus-5-5","max_tokens":2048,"caching":{"type":"ephemeral","ttl":"5m"}}), + BuiltinMessagesTool::Advisor20260301(AdvisorTool { + name: AdvisorToolName::Advisor, + model: String::from("claude-opus-5-5"), + allowed_callers: None, + cache_control: None, + caching: Some(cache_control("ephemeral", "5m")), + defer_loading: None, + max_tokens: Some(2048), + max_uses: None, + strict: None, + extra: Map::new(), + }) +)] +#[case::browser_toolset_20260801( + json!({"type":"browser_toolset_20260801","configs":{"type":{"enabled":false},"javascript_exec":{"defer_loading":true}}}), + BuiltinMessagesTool::BrowserToolset20260801(Toolset { + cache_control: None, + configs: Some(Box::new(BrowserToolsetConfigs { + type_text: Some(toolset_config(Some(false), None)), + javascript_exec: Some(toolset_config(None, Some(true))), + ..browser_toolset_configs() + })), + extra: Map::new(), + }) +)] +#[case::computer_toolset_20260801( + json!({"type":"computer_toolset_20260801","configs":{"zoom":{"enabled":false},"cursor_position":{"enabled":true}}}), + BuiltinMessagesTool::ComputerToolset20260801(Toolset { + cache_control: None, + configs: Some(Box::new(ComputerToolsetConfigs { + zoom: Some(toolset_config(Some(false), None)), + cursor_position: Some(toolset_config(Some(true), None)), + ..computer_toolset_configs() + })), + extra: Map::new(), + }) +)] +#[case::mcp_toolset( + json!({"type":"mcp_toolset","mcp_server_name":"kb","default_config":{"enabled":false}, + "configs":{"search":{"enabled":true,"defer_loading":true}}, + "tools":[{"name":"search","input_schema":{"type":"object"},"description":"Search"}]}), + BuiltinMessagesTool::McpToolset(McpToolset { + mcp_server_name: String::from("kb"), + cache_control: None, + configs: Some(IndexMap::from([(String::from("search"), toolset_config(Some(true), Some(true)))])), + default_config: Some(toolset_config(Some(false), None)), + tools: Some(vec![McpListedTool { + name: String::from("search"), + description: Some(String::from("Search")), + input_schema: JsonSchema::Object(Box::new(JsonSchemaObject { + schema_type: Some(JsonSchemaType::Name(String::from("object"))), + ..JsonSchemaObject::default() + })), + extra: Map::new(), + }]), + extra: Map::new(), + }) +)] +fn builtin_tools_decode_typed_definitions( + #[case] wire: Value, + #[case] expected: BuiltinMessagesTool, +) { + assert_eq!(round_trip::(wire), expected); +} + +#[rstest] +fn builtin_tools_preserve_unknown_fields() { + let tool = round_trip::( + json!({"type":"bash_20250124","name":"bash","extension":[1,null]}), + ); + let BuiltinMessagesTool::Bash20250124(bash) = tool else { + panic!("expected bash tool"); + }; + assert_eq!(bash.extra.get("extension"), Some(&json!([1, null]))); +} + +#[rstest] +#[case::input_tokens("input_tokens", false)] +#[case::tool_uses("tool_uses", true)] +fn context_trigger_exposes_typed_threshold(#[case] tag: &str, #[case] tool_uses: bool) { + let wire = json!({"type":tag,"value":1024,"extension":true}); + let trigger: ContextTrigger = serde_json::from_value(wire.clone()).unwrap(); + match &trigger { + ContextTrigger::InputTokens { value, extra } => { + assert!(!tool_uses); + assert_eq!(*value, 1024); + assert_eq!(extra.get("extension"), Some(&json!(true))); + } + ContextTrigger::ToolUses { value, extra } => { + assert!(tool_uses); + assert_eq!(*value, 1024); + assert_eq!(extra.get("extension"), Some(&json!(true))); + } + } + assert_eq!(serde_json::to_value(trigger).unwrap(), wire); +} + +#[rstest] +#[case::negative(json!({"type":"input_tokens","value":-1}))] +#[case::wrong_shape(json!({"type":"input_tokens","value":"1024"}))] +#[case::missing_value(json!({"type":"input_tokens"}))] +#[case::null_value(json!({"type":"input_tokens","value":null}))] +#[case::tool_uses_negative(json!({"type":"tool_uses","value":-1}))] +#[case::tool_uses_fractional(json!({"type":"tool_uses","value":1.5}))] +#[case::tool_uses_missing_value(json!({"type":"tool_uses"}))] +#[case::missing_discriminator(json!({"value":1}))] +#[case::unknown_discriminator(json!({"type":"other","value":1}))] +fn token_threshold_requires_unsigned_integer(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} + +#[rstest] +#[case::missing_discriminator(json!({"name":"bash"}))] +#[case::unknown_discriminator(json!({"type":"future_tool","name":"bash"}))] +#[case::null_discriminator(json!({"type":null,"name":"bash"}))] +#[case::bash_without_name(json!({"type":"bash_20250124"}))] +#[case::bash_wrong_name(json!({"type":"bash_20241022","name":"shell"}))] +#[case::legacy_editor_with_new_name(json!({"type":"text_editor_20250124","name":"str_replace_based_edit_tool"}))] +#[case::new_editor_with_legacy_name(json!({"type":"text_editor_20250728","name":"str_replace_editor"}))] +#[case::editor_20250429_with_legacy_name(json!({"type":"text_editor_20250429","name":"str_replace_editor"}))] +#[case::memory_without_name(json!({"type":"memory_20250818"}))] +#[case::computer_without_width(json!({"type":"computer_20250124","name":"computer","display_height_px":768}))] +#[case::computer_without_height(json!({"type":"computer_20241022","name":"computer","display_width_px":1024}))] +#[case::zoom_computer_without_display(json!({"type":"computer_20251124","name":"computer"}))] +#[case::code_execution_without_name(json!({"type":"code_execution_20250825"}))] +#[case::regex_search_with_bm25_name(json!({"type":"tool_search_tool_regex","name":"tool_search_tool_bm25"}))] +#[case::bm25_search_with_regex_name(json!({"type":"tool_search_tool_bm25_20251119","name":"tool_search_tool_regex"}))] +#[case::web_search_with_fetch_name(json!({"type":"web_search_20250305","name":"web_fetch"}))] +#[case::web_fetch_without_name(json!({"type":"web_fetch_20260318"}))] +#[case::advisor_without_model(json!({"type":"advisor_20260301","name":"advisor"}))] +#[case::advisor_without_name(json!({"type":"advisor_20260301","model":"claude-opus-5-5"}))] +#[case::mcp_toolset_without_server(json!({"type":"mcp_toolset"}))] +#[case::unknown_allowed_caller(json!({"type":"code_execution_20250825","name":"code_execution","allowed_callers":["code_execution_20250522"]}))] +#[case::unknown_response_inclusion(json!({"type":"web_search_20260318","name":"web_search","response_inclusion":"partial"}))] +#[case::user_input_only_filter(json!({"type":"web_fetch_20250910","name":"web_fetch","url_sources":{"user_input":{"type":"only","tools":[]}}}))] +#[case::negative_limit(json!({"type":"web_search_20250305","name":"web_search","max_uses":-1}))] +#[case::wrong_mcp_config(json!({"type":"mcp_toolset","mcp_server_name":"kb","configs":{"search":{"enabled":"yes"}}}))] +#[case::wrong_browser_config(json!({"type":"browser_toolset_20260801","configs":{"navigate":{"enabled":1}}}))] +fn builtin_tools_reject_malformed_fields(#[case] wire: Value) { + assert!(serde_json::from_value::(wire.clone()).is_err()); + assert!(serde_json::from_value::(wire).is_err()); +} + +#[rstest] +fn existing_content_blocks_preserve_opaque_nested_fields() { + let wire = + json!({"type":"tool_result","content":[{"type":"text","text":7}],"source":{"url":7}}); + let block: litellm_llms_types::formats::messages::ContentBlock = + serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(block.content.as_ref(), wire.get("content")); + assert_eq!(block.extra.get("source"), wire.get("source")); + assert_eq!(serde_json::to_value(block).unwrap(), wire); +} + +fn part(wire: Value) -> MessagesContentPart { + round_trip::(wire) +} + +#[rstest] +fn text_blocks_expose_every_citation_location() { + let block = part(json!({ + "type":"text", + "text":"cited", + "citations":[ + {"type":"char_location","cited_text":"a","document_index":0,"start_char_index":1,"end_char_index":6,"file_id":"file_1"}, + {"type":"page_location","cited_text":"b","document_index":1,"start_page_number":1,"end_page_number":2}, + {"type":"content_block_location","cited_text":"c","document_index":2,"start_block_index":0,"end_block_index":1}, + {"type":"web_search_result_location","cited_text":"d","url":"https://example.test","encrypted_index":"opaque","title":"result"}, + {"type":"search_result_location","cited_text":"e","search_result_index":3,"source":"kb","start_block_index":0,"end_block_index":2} + ], + "cache_control":{"type":"ephemeral"} + })); + let MessagesContentPart::Text(text) = &block else { + panic!("expected text block"); + }; + assert_eq!(text.text, "cited"); + let [ + Citation::CharLocation(chars), + Citation::PageLocation(page), + Citation::ContentBlockLocation(blocks), + Citation::WebSearchResultLocation(search), + Citation::SearchResultLocation(result), + ] = text.citations.as_deref().unwrap() + else { + panic!("expected one citation of each location type"); + }; + assert_eq!((chars.start_char_index, chars.end_char_index), (1, 6)); + assert_eq!(chars.file_id.as_deref(), Some("file_1")); + assert_eq!((page.start_page_number, page.end_page_number), (1, 2)); + assert_eq!(blocks.document_index, 2); + assert_eq!(search.encrypted_index, "opaque"); + assert_eq!(result.search_result_index, 3); + assert_eq!(result.source, "kb"); +} + +#[rstest] +fn text_block_null_optionals_are_omitted() { + let block: MessagesContentPart = serde_json::from_value( + json!({"type":"text","text":"hi","citations":null,"cache_control":null,"future":null}), + ) + .unwrap(); + let MessagesContentPart::Text(text) = &block else { + panic!("expected text block"); + }; + assert!(text.citations.is_none()); + assert!(text.cache_control.is_none()); + assert_eq!(text.extra.get("future"), Some(&Value::Null)); + assert_eq!( + serde_json::to_value(block).unwrap(), + json!({"type":"text","text":"hi","future":null}) + ); +} + +#[rstest] +fn request_media_blocks_expose_sources() { + let MessagesContentPart::Image(image) = part(json!({ + "type":"image", + "source":{"type":"base64","media_type":"image/png","data":"AA=="}, + "transformations":{"oversized_image":"error"} + })) else { + panic!("expected image block"); + }; + assert!( + matches!(&image.source, ContentSource::Base64 { media_type, .. } if media_type == "image/png") + ); + assert_eq!(image.extra["transformations"]["oversized_image"], "error"); + let MessagesContentPart::Document(document) = part(json!({ + "type":"document", + "source":{"type":"content","content":[ + {"type":"text","text":"Section 1"}, + {"type":"image","source":{"type":"url","url":"https://example.test/chart.png"}} + ]}, + "title":"Q3 report", + "context":"quarterly", + "citations":{"enabled":true} + })) else { + panic!("expected document block"); + }; + assert_eq!(document.title.as_deref(), Some("Q3 report")); + assert_eq!(document.citations.as_ref().unwrap().enabled, Some(true)); + let ContentSource::Content { + content: BlockContent::Blocks(blocks), + .. + } = &document.source + else { + panic!("expected content-block source"); + }; + assert!(matches!( + blocks.as_slice(), + [MessagesContentPart::Text(_), MessagesContentPart::Image(_)] + )); + let MessagesContentPart::SearchResult(search) = part(json!({ + "type":"search_result", + "source":"https://example.test/result", + "title":"result", + "content":[{"type":"text","text":"found"}], + "citations":{"enabled":false} + })) else { + panic!("expected search result block"); + }; + assert_eq!(search.source, "https://example.test/result"); + assert!( + matches!(search.content.as_slice(), [MessagesContentPart::Text(text)] if text.text == "found") + ); +} + +#[rstest] +fn reasoning_and_tool_blocks_expose_required_fields() { + let MessagesContentPart::Thinking(thinking) = + part(json!({"type":"thinking","thinking":"plan","signature":"sig"})) + else { + panic!("expected thinking block"); + }; + assert_eq!( + (thinking.thinking.as_str(), thinking.signature.as_str()), + ("plan", "sig") + ); + let MessagesContentPart::RedactedThinking(redacted) = + part(json!({"type":"redacted_thinking","data":"opaque"})) + else { + panic!("expected redacted thinking block"); + }; + assert_eq!(redacted.data, "opaque"); + let MessagesContentPart::ToolUse(tool_use) = part(json!({ + "type":"tool_use", + "id":"toolu_1", + "name":"lookup", + "input":{"query":[1,null]}, + "caller":{"type":"code_execution_20260120","tool_id":"srvtoolu_1"}, + "toolset_name":"browser" + })) else { + panic!("expected tool use block"); + }; + assert_eq!(tool_use.input["query"], json!([1, null])); + assert!( + matches!(&tool_use.caller, Some(ToolCaller::CodeExecution20260120 { tool_id, .. }) if tool_id == "srvtoolu_1") + ); + assert_eq!(tool_use.toolset_name.as_deref(), Some("browser")); + let MessagesContentPart::ToolResult(result) = part(json!({ + "type":"tool_result", + "tool_use_id":"toolu_1", + "is_error":false, + "content":[ + {"type":"tool_reference","tool_name":"lookup"}, + {"type":"browser_state","tabs":[{"tab_id":"1","title":"","url":"","active":true}], + "state_changes":[{"type":"download_completed","download_id":"d1","url":"https://example.test/f","size_bytes":3}]} + ] + })) else { + panic!("expected tool result block"); + }; + let Some(BlockContent::Blocks(blocks)) = &result.content else { + panic!("expected nested result blocks"); + }; + let [ + MessagesContentPart::ToolReference(reference), + MessagesContentPart::BrowserState(browser), + ] = blocks.as_slice() + else { + panic!("expected tool reference and browser state"); + }; + assert_eq!(reference.tool_name, "lookup"); + assert_eq!(browser.tabs[0].active, Some(true)); + assert!(matches!( + browser.state_changes.as_deref(), + Some([BrowserStateChange::DownloadCompleted { + size_bytes: Some(3), + path: None, + .. + }]) + )); + let MessagesContentPart::ToolResult(text_result) = + part(json!({"type":"tool_result","tool_use_id":"toolu_2","content":"done"})) + else { + panic!("expected tool result block"); + }; + assert_eq!(text_result.content, Some(BlockContent::Text("done".into()))); + assert!(text_result.is_error.is_none()); +} + +#[rstest] +fn server_tool_results_expose_nested_result_unions() { + let MessagesContentPart::ServerToolUse(server) = part( + json!({"type":"server_tool_use","id":"srvtoolu_1","name":"web_search","input":{"query":"rust"}}), + ) else { + panic!("expected server tool use"); + }; + assert_eq!(server.name, "web_search"); + let MessagesContentPart::WebSearchToolResult(search) = part(json!({ + "type":"web_search_tool_result", + "tool_use_id":"srvtoolu_1", + "content":[{"type":"web_search_result","url":"https://example.test","title":"t","encrypted_content":"e","page_age":"1d"}] + })) else { + panic!("expected web search result"); + }; + let WebSearchToolResultContent::Results(results) = &search.content else { + panic!("expected search results"); + }; + assert_eq!(results[0].page_age.as_deref(), Some("1d")); + let MessagesContentPart::WebSearchToolResult(search_error) = part(json!({ + "type":"web_search_tool_result", + "tool_use_id":"srvtoolu_1", + "content":{"type":"web_search_tool_result_error","error_code":"max_uses_exceeded"} + })) else { + panic!("expected web search error"); + }; + assert!( + matches!(&search_error.content, WebSearchToolResultContent::Error(error) if error.error_code == "max_uses_exceeded") + ); + let MessagesContentPart::WebFetchToolResult(fetch) = part(json!({ + "type":"web_fetch_tool_result", + "tool_use_id":"srvtoolu_2", + "content":{"type":"web_fetch_result","url":"https://example.test","retrieved_at":"2026-01-01T00:00:00Z", + "content":{"type":"document","source":{"type":"text","media_type":"text/plain","data":"page"}}} + })) else { + panic!("expected web fetch result"); + }; + let WebFetchToolResultContent::WebFetchResult(fetched) = &fetch.content else { + panic!("expected fetched page"); + }; + let WebFetchDocument::Document(document) = &fetched.content; + assert!(matches!(&document.source, ContentSource::Text { data, .. } if data == "page")); + assert!(fetched.extra.is_empty()); + let MessagesContentPart::CodeExecutionToolResult(code) = part(json!({ + "type":"code_execution_tool_result", + "tool_use_id":"srvtoolu_3", + "content":{"type":"encrypted_code_execution_result","encrypted_stdout":"opaque","stderr":"","return_code":-1, + "content":[{"type":"code_execution_output","file_id":"file_1"}]} + })) else { + panic!("expected code execution result"); + }; + let CodeExecutionToolResultContent::EncryptedCodeExecutionResult(encrypted) = &code.content + else { + panic!("expected encrypted execution result"); + }; + assert_eq!(encrypted.return_code, -1); + assert!( + matches!(encrypted.content.as_slice(), [CodeExecutionOutput::CodeExecutionOutput { file_id, .. }] if file_id == "file_1") + ); + let MessagesContentPart::BashCodeExecutionToolResult(bash) = part(json!({ + "type":"bash_code_execution_tool_result", + "tool_use_id":"srvtoolu_7", + "content":{"type":"bash_code_execution_result","stdout":"ok","stderr":"","return_code":0, + "content":[{"type":"bash_code_execution_output","file_id":"file_2"}]} + })) else { + panic!("expected bash code execution result"); + }; + let BashCodeExecutionToolResultContent::BashCodeExecutionResult(ran) = &bash.content else { + panic!("expected bash execution result"); + }; + assert_eq!((ran.stdout.as_str(), ran.return_code), ("ok", 0)); + assert!( + matches!(ran.content.as_slice(), [BashCodeExecutionOutput::BashCodeExecutionOutput { file_id, .. }] if file_id == "file_2") + ); + assert!(ran.extra.is_empty()); + let MessagesContentPart::TextEditorCodeExecutionToolResult(editor) = part(json!({ + "type":"text_editor_code_execution_tool_result", + "tool_use_id":"srvtoolu_4", + "content":{"type":"text_editor_code_execution_view_result","content":"fn main() {}","file_type":"text","num_lines":1} + })) else { + panic!("expected text editor result"); + }; + let TextEditorCodeExecutionToolResultContent::TextEditorCodeExecutionViewResult(view) = + &editor.content + else { + panic!("expected view result"); + }; + assert_eq!(view.file_type, TextEditorFileType::Text); + assert_eq!(view.num_lines, Some(1)); + assert!(view.start_line.is_none()); + let MessagesContentPart::ToolSearchToolResult(tool_search) = part(json!({ + "type":"tool_search_tool_result", + "tool_use_id":"srvtoolu_5", + "content":{"type":"tool_search_tool_result_error","error_code":"unavailable","error_message":"down"} + })) else { + panic!("expected tool search result"); + }; + assert!( + matches!(&tool_search.content, ToolSearchToolResultContent::ToolSearchToolResultError(error) if error.error_message.as_deref() == Some("down")) + ); + let MessagesContentPart::ToolSearchToolResult(found) = part(json!({ + "type":"tool_search_tool_result", + "tool_use_id":"srvtoolu_8", + "content":{"type":"tool_search_tool_search_result","tool_references":[{"type":"tool_reference","tool_name":"lookup"}]} + })) else { + panic!("expected tool search result"); + }; + let ToolSearchToolResultContent::ToolSearchToolSearchResult(result) = &found.content else { + panic!("expected tool search hits"); + }; + let [ToolSearchReference::ToolReference(reference)] = result.tool_references.as_slice() else { + panic!("expected one tool reference"); + }; + assert_eq!(reference.tool_name, "lookup"); + assert!(reference.extra.is_empty() && result.extra.is_empty()); + let MessagesContentPart::AdvisorToolResult(advisor) = part(json!({ + "type":"advisor_tool_result", + "tool_use_id":"srvtoolu_6", + "content":{"type":"advisor_result","text":"advice"} + })) else { + panic!("expected advisor result"); + }; + assert!( + matches!(&advisor.content, AdvisorToolResultContent::AdvisorResult { text, stop_reason: None, .. } if text == "advice") + ); +} + +#[rstest] +fn beta_blocks_expose_mcp_compaction_and_fallback_fields() { + let MessagesContentPart::McpToolUse(mcp) = part( + json!({"type":"mcp_tool_use","id":"mcptoolu_1","name":"search","server_name":"kb","input":{}}), + ) else { + panic!("expected MCP tool use"); + }; + assert_eq!(mcp.server_name, "kb"); + let MessagesContentPart::McpToolResult(mcp_result) = part( + json!({"type":"mcp_tool_result","tool_use_id":"mcptoolu_1","is_error":true,"content":"failed"}), + ) else { + panic!("expected MCP tool result"); + }; + assert_eq!(mcp_result.is_error, Some(true)); + assert_eq!( + mcp_result.content, + Some(McpToolResultContent::Text("failed".into())) + ); + let MessagesContentPart::McpToolResult(text_result) = part(json!({ + "type":"mcp_tool_result", + "tool_use_id":"mcptoolu_3", + "content":[{"type":"text","text":"found"}] + })) else { + panic!("expected MCP tool result"); + }; + let Some(McpToolResultContent::Blocks(blocks)) = &text_result.content else { + panic!("expected MCP text blocks"); + }; + let [McpToolResultText::Text(text)] = blocks.as_slice() else { + panic!("expected one text block"); + }; + assert_eq!(text.text, "found"); + let MessagesContentPart::McpToolResult(empty_result) = + part(json!({"type":"mcp_tool_result","tool_use_id":"mcptoolu_2"})) + else { + panic!("expected MCP tool result"); + }; + assert!(empty_result.content.is_none()); + assert!(empty_result.extra.is_empty()); + let MessagesContentPart::McpToolListing(listing) = part(json!({ + "type":"mcp_tool_listing", + "mcp_server_name":"kb", + "tools":[{"name":"search","input_schema":{"type":"object"}}] + })) else { + panic!("expected MCP tool listing"); + }; + assert!(listing.tools[0].description.is_none()); + let MessagesContentPart::Compaction(compaction) = part(json!({ + "type":"compaction", + "content":"summary", + "encrypted_content":"opaque", + "tool_changes":[ + {"type":"tool_addition","tool":{"type":"tool_definition","definition":{"name":"lookup","input_schema":{"type":"object"}}}}, + {"type":"tool_removal","tool":{"type":"mcp_tool_reference","server_name":"kb","name":"search"}} + ] + })) else { + panic!("expected compaction block"); + }; + let [ + ToolChange::ToolAddition(addition), + ToolChange::ToolRemoval(removal), + ] = compaction.tool_changes.as_deref().unwrap() + else { + panic!("expected tool addition and removal"); + }; + let ToolChangeTarget::ToolDefinition { definition, .. } = &addition.tool else { + panic!("expected inline tool definition"); + }; + assert!( + matches!(definition.as_ref(), MessagesToolParam::Custom(tool) if tool.name == "lookup") + ); + assert!( + matches!(&removal.tool, ToolChangeTarget::McpToolReference { server_name, .. } if server_name == "kb") + ); + let failed: MessagesContentPart = + serde_json::from_value(json!({"type":"compaction","content":null})).unwrap(); + assert_eq!(failed, MessagesContentPart::Compaction(Default::default())); + let MessagesContentPart::Fallback(fallback) = part(json!({ + "type":"fallback", + "from":{"model":"claude-opus-5-5"}, + "to":{"model":"claude-sonnet-5-5"}, + "trigger":{"type":"refusal","category":"cyber"} + })) else { + panic!("expected fallback block"); + }; + assert_eq!(fallback.from.model, "claude-opus-5-5"); + assert_eq!(fallback.to.model, "claude-sonnet-5-5"); + let Some(FallbackTrigger::Refusal { category, extra }) = &fallback.trigger else { + panic!("expected refusal trigger"); + }; + assert_eq!(category.as_deref(), Some("cyber")); + assert!(extra.is_empty() && fallback.extra.is_empty()); + let MessagesContentPart::Fallback(untriggered) = part(json!({ + "type":"fallback", + "from":{"model":"claude-opus-5-5"}, + "to":{"model":"claude-sonnet-5-5"} + })) else { + panic!("expected fallback block"); + }; + assert!(untriggered.trigger.is_none()); +} + +#[rstest] +#[case::missing_tag(json!({"text":"hello"}))] +#[case::unknown_tag(json!({"type":"future_block","text":"hello"}))] +#[case::text_without_text(json!({"type":"text"}))] +#[case::wrong_text(json!({"type":"text","text":7}))] +#[case::citation_missing_location(json!({"type":"text","text":"a","citations":[{"type":"char_location","cited_text":"a","document_index":0}]}))] +#[case::unknown_citation(json!({"type":"text","text":"a","citations":[{"type":"future_location"}]}))] +#[case::thinking_without_signature(json!({"type":"thinking","thinking":"plan"}))] +#[case::tool_use_without_id(json!({"type":"tool_use","name":"lookup","input":{}}))] +#[case::tool_use_array_input(json!({"type":"tool_use","id":"t","name":"lookup","input":[]}))] +#[case::tool_result_without_id(json!({"type":"tool_result","content":"done"}))] +#[case::malformed_nested_block(json!({"type":"tool_result","tool_use_id":"t","content":[{"type":"text","text":7}]}))] +#[case::malformed_source(json!({"type":"document","source":{"type":"url","url":7}}))] +#[case::image_without_source(json!({"type":"image"}))] +#[case::search_result_without_title(json!({"type":"search_result","source":"s","content":[]}))] +#[case::web_search_result_without_url(json!({"type":"web_search_tool_result","tool_use_id":"t","content":[{"type":"web_search_result","title":"t","encrypted_content":"e"}]}))] +#[case::unknown_fetch_result(json!({"type":"web_fetch_tool_result","tool_use_id":"t","content":{"type":"future"}}))] +#[case::fractional_return_code(json!({"type":"bash_code_execution_tool_result","tool_use_id":"t","content":{"type":"bash_code_execution_result","stdout":"","stderr":"","return_code":0.5,"content":[]}}))] +#[case::unknown_file_type(json!({"type":"text_editor_code_execution_tool_result","tool_use_id":"t","content":{"type":"text_editor_code_execution_view_result","content":"","file_type":"video"}}))] +#[case::negative_line_count(json!({"type":"text_editor_code_execution_tool_result","tool_use_id":"t","content":{"type":"text_editor_code_execution_str_replace_result","new_lines":-1}}))] +#[case::unknown_tool_change(json!({"type":"compaction","tool_changes":[{"type":"tool_addition","tool":{"type":"future"}}]}))] +#[case::tool_search_non_reference(json!({"type":"tool_search_tool_result","tool_use_id":"t","content":{"type":"tool_search_tool_search_result","tool_references":[{"type":"text","text":"lookup"}]}}))] +#[case::web_fetch_text_content(json!({"type":"web_fetch_tool_result","tool_use_id":"t","content":{"type":"web_fetch_result","url":"https://example.test","content":{"type":"text","text":"page"}}}))] +#[case::bash_result_with_code_output(json!({"type":"bash_code_execution_tool_result","tool_use_id":"t","content":{"type":"bash_code_execution_result","stdout":"","stderr":"","return_code":0,"content":[{"type":"code_execution_output","file_id":"f"}]}}))] +#[case::code_result_with_bash_output(json!({"type":"code_execution_tool_result","tool_use_id":"t","content":{"type":"code_execution_result","stdout":"","stderr":"","return_code":0,"content":[{"type":"bash_code_execution_output","file_id":"f"}]}}))] +#[case::fallback_unknown_trigger(json!({"type":"fallback","from":{"model":"a"},"to":{"model":"b"},"trigger":{"type":"overload"}}))] +#[case::unknown_state_change(json!({"type":"browser_state","tabs":[],"state_changes":[{"type":"tab_closed","tab_id":"1"}]}))] +fn content_parts_reject_malformed_known_fields(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} + +#[rstest] +fn custom_tool_exposes_schema_and_preserves_extensions() { + let tool = round_trip::(json!({ + "type":"custom", + "name":"lookup", + "input_schema":{"type":"object","properties":{"query":{"type":"string"}}}, + "strict":false, + "defer_loading":true, + "allowed_callers":["direct","code_execution_20260120"], + "extension":{"nested":[1,null]} + })); + assert_eq!( + tool.allowed_callers.as_deref(), + Some([AllowedCaller::Direct, AllowedCaller::CodeExecution20260120].as_slice()) + ); + assert_eq!(tool.tool_type, Some(CustomToolType::Custom)); + assert_eq!(tool.name, "lookup"); + assert_eq!((tool.strict, tool.defer_loading), (Some(false), Some(true))); + let JsonSchema::Object(schema) = &tool.input_schema else { + panic!("expected an object schema"); + }; + assert!(schema.properties.as_ref().unwrap().contains_key("query")); + assert_eq!(tool.extra["extension"], json!({"nested":[1,null]})); +} + +#[rstest] +fn custom_tool_omits_null_discriminator() { + let tool: CustomTool = + serde_json::from_value(json!({"type":null,"name":"lookup","input_schema":true})).unwrap(); + assert!(tool.tool_type.is_none()); + assert!(tool.description.is_none()); + assert_eq!( + serde_json::to_value(tool).unwrap(), + json!({"name":"lookup","input_schema":true}) + ); +} + +#[rstest] +#[case::missing_name(json!({"input_schema":{}}))] +#[case::missing_schema(json!({"name":"lookup"}))] +#[case::wrong_name_shape(json!({"name":7,"input_schema":{}}))] +#[case::builtin_tool(json!({"name":"lookup","input_schema":{},"type":"bash_20250124"}))] +#[case::unknown_discriminator(json!({"name":"lookup","input_schema":{},"type":"future_tool"}))] +#[case::unknown_allowed_caller(json!({"name":"lookup","input_schema":{},"allowed_callers":["anyone"]}))] +fn custom_tool_rejects_invalid_shapes(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} + +#[rstest] +#[case::builtin(json!({"type":"web_search_20250305","name":"web_search","max_uses":2}), true)] +#[case::builtin_toolset(json!({"type":"mcp_toolset","mcp_server_name":"kb"}), true)] +#[case::custom(json!({"name":"lookup","input_schema":{"type":"object"}}), false)] +fn tool_params_dispatch_on_discriminator(#[case] wire: Value, #[case] builtin: bool) { + let tool = round_trip::(wire); + assert_eq!(matches!(tool, MessagesToolParam::Builtin(_)), builtin); +} + +#[rstest] +fn tool_params_reject_unknown_tool_types() { + assert!( + serde_json::from_value::( + json!({"type":"future_tool","name":"lookup","input_schema":{}}) + ) + .is_err() + ); +} + +#[rstest] +#[case::image(json!([{"type":"image","source":{"type":"url","url":"https://example.test/a.png"}}]))] +#[case::tool_reference(json!([{"type":"tool_reference","tool_name":"lookup"}]))] +#[case::untagged(json!([{"text":"found"}]))] +fn mcp_tool_result_content_rejects_non_text_blocks(#[case] content: Value) { + assert!( + serde_json::from_value::( + json!({"type":"mcp_tool_result","tool_use_id":"mcptoolu_1","content":content}) + ) + .is_err() + ); +} + +#[rstest] +#[case::model(json!({"type":"model_changed","cache_missed_input_tokens":12}), Some(12))] +#[case::system(json!({"type":"system_changed","cache_missed_input_tokens":3}), Some(3))] +#[case::tools(json!({"type":"tools_changed","cache_missed_input_tokens":0}), Some(0))] +#[case::messages(json!({"type":"messages_changed","cache_missed_input_tokens":9}), Some(9))] +#[case::not_found(json!({"type":"previous_message_not_found"}), None)] +#[case::unavailable(json!({"type":"unavailable"}), None)] +fn diagnostics_expose_cache_miss_reason(#[case] reason: Value, #[case] missed: Option) { + let diagnostics = round_trip::(json!({"cache_miss_reason":reason})); + let Some(Recognized::Known(reason)) = &diagnostics.cache_miss_reason else { + panic!("expected known cache miss reason"); + }; + let actual = match reason { + CacheMissReason::ModelChanged(tokens) + | CacheMissReason::SystemChanged(tokens) + | CacheMissReason::ToolsChanged(tokens) + | CacheMissReason::MessagesChanged(tokens) => { + assert!(tokens.extra.is_empty()); + Some(tokens.cache_missed_input_tokens) + } + CacheMissReason::PreviousMessageNotFound { extra } + | CacheMissReason::Unavailable { extra } => { + assert!(extra.is_empty()); + None + } + }; + assert_eq!(actual, missed); + assert!(diagnostics.extra.is_empty()); +} + +#[rstest] +#[case::model(json!({"type":"model_changed","cache_missed_input_tokens":12}), "model_changed")] +#[case::system(json!({"type":"system_changed","cache_missed_input_tokens":12}), "system_changed")] +#[case::tools(json!({"type":"tools_changed","cache_missed_input_tokens":12}), "tools_changed")] +#[case::messages(json!({"type":"messages_changed","cache_missed_input_tokens":12}), "messages_changed")] +fn cache_miss_reason_variants_keep_their_tags(#[case] reason: Value, #[case] tag: &str) { + let parsed: CacheMissReason = serde_json::from_value(reason).unwrap(); + assert_eq!(serde_json::to_value(parsed).unwrap()["type"], json!(tag)); +} + +#[rstest] +#[case::pending(json!({"cache_miss_reason":null}))] +#[case::future_reason(json!({"cache_miss_reason":{"type":"future_changed","cache_missed_input_tokens":1}}))] +#[case::missing_tokens(json!({"cache_miss_reason":{"type":"model_changed"}}))] +fn diagnostics_keep_pending_and_unmodeled_reasons(#[case] wire: Value) { + let diagnostics: MessagesDiagnostics = serde_json::from_value(wire.clone()).unwrap(); + match &diagnostics.cache_miss_reason { + None => assert_eq!(serde_json::to_value(&diagnostics).unwrap(), json!({})), + Some(Recognized::Unrecognized(kept)) => { + assert_eq!(kept, &wire["cache_miss_reason"]); + assert_eq!(serde_json::to_value(&diagnostics).unwrap(), wire); + } + Some(Recognized::Known(reason)) => panic!("unexpected known reason {reason:?}"), + } +} + +#[rstest] +#[case::previous(json!({"previous_message_id":"msg_1"}), Some(Some("msg_1")))] +#[case::first_turn(json!({"previous_message_id":null}), Some(None))] +#[case::absent(json!({}), None)] +fn diagnostics_param_distinguishes_null_from_absent_previous_message( + #[case] wire: Value, + #[case] expected: Option>, +) { + let param = round_trip::(wire); + assert_eq!( + param.previous_message_id.as_ref().map(Option::as_deref), + expected + ); + assert!(param.extra.is_empty()); +} + +#[rstest] +fn diagnostics_param_rejects_non_string_previous_message() { + assert!( + serde_json::from_value::(json!({"previous_message_id":7})) + .is_err() + ); +} + +#[rstest] +fn cache_miss_reason_rejects_non_numeric_tokens() { + assert!( + serde_json::from_value::( + json!({"type":"model_changed","cache_missed_input_tokens":"many"}) + ) + .is_err() + ); +} diff --git a/litellm-rust/crates/llms-types/tests/minimax.rs b/litellm-rust/crates/llms-types/tests/minimax.rs new file mode 100644 index 00000000000..a7bd46f0934 --- /dev/null +++ b/litellm-rust/crates/llms-types/tests/minimax.rs @@ -0,0 +1,82 @@ +use litellm_llms_types::providers::minimax::{ + MinimaxMediaDetail, MinimaxMediaSource, MinimaxMessagesContentBlock, +}; +use rstest::rstest; +use serde_json::{Value, json}; + +#[rstest] +#[case::image(json!({"type":"image","source":{"type":"base64","media_type":"image/png","data":"AA==","detail":"low"}}))] +#[case::video(json!({"type":"video","source":{"type":"url","url":"https://example.test/video","detail":"high","fps":1,"max_long_side_pixel":1024,"future":null},"cache_control":{"type":"ephemeral"}}))] +#[case::mid_conversation_system(json!({"type":"mid_conv_system","text":"instruction","future":null}))] +fn provider_content_blocks_round_trip(#[case] wire: Value) { + let block: MinimaxMessagesContentBlock = serde_json::from_value(wire.clone()).unwrap(); + match &block { + MinimaxMessagesContentBlock::Image(image) => { + let MinimaxMediaSource::Base64 { + media_type, + data, + options, + } = &image.source + else { + panic!("expected base64 image source"); + }; + assert_eq!((media_type.as_str(), data.as_str()), ("image/png", "AA==")); + assert_eq!(options.detail, Some(MinimaxMediaDetail::Low)); + assert!(options.extra.is_empty()); + } + MinimaxMessagesContentBlock::Video(video) => { + let MinimaxMediaSource::Url { url, options } = &video.source else { + panic!("expected URL video source"); + }; + assert_eq!(url, "https://example.test/video"); + assert_eq!(options.detail, Some(MinimaxMediaDetail::High)); + assert_eq!(options.fps, Some(1.into())); + assert_eq!(options.max_long_side_pixel, Some(1024)); + assert_eq!(options.extra["future"], Value::Null); + assert_eq!( + video.cache_control.as_ref().unwrap().cache_type.as_deref(), + Some("ephemeral") + ); + } + MinimaxMessagesContentBlock::MidConvSystem { text, extra } => { + assert_eq!(text, "instruction"); + assert_eq!(extra["future"], Value::Null); + } + } + assert_eq!(serde_json::to_value(block).unwrap(), wire); +} + +#[rstest] +#[case::missing_source(json!({"type":"video"}))] +#[case::bad_source_tag(json!({"type":"image","source":{"type":"future"}}))] +#[case::url_without_url(json!({"type":"video","source":{"type":"url"}}))] +#[case::base64_without_data(json!({"type":"image","source":{"type":"base64","media_type":"image/png"}}))] +#[case::bad_detail(json!({"type":"image","source":{"type":"url","url":"u","detail":7}}))] +#[case::bad_fps(json!({"type":"video","source":{"type":"url","url":"u","fps":"fast"}}))] +#[case::missing_text(json!({"type":"mid_conv_system"}))] +fn provider_content_rejects_malformed_fields(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} + +#[rstest] +fn partial_media_options_omit_null_optionals() { + let block: MinimaxMessagesContentBlock = serde_json::from_value(json!({ + "type":"video", + "source":{"type":"url","url":"u","detail":null,"fps":null,"future":null}, + "cache_control":null + })) + .unwrap(); + let MinimaxMessagesContentBlock::Video(video) = &block else { + panic!("expected video") + }; + let MinimaxMediaSource::Url { options, .. } = &video.source else { + panic!("expected URL source"); + }; + assert!(options.detail.is_none()); + assert!(options.fps.is_none()); + assert!(video.cache_control.is_none()); + assert_eq!( + serde_json::to_value(block).unwrap(), + json!({"type":"video","source":{"type":"url","url":"u","future":null}}) + ); +} diff --git a/litellm-rust/crates/llms-types/tests/ocr.rs b/litellm-rust/crates/llms-types/tests/ocr.rs index 48c816819f2..03e8a586a09 100644 --- a/litellm-rust/crates/llms-types/tests/ocr.rs +++ b/litellm-rust/crates/llms-types/tests/ocr.rs @@ -1,4 +1,4 @@ -use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrDocument, OcrPage}; +use litellm_llms_types::formats::ocr::{LiteLLMOcrResponse, OcrBoundingBox, OcrDocument, OcrPage}; use rstest::rstest; use serde_json::{Map, Value, json}; @@ -102,3 +102,39 @@ fn response_serialization_preserves_extensions_and_native_presence( assert_eq!(decoded.provider_native_response, native); assert_eq!(decoded.into_json(), serialized); } + +#[rstest] +fn bounding_box_exposes_corner_coordinates_and_keeps_extensions() { + let wire = json!({"top_left_x":1,"top_left_y":2.5,"bottom_right_x":30,"bottom_right_y":40,"future":true}); + let bounds: OcrBoundingBox = serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(bounds.top_left_x, Some(1.into())); + assert_eq!( + bounds + .top_left_y + .as_ref() + .and_then(serde_json::Number::as_f64), + Some(2.5) + ); + assert_eq!(bounds.bottom_right_x, Some(30.into())); + assert_eq!(bounds.bottom_right_y, Some(40.into())); + assert_eq!(Value::Object(bounds.extra.clone()), json!({"future":true})); + assert_eq!(serde_json::to_value(bounds).unwrap(), wire); +} + +#[rstest] +fn partial_bounding_box_omits_null_corners() { + let bounds: OcrBoundingBox = + serde_json::from_value(json!({"top_left_x":null,"bottom_right_y":4})).unwrap(); + assert!(bounds.top_left_x.is_none()); + assert_eq!( + serde_json::to_value(bounds).unwrap(), + json!({"bottom_right_y":4}) + ); +} + +#[rstest] +#[case::string_corner(json!({"top_left_x":"1"}))] +#[case::array_corner(json!({"bottom_right_y":[4]}))] +fn bounding_box_rejects_non_numeric_corners(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} diff --git a/litellm-rust/crates/llms-types/tests/responses.rs b/litellm-rust/crates/llms-types/tests/responses.rs index 92271757b3f..73cf5e15f6f 100644 --- a/litellm-rust/crates/llms-types/tests/responses.rs +++ b/litellm-rust/crates/llms-types/tests/responses.rs @@ -1,5 +1,31 @@ -use litellm_llms_types::formats::responses::streaming_websocket::ResponsesWsEventType; +use litellm_llms_types::formats::chat_completions::{PromptCacheBreakpoint, PromptCacheMode}; +use litellm_llms_types::formats::responses::{ + ResponsesAdditionalTools, ResponsesApplyPatchCall, ResponsesApplyPatchCallOutput, + ResponsesApplyPatchOperation, ResponsesCodeOutput, ResponsesCompaction, + ResponsesComputerAction, ResponsesComputerCall, ResponsesComputerCallOutput, + ResponsesComputerOutput, ResponsesContentPart, ResponsesCoordinate, ResponsesCustomToolCall, + ResponsesCustomToolCallOutput, ResponsesFunctionCall, ResponsesFunctionCallOutput, + ResponsesImageGenerationCall, ResponsesInputContent, ResponsesLocalShellAction, + ResponsesLocalShellCall, ResponsesLocalShellCallOutput, ResponsesMcpApprovalRequest, + ResponsesMcpApprovalResponse, ResponsesMcpError, ResponsesMcpErrorDetail, ResponsesOutputItem, + ResponsesProgram, ResponsesProgramOutput, ResponsesSafetyCheck, ResponsesShellAction, + ResponsesShellCall, ResponsesShellCallOutput, ResponsesShellEnvironment, ResponsesShellOutcome, + ResponsesShellOutputChunk, ResponsesToolCaller, ResponsesToolOutput, ResponsesToolSearchCall, + ResponsesToolSearchOutput, ResponsesWebSearchAction, + streaming_websocket::{ResponsesEventResponse, ResponsesWsEventType}, +}; +use litellm_llms_types::recognized::Recognized; use rstest::rstest; +use serde::{Serialize, de::DeserializeOwned}; +use serde_json::{Value, json}; + +fn round_trip(wire: Value) +where + T: DeserializeOwned + Serialize, +{ + let parsed: T = serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(serde_json::to_value(parsed).unwrap(), wire); +} #[rstest] #[case::create("response.create", ResponsesWsEventType::ResponseCreate)] @@ -46,3 +72,646 @@ fn websocket_event_type_schema_is_open_string() { Some(&serde_json::json!("string")) ); } + +#[rstest] +#[case::message(json!({"type":"message","role":"assistant","content":[{"type":"output_text","text":"answer","annotations":[{"type":"url_citation","url":"https://example.test","start_index":0,"end_index":6}]}]}))] +#[case::function_call(json!({"type":"function_call","call_id":"call_1","name":"lookup","arguments":"{\"query\":\"q\"}"}))] +#[case::custom_tool(json!({"type":"custom_tool_call","name":"lookup","input":"q"}))] +#[case::reasoning(json!({"type":"reasoning","summary":[{"type":"summary_text","text":"summary"}],"encrypted_content":"opaque"}))] +#[case::web_search(json!({"type":"web_search_call","action":{"type":"search","queries":["q"],"sources":[{"type":"url","url":"https://example.test"}]}}))] +#[case::file_search(json!({"type":"file_search_call","queries":["q"],"results":[{"file_id":"file_1","score":1,"attributes":{"custom":[1,null]}}]}))] +#[case::code(json!({"type":"code_interpreter_call","outputs":[{"type":"logs","logs":"done"},{"type":"image","url":"https://example.test"}]}))] +#[case::image(json!({"type":"image_generation_call","result":"generated"}))] +#[case::mcp(json!({"type":"mcp_call","server_label":"server","name":"lookup","arguments":"{}","output":"done"}))] +fn output_items_round_trip(#[case] wire: Value) { + let item: ResponsesOutputItem = serde_json::from_value(wire.clone()).unwrap(); + match &item { + ResponsesOutputItem::Message(message) => { + assert_eq!(message.role.as_deref(), Some("assistant")); + let Some(content) = &message.content else { + panic!("expected content") + }; + let [ + ResponsesContentPart::OutputText { + text, + annotations: Some(annotations), + .. + }, + ] = content.as_slice() + else { + panic!("expected output text and annotations") + }; + assert_eq!(text, "answer"); + let [ + litellm_llms_types::formats::responses::ResponsesAnnotation::UrlCitation(citation), + ] = annotations.as_slice() + else { + panic!("expected URL citation") + }; + assert_eq!(citation.url.as_deref(), Some("https://example.test")); + assert_eq!(citation.start_index, Some(0)); + } + ResponsesOutputItem::FunctionCall(call) => { + assert_eq!(call.call_id.as_deref(), Some("call_1")); + assert_eq!(call.name.as_deref(), Some("lookup")); + assert_eq!(call.arguments.as_deref(), Some("{\"query\":\"q\"}")); + } + ResponsesOutputItem::CustomToolCall(call) => { + assert_eq!(call.name.as_deref(), Some("lookup")); + assert_eq!(call.input.as_deref(), Some("q")); + } + ResponsesOutputItem::Reasoning(reasoning) => { + assert_eq!(reasoning.encrypted_content.as_deref(), Some("opaque")); + let Some(summary) = &reasoning.summary else { + panic!("expected summary") + }; + let [ResponsesContentPart::SummaryText { text, .. }] = summary.as_slice() else { + panic!("expected summary text") + }; + assert_eq!(text, "summary"); + } + ResponsesOutputItem::WebSearchCall(call) => { + let Some(ResponsesWebSearchAction::Search { + queries: Some(queries), + sources: Some(sources), + .. + }) = &call.action + else { + panic!("expected search action") + }; + assert_eq!(queries, &["q"]); + assert_eq!(sources[0].url.as_deref(), Some("https://example.test")); + } + ResponsesOutputItem::FileSearchCall(call) => { + assert_eq!(call.queries.as_deref(), Some(["q".to_owned()].as_slice())); + let Some(results) = &call.results else { + panic!("expected search results") + }; + assert_eq!(results[0].file_id.as_deref(), Some("file_1")); + assert_eq!(results[0].score, Some(1.into())); + assert_eq!( + results[0].attributes.as_ref().unwrap()["custom"], + json!([1, null]) + ); + } + ResponsesOutputItem::CodeInterpreterCall(call) => { + let Some(outputs) = &call.outputs else { + panic!("expected code outputs") + }; + let [ + ResponsesCodeOutput::Logs { logs, .. }, + ResponsesCodeOutput::Image { url, .. }, + ] = outputs.as_slice() + else { + panic!("expected logs and image") + }; + assert_eq!(logs, "done"); + assert_eq!(url, "https://example.test"); + } + ResponsesOutputItem::ImageGenerationCall(call) => { + assert_eq!(call.result.as_deref(), Some("generated")) + } + ResponsesOutputItem::McpCall(call) => { + assert_eq!(call.server_label.as_deref(), Some("server")); + assert_eq!(call.name.as_deref(), Some("lookup")); + assert_eq!(call.arguments.as_deref(), Some("{}")); + assert_eq!(call.output.as_deref(), Some("done")); + } + other => panic!("unexpected variant {other:?}"), + } + assert_eq!(serde_json::to_value(item).unwrap(), wire); +} + +#[rstest] +fn nested_event_response_exposes_typed_output_and_preserves_extensions() { + let wire = json!({ + "id":"response_1", + "model":"example-model", + "status":"completed", + "output":[{"type":"function_call","call_id":"call_1","name":"lookup","arguments":"{}","extension":true}], + "extension":{"nested":[1,null]} + }); + let response: ResponsesEventResponse = serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(response.id.as_deref(), Some("response_1")); + assert_eq!(response.model.as_deref(), Some("example-model")); + let Some(output) = &response.output else { + panic!("expected typed output"); + }; + let [Recognized::Known(ResponsesOutputItem::FunctionCall(call))] = output.as_slice() else { + panic!("expected function call"); + }; + assert_eq!(call.name.as_deref(), Some("lookup")); + assert_eq!(call.arguments.as_deref(), Some("{}")); + assert_eq!(serde_json::to_value(response).unwrap(), wire); +} + +#[rstest] +#[case::empty(json!({}))] +#[case::partial(json!({"id":"response_1","output":[]}))] +fn nested_event_response_accepts_partial_metadata(#[case] wire: Value) { + round_trip::(wire); +} + +#[rstest] +#[case::wrong_model(json!({"model":7}))] +#[case::wrong_status(json!({"status":false}))] +#[case::wrong_output(json!({"output":{}}))] +fn event_response_rejects_malformed_typed_fields(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} + +#[rstest] +#[case::unknown_type(json!({"type":"future_item","id":"item_1","payload":[1,null]}))] +#[case::missing_tag(json!({"id":"item_1"}))] +#[case::malformed_known(json!({"type":"message","content":[{"type":"output_text","text":7}]}))] +fn event_response_keeps_unmodeled_output_items_beside_typed_ones(#[case] item: Value) { + let wire = json!({"output":[{"type":"function_call","name":"lookup"}, item.clone()]}); + let response: ResponsesEventResponse = serde_json::from_value(wire.clone()).unwrap(); + let Some( + [ + Recognized::Known(ResponsesOutputItem::FunctionCall(call)), + Recognized::Unrecognized(kept), + ], + ) = response.output.as_deref() + else { + panic!("expected one typed item and one preserved item"); + }; + assert_eq!(call.name.as_deref(), Some("lookup")); + assert_eq!(kept, &item); + assert_eq!(serde_json::to_value(response).unwrap(), wire); +} + +#[rstest] +#[case::find_in_page( + json!({"type":"find_in_page","url":"https://example.test","pattern":"needle"}), + json!({"type":"find_in_page","url":"https://example.test","pattern":"needle"}) +)] +#[case::legacy_find( + json!({"type":"find","url":"https://example.test","pattern":"needle"}), + json!({"type":"find_in_page","url":"https://example.test","pattern":"needle"}) +)] +fn web_search_find_action_exposes_url_and_pattern(#[case] wire: Value, #[case] serialized: Value) { + let action: ResponsesWebSearchAction = serde_json::from_value(wire).unwrap(); + let ResponsesWebSearchAction::FindInPage { + url, + pattern, + extra, + } = &action + else { + panic!("expected find_in_page action"); + }; + assert_eq!(url, "https://example.test"); + assert_eq!(pattern, "needle"); + assert!(extra.is_empty()); + assert_eq!(serde_json::to_value(action).unwrap(), serialized); +} + +#[rstest] +#[case::with_url(json!({"type":"open_page","url":"https://example.test"}), Some("https://example.test"))] +#[case::without_url(json!({"type":"open_page"}), None)] +fn web_search_open_page_url_is_optional(#[case] wire: Value, #[case] expected: Option<&str>) { + let action: ResponsesWebSearchAction = serde_json::from_value(wire.clone()).unwrap(); + let ResponsesWebSearchAction::OpenPage { url, .. } = &action else { + panic!("expected open_page action"); + }; + assert_eq!(url.as_deref(), expected); + assert_eq!(serde_json::to_value(action).unwrap(), wire); +} + +#[rstest] +#[case::find_pattern(json!({"type":"find_in_page","url":"u","pattern":false}))] +#[case::find_missing_url(json!({"type":"find_in_page","pattern":"p"}))] +#[case::open_page_url(json!({"type":"open_page","url":7}))] +fn web_search_action_rejects_malformed_fields(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} + +#[rstest] +fn event_response_optional_fields_omit_missing_and_null() { + let response: ResponsesEventResponse = serde_json::from_value(json!({ + "id":null,"model":null,"status":null,"output":null,"future":null + })) + .unwrap(); + assert!(response.id.is_none()); + assert!(response.model.is_none()); + assert!(response.status.is_none()); + assert!(response.output.is_none()); + assert_eq!( + serde_json::to_value(response).unwrap(), + json!({"future":null}) + ); +} + +#[rstest] +#[case::refusal(json!({"type":"refusal","refusal":"refused","future":null}), ResponsesContentPart::Refusal { refusal:"refused".into(), extra:serde_json::Map::from_iter([("future".into(), Value::Null)]) })] +#[case::reasoning(json!({"type":"reasoning_text","text":"reason"}), ResponsesContentPart::ReasoningText { text:"reason".into(), extra:Default::default() })] +fn content_parts_expose_typed_variants( + #[case] wire: Value, + #[case] expected: ResponsesContentPart, +) { + let content: ResponsesContentPart = serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(content, expected); + assert_eq!(serde_json::to_value(content).unwrap(), wire); +} + +#[rstest] +fn mcp_discovery_exposes_tools_and_nested_schemas() { + let wire = json!({ + "type":"mcp_list_tools","id":"item_1","server_label":"tools", + "tools":[{"name":"lookup","description":"Look up a value","input_schema":{"type":"object","properties":{"query":{"type":"string"}}},"annotations":{"read_only":false},"future":null}], + "extension":[1,null] + }); + let item: ResponsesOutputItem = serde_json::from_value(wire.clone()).unwrap(); + let ResponsesOutputItem::McpListTools(discovery) = &item else { + panic!("expected discovery") + }; + assert_eq!(discovery.server_label.as_deref(), Some("tools")); + let tool = &discovery.tools.as_ref().unwrap()[0]; + assert_eq!(tool.name.as_deref(), Some("lookup")); + let Some(litellm_llms_types::json_schema::JsonSchema::Object(schema)) = &tool.input_schema + else { + panic!("expected tool schema") + }; + assert!(schema.properties.as_ref().unwrap().contains_key("query")); + assert_eq!(tool.annotations, Some(json!({"read_only":false}))); + assert_eq!(tool.extra["future"], Value::Null); + assert_eq!(serde_json::to_value(item).unwrap(), wire); +} + +#[rstest] +#[case::partial(json!({"server_label":"tools","tools":[]}))] +#[case::failure(json!({"server_label":"tools","error":"unavailable"}))] +fn mcp_discovery_accepts_partial_payloads(#[case] wire: Value) { + let discovery: litellm_llms_types::formats::responses::ResponsesMcpListTools = + serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(discovery.server_label.as_deref(), Some("tools")); + assert!(discovery.id.is_none()); + assert_eq!(serde_json::to_value(discovery).unwrap(), wire); +} + +#[rstest] +#[case::message(json!("unavailable"), ResponsesMcpError::Message("unavailable".into()))] +#[case::protocol(json!({"type":"mcp_protocol_error","code":-32600,"message":"invalid","extension":null}), ResponsesMcpError::Detail(ResponsesMcpErrorDetail::McpProtocolError {code:-32600,message:"invalid".into(),extra:serde_json::Map::from_iter([("extension".into(), Value::Null)])}))] +#[case::http(json!({"type":"http_error","code":503,"message":"unavailable"}), ResponsesMcpError::Detail(ResponsesMcpErrorDetail::HttpError {code:503,message:"unavailable".into(),extra:Default::default()}))] +#[case::tool(json!({"type":"mcp_tool_execution_error","content":{"arbitrary":[1,null]}}), ResponsesMcpError::Detail(ResponsesMcpErrorDetail::McpToolExecutionError {content:json!({"arbitrary":[1,null]}),extra:Default::default()}))] +fn mcp_call_errors_expose_typed_variants(#[case] wire: Value, #[case] expected: ResponsesMcpError) { + let error: ResponsesMcpError = serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(error, expected); + assert_eq!(serde_json::to_value(&error).unwrap(), wire); + let call: ResponsesOutputItem = + serde_json::from_value(json!({"type":"mcp_call","error":wire})).unwrap(); + let ResponsesOutputItem::McpCall(call) = call else { + panic!("expected call") + }; + assert_eq!(call.error, Some(error)); +} + +#[rstest] +#[case::tools_not_array(json!({"type":"mcp_list_tools","tools":{}}))] +#[case::invalid_name(json!({"type":"mcp_list_tools","tools":[{"name":7}]}))] +#[case::invalid_schema(json!({"type":"mcp_list_tools","tools":[{"input_schema":{"properties":{"query":7}}}]}))] +#[case::invalid_error_code(json!({"type":"mcp_call","error":{"type":"http_error","code":"503","message":"unavailable"}}))] +#[case::missing_error_content(json!({"type":"mcp_call","error":{"type":"mcp_tool_execution_error"}}))] +#[case::unknown_error_tag(json!({"type":"mcp_call","error":{"type":"future"}}))] +fn mcp_output_rejects_malformed_known_fields(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} + +fn text(value: &str) -> Option { + Some(value.to_owned()) +} + +fn program_caller() -> Option { + Some(ResponsesToolCaller::Program { + caller_id: "prog_1".into(), + extra: Default::default(), + }) +} + +fn direct_caller() -> Option { + Some(ResponsesToolCaller::Direct { + extra: Default::default(), + }) +} + +fn safety_check() -> ResponsesSafetyCheck { + ResponsesSafetyCheck { + id: text("sc_1"), + code: text("malicious_instructions"), + message: text("check"), + ..Default::default() + } +} + +#[rstest] +#[case::function_call( + json!({"type":"function_call","id":"fc_1","status":"completed","call_id":"call_1","name":"lookup","arguments":"{}","namespace":"ns","async":true,"caller":{"type":"program","caller_id":"prog_1"}}), + ResponsesOutputItem::FunctionCall(ResponsesFunctionCall { + id: text("fc_1"), status: text("completed"), call_id: text("call_1"), name: text("lookup"), + arguments: text("{}"), namespace: text("ns"), r#async: Some(true), caller: program_caller(), ..Default::default() + }) +)] +#[case::custom_tool_call( + json!({"type":"custom_tool_call","call_id":"call_1","name":"lookup","input":"q","namespace":"ns","async":false,"caller":{"type":"direct"}}), + ResponsesOutputItem::CustomToolCall(ResponsesCustomToolCall { + call_id: text("call_1"), name: text("lookup"), input: text("q"), namespace: text("ns"), + r#async: Some(false), caller: direct_caller(), ..Default::default() + }) +)] +#[case::image_generation_call( + json!({"type":"image_generation_call","id":"ig_1","status":"completed","result":"b64","action":"edit","background":"opaque","output_format":"webp","quality":"high","revised_prompt":"a cat","size":"1536x864"}), + ResponsesOutputItem::ImageGenerationCall(ResponsesImageGenerationCall { + id: text("ig_1"), status: text("completed"), result: text("b64"), action: text("edit"), background: text("opaque"), + output_format: text("webp"), quality: text("high"), revised_prompt: text("a cat"), size: text("1536x864"), ..Default::default() + }) +)] +#[case::function_call_output_text( + json!({"type":"function_call_output","id":"fco_1","status":"completed","call_id":"call_1","output":"done","caller":{"type":"direct"},"created_by":"user_1","name":"lookup","namespace":"ns"}), + ResponsesOutputItem::FunctionCallOutput(ResponsesFunctionCallOutput { + id: text("fco_1"), status: text("completed"), call_id: text("call_1"), output: Some(ResponsesToolOutput::Text("done".into())), + caller: direct_caller(), created_by: text("user_1"), name: text("lookup"), namespace: text("ns"), ..Default::default() + }) +)] +#[case::function_call_output_content( + json!({"type":"function_call_output","call_id":"call_1","output":[ + {"type":"input_text","text":"t","prompt_cache_breakpoint":{"mode":"explicit"}}, + {"type":"input_image","detail":"low","file_id":"file_1","image_url":"https://example.test/i.png"}, + {"type":"input_file","detail":"high","file_data":"data","file_id":"file_2","file_url":"https://example.test/f","filename":"f.pdf"} + ]}), + ResponsesOutputItem::FunctionCallOutput(ResponsesFunctionCallOutput { + call_id: text("call_1"), + output: Some(ResponsesToolOutput::Content(vec![ + ResponsesInputContent::InputText { + text: "t".into(), + prompt_cache_breakpoint: Some(PromptCacheBreakpoint { mode: PromptCacheMode::Explicit, extra: Default::default() }), + extra: Default::default(), + }, + ResponsesInputContent::InputImage { + detail: "low".into(), file_id: text("file_1"), image_url: text("https://example.test/i.png"), + prompt_cache_breakpoint: None, extra: Default::default(), + }, + ResponsesInputContent::InputFile { + detail: text("high"), file_data: text("data"), file_id: text("file_2"), file_url: text("https://example.test/f"), + filename: text("f.pdf"), prompt_cache_breakpoint: None, extra: Default::default(), + }, + ])), + ..Default::default() + }) +)] +#[case::custom_tool_call_output( + json!({"type":"custom_tool_call_output","id":"cto_1","status":"completed","call_id":"call_1","output":[{"type":"input_text","text":"t"}],"caller":{"type":"program","caller_id":"prog_1"},"created_by":"user_1"}), + ResponsesOutputItem::CustomToolCallOutput(ResponsesCustomToolCallOutput { + id: text("cto_1"), status: text("completed"), call_id: text("call_1"), + output: Some(ResponsesToolOutput::Content(vec![ResponsesInputContent::InputText { + text: "t".into(), prompt_cache_breakpoint: None, extra: Default::default(), + }])), + caller: program_caller(), created_by: text("user_1"), ..Default::default() + }) +)] +#[case::computer_call( + json!({"type":"computer_call","id":"cu_1","status":"completed","call_id":"call_1", + "pending_safety_checks":[{"id":"sc_1","code":"malicious_instructions","message":"check"}], + "action":{"type":"click","button":"left","x":1,"y":2,"keys":["shift"]}, + "actions":[ + {"type":"double_click","x":3,"y":4}, + {"type":"drag","path":[{"x":5,"y":6},{"x":7,"y":8}],"keys":["ctrl"]}, + {"type":"keypress","keys":["enter"]}, + {"type":"move","x":-1,"y":9}, + {"type":"screenshot"}, + {"type":"scroll","scroll_x":0,"scroll_y":-10,"x":11,"y":12}, + {"type":"type","text":"hello"}, + {"type":"wait"} + ]}), + ResponsesOutputItem::ComputerCall(ResponsesComputerCall { + id: text("cu_1"), status: text("completed"), call_id: text("call_1"), pending_safety_checks: Some(vec![safety_check()]), + action: Some(ResponsesComputerAction::Click { button: "left".into(), x: 1, y: 2, keys: Some(vec!["shift".into()]), extra: Default::default() }), + actions: Some(vec![ + ResponsesComputerAction::DoubleClick { x: 3, y: 4, keys: None, extra: Default::default() }, + ResponsesComputerAction::Drag { + path: vec![ + ResponsesCoordinate { x: 5, y: 6, extra: Default::default() }, + ResponsesCoordinate { x: 7, y: 8, extra: Default::default() }, + ], + keys: Some(vec!["ctrl".into()]), + extra: Default::default(), + }, + ResponsesComputerAction::Keypress { keys: vec!["enter".into()], extra: Default::default() }, + ResponsesComputerAction::Move { x: -1, y: 9, keys: None, extra: Default::default() }, + ResponsesComputerAction::Screenshot { extra: Default::default() }, + ResponsesComputerAction::Scroll { scroll_x: 0, scroll_y: -10, x: 11, y: 12, keys: None, extra: Default::default() }, + ResponsesComputerAction::Type { text: "hello".into(), extra: Default::default() }, + ResponsesComputerAction::Wait { extra: Default::default() }, + ]), + ..Default::default() + }) +)] +#[case::computer_call_output( + json!({"type":"computer_call_output","id":"cuo_1","status":"completed","call_id":"call_1","output":{"type":"computer_screenshot","file_id":"file_1","image_url":"https://example.test/s.png"},"acknowledged_safety_checks":[{"id":"sc_1","code":"malicious_instructions","message":"check"}],"created_by":"user_1"}), + ResponsesOutputItem::ComputerCallOutput(ResponsesComputerCallOutput { + id: text("cuo_1"), status: text("completed"), call_id: text("call_1"), + output: Some(ResponsesComputerOutput::ComputerScreenshot { file_id: text("file_1"), image_url: text("https://example.test/s.png"), extra: Default::default() }), + acknowledged_safety_checks: Some(vec![safety_check()]), created_by: text("user_1"), ..Default::default() + }) +)] +#[case::program( + json!({"type":"program","id":"prog_item","call_id":"prog_1","code":"run()","fingerprint":"fp"}), + ResponsesOutputItem::Program(ResponsesProgram { + id: text("prog_item"), call_id: text("prog_1"), code: text("run()"), fingerprint: text("fp"), ..Default::default() + }) +)] +#[case::program_output( + json!({"type":"program_output","id":"po_1","status":"completed","call_id":"prog_1","result":"42"}), + ResponsesOutputItem::ProgramOutput(ResponsesProgramOutput { + id: text("po_1"), status: text("completed"), call_id: text("prog_1"), result: text("42"), ..Default::default() + }) +)] +#[case::tool_search_call( + json!({"type":"tool_search_call","id":"ts_1","status":"completed","call_id":"call_1","arguments":{"query":["weather",null]},"execution":"server","created_by":"user_1"}), + ResponsesOutputItem::ToolSearchCall(ResponsesToolSearchCall { + id: text("ts_1"), status: text("completed"), call_id: text("call_1"), arguments: Some(json!({"query":["weather",null]})), + execution: text("server"), created_by: text("user_1"), ..Default::default() + }) +)] +#[case::tool_search_output( + json!({"type":"tool_search_output","id":"tso_1","status":"completed","call_id":"call_1","execution":"client","tools":[{"type":"function","name":"lookup","parameters":null,"strict":true}],"created_by":"user_1"}), + ResponsesOutputItem::ToolSearchOutput(ResponsesToolSearchOutput { + id: text("tso_1"), status: text("completed"), call_id: text("call_1"), execution: text("client"), + tools: Some(vec![serde_json::Map::from_iter([ + ("type".into(), json!("function")), ("name".into(), json!("lookup")), + ("parameters".into(), Value::Null), ("strict".into(), json!(true)), + ])]), + created_by: text("user_1"), ..Default::default() + }) +)] +#[case::additional_tools( + json!({"type":"additional_tools","id":"at_1","role":"developer","tools":[{"type":"local_shell"}]}), + ResponsesOutputItem::AdditionalTools(ResponsesAdditionalTools { + id: text("at_1"), role: text("developer"), + tools: Some(vec![serde_json::Map::from_iter([("type".into(), json!("local_shell"))])]), ..Default::default() + }) +)] +#[case::compaction( + json!({"type":"compaction","id":"cmp_1","encrypted_content":"opaque","created_by":"user_1"}), + ResponsesOutputItem::Compaction(ResponsesCompaction { + id: text("cmp_1"), encrypted_content: text("opaque"), created_by: text("user_1"), ..Default::default() + }) +)] +#[case::local_shell_call( + json!({"type":"local_shell_call","id":"ls_1","status":"completed","call_id":"call_1","action":{"type":"exec","command":["ls","-a"],"env":{"HOME":"/home/u"},"timeout_ms":1000,"user":"u","working_directory":"/tmp"}}), + ResponsesOutputItem::LocalShellCall(ResponsesLocalShellCall { + id: text("ls_1"), status: text("completed"), call_id: text("call_1"), + action: Some(ResponsesLocalShellAction::Exec { + command: vec!["ls".into(), "-a".into()], env: [("HOME".into(), "/home/u".into())].into(), + timeout_ms: Some(1000), user: text("u"), working_directory: text("/tmp"), extra: Default::default(), + }), + ..Default::default() + }) +)] +#[case::local_shell_call_output( + json!({"type":"local_shell_call_output","id":"lso_1","status":"completed","output":"{\"stdout\":\"x\"}"}), + ResponsesOutputItem::LocalShellCallOutput(ResponsesLocalShellCallOutput { + id: text("lso_1"), status: text("completed"), output: text("{\"stdout\":\"x\"}"), ..Default::default() + }) +)] +#[case::shell_call_local( + json!({"type":"shell_call","id":"sh_1","status":"in_progress","call_id":"call_1","action":{"commands":["ls"],"max_output_length":100,"timeout_ms":500},"environment":{"type":"local"},"caller":{"type":"direct"},"created_by":"user_1"}), + ResponsesOutputItem::ShellCall(ResponsesShellCall { + id: text("sh_1"), status: text("in_progress"), call_id: text("call_1"), + action: Some(ResponsesShellAction { commands: Some(vec!["ls".into()]), max_output_length: Some(100), timeout_ms: Some(500), ..Default::default() }), + environment: Some(ResponsesShellEnvironment::Local { extra: Default::default() }), + caller: direct_caller(), created_by: text("user_1"), ..Default::default() + }) +)] +#[case::shell_call_container( + json!({"type":"shell_call","call_id":"call_1","environment":{"type":"container_reference","container_id":"cntr_1"}}), + ResponsesOutputItem::ShellCall(ResponsesShellCall { + call_id: text("call_1"), + environment: Some(ResponsesShellEnvironment::ContainerReference { container_id: "cntr_1".into(), extra: Default::default() }), + ..Default::default() + }) +)] +#[case::shell_call_output( + json!({"type":"shell_call_output","id":"sho_1","status":"completed","call_id":"call_1","max_output_length":100,"output":[ + {"outcome":{"type":"exit","exit_code":2},"stdout":"out","stderr":"err","created_by":"user_1"}, + {"outcome":{"type":"timeout"},"stdout":"","stderr":""} + ],"caller":{"type":"program","caller_id":"prog_1"},"created_by":"user_1"}), + ResponsesOutputItem::ShellCallOutput(ResponsesShellCallOutput { + id: text("sho_1"), status: text("completed"), call_id: text("call_1"), max_output_length: Some(100), + output: Some(vec![ + ResponsesShellOutputChunk { + outcome: Some(ResponsesShellOutcome::Exit { exit_code: 2, extra: Default::default() }), + stdout: text("out"), stderr: text("err"), created_by: text("user_1"), ..Default::default() + }, + ResponsesShellOutputChunk { + outcome: Some(ResponsesShellOutcome::Timeout { extra: Default::default() }), + stdout: text(""), stderr: text(""), ..Default::default() + }, + ]), + caller: program_caller(), created_by: text("user_1"), ..Default::default() + }) +)] +#[case::apply_patch_create( + json!({"type":"apply_patch_call","id":"ap_1","status":"completed","call_id":"call_1","operation":{"type":"create_file","path":"a.txt","diff":"+a"},"caller":{"type":"direct"},"created_by":"user_1"}), + ResponsesOutputItem::ApplyPatchCall(ResponsesApplyPatchCall { + id: text("ap_1"), status: text("completed"), call_id: text("call_1"), + operation: Some(ResponsesApplyPatchOperation::CreateFile { path: "a.txt".into(), diff: "+a".into(), extra: Default::default() }), + caller: direct_caller(), created_by: text("user_1"), ..Default::default() + }) +)] +#[case::apply_patch_delete( + json!({"type":"apply_patch_call","call_id":"call_1","operation":{"type":"delete_file","path":"a.txt"}}), + ResponsesOutputItem::ApplyPatchCall(ResponsesApplyPatchCall { + call_id: text("call_1"), + operation: Some(ResponsesApplyPatchOperation::DeleteFile { path: "a.txt".into(), extra: Default::default() }), + ..Default::default() + }) +)] +#[case::apply_patch_update( + json!({"type":"apply_patch_call","call_id":"call_1","operation":{"type":"update_file","path":"a.txt","diff":"-a\n+b"}}), + ResponsesOutputItem::ApplyPatchCall(ResponsesApplyPatchCall { + call_id: text("call_1"), + operation: Some(ResponsesApplyPatchOperation::UpdateFile { path: "a.txt".into(), diff: "-a\n+b".into(), extra: Default::default() }), + ..Default::default() + }) +)] +#[case::apply_patch_call_output( + json!({"type":"apply_patch_call_output","id":"apo_1","status":"failed","call_id":"call_1","output":"conflict","caller":{"type":"direct"},"created_by":"user_1"}), + ResponsesOutputItem::ApplyPatchCallOutput(ResponsesApplyPatchCallOutput { + id: text("apo_1"), status: text("failed"), call_id: text("call_1"), output: text("conflict"), + caller: direct_caller(), created_by: text("user_1"), ..Default::default() + }) +)] +#[case::mcp_approval_request( + json!({"type":"mcp_approval_request","id":"apr_1","server_label":"s","name":"n","arguments":"{}"}), + ResponsesOutputItem::McpApprovalRequest(ResponsesMcpApprovalRequest { + id: text("apr_1"), server_label: text("s"), name: text("n"), arguments: text("{}"), ..Default::default() + }) +)] +#[case::mcp_approval_response( + json!({"type":"mcp_approval_response","id":"aprr_1","approval_request_id":"apr_1","approve":false,"reason":"denied"}), + ResponsesOutputItem::McpApprovalResponse(ResponsesMcpApprovalResponse { + id: text("aprr_1"), approval_request_id: text("apr_1"), approve: Some(false), reason: text("denied"), ..Default::default() + }) +)] +fn documented_output_items_parse_into_typed_variants( + #[case] wire: Value, + #[case] expected: ResponsesOutputItem, +) { + let item: ResponsesOutputItem = serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(item, expected); + assert_eq!(serde_json::to_value(item).unwrap(), wire); +} + +#[rstest] +#[case::compaction(json!({"type":"compaction","id":"cmp_1","encrypted_content":"opaque"}))] +#[case::approval(json!({"type":"mcp_approval_request","id":"apr_1","server_label":"s","name":"n","arguments":"{}"}))] +#[case::shell(json!({"type":"shell_call","call_id":"call_1","action":{"commands":["ls"]}}))] +fn event_response_types_documented_output_items(#[case] item: Value) { + let wire = json!({"output":[item.clone()]}); + let response: ResponsesEventResponse = serde_json::from_value(wire.clone()).unwrap(); + let Some([Recognized::Known(known)]) = response.output.as_deref() else { + panic!("expected one typed item"); + }; + assert_eq!( + known, + &serde_json::from_value::(item).unwrap() + ); + assert_eq!(serde_json::to_value(response).unwrap(), wire); +} + +#[rstest] +#[case::unknown_caller(json!({"type":"function_call","caller":{"type":"future"}}))] +#[case::program_caller_missing_id(json!({"type":"shell_call","caller":{"type":"program"}}))] +#[case::untagged_caller(json!({"type":"custom_tool_call","caller":{}}))] +#[case::async_not_bool(json!({"type":"function_call","async":"yes"}))] +#[case::output_not_text_or_list(json!({"type":"function_call_output","output":{"text":"t"}}))] +#[case::unknown_input_content(json!({"type":"function_call_output","output":[{"type":"input_audio"}]}))] +#[case::input_text_missing_text(json!({"type":"custom_tool_call_output","output":[{"type":"input_text"}]}))] +#[case::input_image_missing_detail(json!({"type":"custom_tool_call_output","output":[{"type":"input_image","file_id":"f"}]}))] +#[case::bad_cache_mode(json!({"type":"custom_tool_call_output","output":[{"type":"input_text","text":"t","prompt_cache_breakpoint":{"mode":"implicit"}}]}))] +#[case::unknown_computer_action(json!({"type":"computer_call","action":{"type":"teleport"}}))] +#[case::click_missing_button(json!({"type":"computer_call","action":{"type":"click","x":1,"y":2}}))] +#[case::click_fractional_x(json!({"type":"computer_call","action":{"type":"click","button":"left","x":1.5,"y":2}}))] +#[case::drag_point_missing_y(json!({"type":"computer_call","actions":[{"type":"drag","path":[{"x":1}]}]}))] +#[case::keypress_keys_not_list(json!({"type":"computer_call","actions":[{"type":"keypress","keys":"enter"}]}))] +#[case::safety_check_not_object(json!({"type":"computer_call","pending_safety_checks":["sc_1"]}))] +#[case::unknown_screenshot_tag(json!({"type":"computer_call_output","output":{"type":"screenshot"}}))] +#[case::untagged_screenshot(json!({"type":"computer_call_output","output":{"file_id":"f"}}))] +#[case::tool_not_object(json!({"type":"tool_search_output","tools":["lookup"]}))] +#[case::unknown_local_shell_action(json!({"type":"local_shell_call","action":{"type":"spawn","command":[],"env":{}}}))] +#[case::local_shell_missing_env(json!({"type":"local_shell_call","action":{"type":"exec","command":["ls"]}}))] +#[case::local_shell_env_value(json!({"type":"local_shell_call","action":{"type":"exec","command":["ls"],"env":{"A":1}}}))] +#[case::shell_commands_not_list(json!({"type":"shell_call","action":{"commands":"ls"}}))] +#[case::unknown_shell_environment(json!({"type":"shell_call","environment":{"type":"container_auto"}}))] +#[case::container_missing_id(json!({"type":"shell_call","environment":{"type":"container_reference"}}))] +#[case::unknown_shell_outcome(json!({"type":"shell_call_output","output":[{"outcome":{"type":"killed"}}]}))] +#[case::exit_missing_code(json!({"type":"shell_call_output","output":[{"outcome":{"type":"exit"}}]}))] +#[case::negative_max_output(json!({"type":"shell_call_output","max_output_length":-1}))] +#[case::unknown_patch_operation(json!({"type":"apply_patch_call","operation":{"type":"rename_file","path":"a"}}))] +#[case::update_missing_diff(json!({"type":"apply_patch_call","operation":{"type":"update_file","path":"a"}}))] +#[case::approve_not_bool(json!({"type":"mcp_approval_response","approve":"true"}))] +#[case::compaction_content_not_string(json!({"type":"compaction","encrypted_content":7}))] +#[case::program_code_not_string(json!({"type":"program","code":["run()"]}))] +fn output_items_reject_malformed_typed_fields(#[case] wire: Value) { + assert!(serde_json::from_value::(wire).is_err()); +} diff --git a/litellm-rust/crates/llms/src/anthropic/beta_headers.rs b/litellm-rust/crates/llms/src/anthropic/beta_headers.rs new file mode 100644 index 00000000000..3b83b12d727 --- /dev/null +++ b/litellm-rust/crates/llms/src/anthropic/beta_headers.rs @@ -0,0 +1,36 @@ +use litellm_http::request::{with_header, without_headers}; +use litellm_llms_types::providers::anthropic::{BetaProvider, BetaSet}; + +use crate::{anthropic::common_utils::existing_betas, base_llm::auth::Headers}; + +const BETA_HEADER: &str = "anthropic-beta"; + +/// How the `anthropic-beta` header is treated before a request leaves for a host. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum BetaPolicy { + /// Send the header untouched and let the host decide. + Forward, + /// Keep only the betas the provider accepts, under the names it expects. + Filter(BetaProvider), + /// Remove every `anthropic-beta` header. + Drop, +} + +impl BetaPolicy { + pub fn apply(self, headers: Headers) -> Headers { + let provider = match self { + Self::Forward => return headers, + Self::Drop => return without_headers(headers, &[BETA_HEADER]), + Self::Filter(provider) => provider, + }; + let accepted: BetaSet = existing_betas(&headers) + .iter() + .filter_map(|beta| beta.on(provider)) + .collect(); + let headers = without_headers(headers, &[BETA_HEADER]); + if accepted.is_empty() { + return headers; + } + with_header(headers, BETA_HEADER, accepted.to_string()) + } +} diff --git a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs index ecc6cfaf83d..e1fd1956167 100644 --- a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs @@ -241,7 +241,7 @@ fn anthropic_body( .iter() .map(|turn| { json!({ - "role": turn.role.as_str(), + "role": <&'static str>::from(turn.role), "content": turn.texts.iter().map(|text| text_block(text)).collect::>(), }) }) diff --git a/litellm-rust/crates/llms/src/anthropic/common_utils.rs b/litellm-rust/crates/llms/src/anthropic/common_utils.rs index 6e2fa3b8785..ccf3a6d15f2 100644 --- a/litellm-rust/crates/llms/src/anthropic/common_utils.rs +++ b/litellm-rust/crates/llms/src/anthropic/common_utils.rs @@ -596,17 +596,10 @@ mod tests { use crate::base_llm::messages::context::SupportedEffortTiers; use rstest::{fixture, rstest}; use serde_json::json; + use strum::VariantArray; use super::*; - const ALL_LEVELS: [EffortLevel; 5] = [ - EffortLevel::Low, - EffortLevel::Medium, - EffortLevel::High, - EffortLevel::Xhigh, - EffortLevel::Max, - ]; - fn apply(sanitizer: fn(Vec) -> Vec, messages: Value) -> Value { let parsed: Vec = serde_json::from_value(messages).unwrap(); serde_json::to_value(sanitizer(parsed)).unwrap() @@ -1712,7 +1705,10 @@ mod tests { ..unmapped }; assert_eq!( - ALL_LEVELS.map(|level| supports_effort_tier(&capabilities, level)), + EffortLevel::VARIANTS + .iter() + .map(|level| supports_effort_tier(&capabilities, *level)) + .collect::>(), expected ); } diff --git a/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs b/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs index 2162c39c229..df1ae13fd28 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs @@ -10,6 +10,7 @@ use litellm_llms_types::{ }; use litellm_python_compat::{json::from_json, repr::repr, truthy::truthy}; use serde_json::Value; +use strum::VariantArray; use crate::base_llm::messages::context::{ MessagesModelCapabilities, ThinkingBudgets, ThinkingContext, @@ -54,14 +55,17 @@ fn unmapped_effort(effort: &Value) -> Error { Error::InvalidRequest(crate::ErrorDetail::InvalidChoice { field: "reasoning effort", actual: repr(&from_json(effort.clone())), - choices: ReasoningEffort::ALL.map(|effort| effort.as_str()).into(), + choices: ReasoningEffort::VARIANTS + .iter() + .map(|effort| <&'static str>::from(*effort)) + .collect(), }) } fn unsupported_effort(level: EffortLevel, model: &str) -> Error { Error::InvalidRequest(crate::ErrorDetail::UnsupportedValue { field: "effort", - value: level.as_str(), + value: <&'static str>::from(level), model: model.into(), }) } @@ -134,7 +138,7 @@ fn legacy_reasoning_effort( Some(Recognized::Known(level)) => Ok((*level).into()), Some(Recognized::Unrecognized(value)) if truthy(&from_json(value.clone())) => value .as_str() - .and_then(ReasoningEffort::parse) + .and_then(|text| text.parse().ok()) .ok_or_else(|| unmapped_effort(value)), None | Some(Recognized::Unrecognized(_)) => Ok(ReasoningEffort::Medium), } diff --git a/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs index b9cf6c37272..74ac35752ab 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs @@ -270,7 +270,7 @@ fn drop_unsupported_params( fn speed_text(speed: &Recognized) -> String { match speed { - Recognized::Known(speed) => speed.as_str().to_string(), + Recognized::Known(speed) => <&'static str>::from(*speed).to_string(), Recognized::Unrecognized(Value::String(text)) => text.clone(), Recognized::Unrecognized(other) => other.to_string(), } diff --git a/litellm-rust/crates/llms/src/anthropic/mod.rs b/litellm-rust/crates/llms/src/anthropic/mod.rs index a884c146dca..72c390cfe21 100644 --- a/litellm-rust/crates/llms/src/anthropic/mod.rs +++ b/litellm-rust/crates/llms/src/anthropic/mod.rs @@ -1,3 +1,4 @@ +pub mod beta_headers; pub mod common_utils; pub mod batches; diff --git a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs index 90f5b97322f..4139ee860cf 100644 --- a/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs +++ b/litellm-rust/crates/llms/src/aws_textract/ocr/common_utils.rs @@ -23,9 +23,11 @@ const HEALTH_CHECK_IMAGE_DATA_URI: &str = "data:image/png;base64,iVBORw0KGgoAAAA /// Textract has operations rather than models; the model slot of /// `aws_textract/` names the one to call. #[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, VariantNames, PartialEq, Eq)] -#[strum(serialize_all = "kebab-case", ascii_case_insensitive)] +#[strum(ascii_case_insensitive)] pub enum TextractOperation { + #[strum(serialize = "detect-document-text")] DetectDocumentText, + #[strum(serialize = "analyze-document")] AnalyzeDocument, } diff --git a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs index cee906e77a4..c0b2c7aa3b7 100644 --- a/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs +++ b/litellm-rust/crates/llms/src/bedrock/audio_transcription/mod.rs @@ -30,20 +30,17 @@ pub struct BedrockAudioTranscriptionConfig; #[derive(Clone, Copy, Deserialize, IntoStaticStr)] #[serde(rename_all = "lowercase")] -#[strum(serialize_all = "lowercase")] enum AudioFormat { + #[strum(serialize = "wav")] Wav, + #[strum(serialize = "mp3")] Mp3, + #[strum(serialize = "flac")] Flac, + #[strum(serialize = "ogg")] Ogg, } -impl AudioFormat { - fn as_str(self) -> &'static str { - self.into() - } -} - struct AudioInput { data: String, format: AudioFormat, @@ -118,7 +115,7 @@ impl BaseAudioTranscriptionConfig for BedrockAudioTranscriptionConfig { "messages": [{ "role": "user", "content": [ - {"audio": {"format": audio.format.as_str(), "source": {"bytes": audio.data}}}, + {"audio": {"format": <&'static str>::from(audio.format), "source": {"bytes": audio.data}}}, {"text": instruction} ] }], diff --git a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs index 3a13e388a4b..bcd1077b830 100644 --- a/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs +++ b/litellm-rust/crates/llms/src/bedrock/chat/converse_transformation.rs @@ -385,7 +385,7 @@ fn converse_body(conversation: &Conversation, optional_params: &Map::from(turn.role), "content": turn.texts.iter().map(|text| json!({"text": text})).collect::>(), }) }) diff --git a/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs b/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs index ad7025ab6af..373ad4ca4f5 100644 --- a/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs +++ b/litellm-rust/crates/llms/src/bedrock/chat/invoke_handler.rs @@ -64,27 +64,23 @@ pub fn invoke_anthropic_event_stream(bytes: ByteStream) -> EventStream { Box::pin(invoke_chunk_stream(bytes).map(|chunk| decode_invoke_anthropic_chunk(chunk?))) } -#[derive(Clone, Copy)] +#[derive(Clone, Copy, Debug, PartialEq, Eq, strum::EnumString)] enum InvokeProvider { + #[strum(serialize = "anthropic")] Anthropic, + #[strum(serialize = "deepseek_r1")] DeepseekR1, + #[strum(serialize = "moonshot")] Moonshot, + #[strum(disabled)] Unsupported, } -impl From<&str> for InvokeProvider { - fn from(value: &str) -> Self { - match value { - "anthropic" => Self::Anthropic, - "deepseek_r1" => Self::DeepseekR1, - "moonshot" => Self::Moonshot, - _ => Self::Unsupported, - } - } -} - pub fn invoke_chat_stream(invoke_provider: &str, shape: StreamShape) -> Result { - match InvokeProvider::from(invoke_provider) { + match invoke_provider + .parse() + .unwrap_or(InvokeProvider::Unsupported) + { InvokeProvider::Anthropic => Ok(ChatStream::new( invoke_anthropic_event_stream, ModelResponseIterator::new(shape), @@ -116,6 +112,22 @@ mod tests { const TEXT_DELTA: &str = r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hello"}}"#; + #[rstest::rstest] + #[case::anthropic("anthropic", InvokeProvider::Anthropic)] + #[case::deepseek_r1("deepseek_r1", InvokeProvider::DeepseekR1)] + #[case::moonshot("moonshot", InvokeProvider::Moonshot)] + #[case::unknown("qwen", InvokeProvider::Unsupported)] + #[case::case_sensitive("Anthropic", InvokeProvider::Unsupported)] + fn invoke_provider_maps_each_model_family_name( + #[case] name: &str, + #[case] expected: InvokeProvider, + ) { + assert_eq!( + name.parse().unwrap_or(InvokeProvider::Unsupported), + expected + ); + } + fn aws_wire(chunk: &str) -> Vec { let payload = serde_json::json!({"bytes": STANDARD.encode(chunk)}); let message = Message::new(Bytes::from(serde_json::to_vec(&payload).unwrap())).add_header( diff --git a/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs b/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs index 67b81ec6feb..dadab413fe4 100644 --- a/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs @@ -25,11 +25,13 @@ const COHERE_API_KEY_ENV: &str = "COHERE_API_KEY"; const COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI: &str = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"; -#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)] +#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize, strum::IntoStaticStr)] #[serde(rename_all = "lowercase")] pub enum OutputFormat { #[default] + #[strum(serialize = "markdown")] Markdown, + #[strum(serialize = "blocks")] Blocks, } @@ -267,11 +269,7 @@ fn build_request(model: &str, image_url: String, params: &CohereOptions) -> Cohe CohereRequest { model: model.into(), document: CohereParseDocument::ImageUrl { image_url }, - output_format: match params.output_format.unwrap_or_default() { - OutputFormat::Markdown => "markdown", - OutputFormat::Blocks => "blocks", - } - .into(), + output_format: <&'static str>::from(params.output_format.unwrap_or_default()).into(), } } diff --git a/litellm-rust/crates/llms/tests/anthropic_beta_headers.rs b/litellm-rust/crates/llms/tests/anthropic_beta_headers.rs new file mode 100644 index 00000000000..422c3d4b5f0 --- /dev/null +++ b/litellm-rust/crates/llms/tests/anthropic_beta_headers.rs @@ -0,0 +1,111 @@ +use litellm_llms::anthropic::beta_headers::BetaPolicy; +use litellm_llms_types::providers::anthropic::BetaProvider; +use rstest::rstest; + +fn headers(pairs: &[(&str, &str)]) -> Vec<(String, String)> { + pairs + .iter() + .map(|(name, value)| (name.to_string(), value.to_string())) + .collect() +} + +#[rstest] +fn forward_keeps_headers_as_is() { + let input = headers(&[ + ( + "Anthropic-Beta", + "example-beta-2099-01-01,web-search-2025-03-05", + ), + ("x-api-key", "k"), + ]); + assert_eq!(BetaPolicy::Forward.apply(input.clone()), input); +} + +#[rstest] +fn drop_removes_every_casing() { + assert_eq!( + BetaPolicy::Drop.apply(headers(&[ + ("anthropic-beta", "a"), + ("x-api-key", "k"), + ("ANTHROPIC-BETA", "b"), + ])), + headers(&[("x-api-key", "k")]) + ); +} + +#[rstest] +#[case::bedrock(BetaProvider::Bedrock, Some("tool-search-tool-2025-10-19"))] +#[case::bedrock_mantle(BetaProvider::BedrockMantle, Some("tool-search-tool-2025-10-19"))] +#[case::vertex_ai(BetaProvider::VertexAi, Some("tool-search-tool-2025-10-19"))] +#[case::anthropic(BetaProvider::Anthropic, Some("advanced-tool-use-2025-11-20"))] +#[case::azure_ai(BetaProvider::AzureAi, Some("advanced-tool-use-2025-11-20"))] +#[case::databricks(BetaProvider::Databricks, Some("advanced-tool-use-2025-11-20"))] +#[case::bedrock_converse(BetaProvider::BedrockConverse, None)] +fn filter_renames_per_host(#[case] provider: BetaProvider, #[case] expected_beta: Option<&str>) { + let expected = match expected_beta { + Some(beta) => headers(&[("x-api-key", "k"), ("anthropic-beta", beta)]), + None => headers(&[("x-api-key", "k")]), + }; + assert_eq!( + BetaPolicy::Filter(provider).apply(headers(&[ + ("Anthropic-Beta", "advanced-tool-use-2025-11-20"), + ("x-api-key", "k"), + ])), + expected + ); +} + +#[rstest] +#[case::azure_ai(BetaProvider::AzureAi, "fast-mode-2026-02-01,example-beta-2099-01-01")] +#[case::anthropic(BetaProvider::Anthropic, "bash_20241022")] +#[case::vertex_ai(BetaProvider::VertexAi, " , ")] +fn filter_removes_header_when_nothing_survives( + #[case] provider: BetaProvider, + #[case] beta_value: &str, +) { + assert_eq!( + BetaPolicy::Filter(provider).apply(headers(&[ + ("x-api-key", "k"), + ("anthropic-beta", beta_value) + ])), + headers(&[("x-api-key", "k")]) + ); +} + +#[rstest] +#[case::anthropic( + BetaProvider::Anthropic, + "web-search-2025-03-05", + "oauth-2025-04-20,web-search-2025-03-05,example-beta-2099-01-01", + "oauth-2025-04-20,web-search-2025-03-05" +)] +#[case::vertex_ai_renames_and_deduplicates( + BetaProvider::VertexAi, + "advanced-tool-use-2025-11-20", + "tool-search-tool-2025-10-19", + "tool-search-tool-2025-10-19" +)] +fn filter_deduplicates_and_sorts( + #[case] provider: BetaProvider, + #[case] first_beta: &str, + #[case] second_beta: &str, + #[case] expected_beta: &str, +) { + assert_eq!( + BetaPolicy::Filter(provider).apply(headers(&[ + ("ANTHROPIC-BETA", first_beta), + ("x-api-key", "k"), + ("anthropic-beta", second_beta), + ])), + headers(&[("x-api-key", "k"), ("anthropic-beta", expected_beta)]) + ); +} + +#[rstest] +#[case::forward(BetaPolicy::Forward)] +#[case::drop(BetaPolicy::Drop)] +#[case::filter(BetaPolicy::Filter(BetaProvider::Bedrock))] +fn headers_without_a_beta_are_untouched(#[case] policy: BetaPolicy) { + let input = headers(&[("x-api-key", "k"), ("anthropic-version", "2023")]); + assert_eq!(policy.apply(input.clone()), input); +} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index cf1080a3881..92acf796e6b 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -17,6 +17,10 @@ mod tokenizer; #[pymodule(gil_used = true)] mod _native { + #[pymodule_export] + use litellm_host_python::{ForkedAfterNativeRuntimeStarted, ProcessReservedForForking}; + use pyo3::{prelude::*, types::PyModule}; + #[cfg(feature = "panic-test")] #[pymodule_export] use crate::diagnostics::_panic_for_test; @@ -29,9 +33,7 @@ mod _native { #[pymodule_export] use crate::routes::audio_transcription::{atranscription, transcription}; #[pymodule_export] - use crate::routes::chat_completions::{ - achat_completions, acompletion, chat_completions, completion, - }; + use crate::routes::chat_completions::{acompletion, completion}; #[pymodule_export] use crate::routes::embeddings::{aembedding, embedding}; #[pymodule_export] @@ -51,9 +53,6 @@ mod _native { use crate::tokenizer::HuggingFaceEncoding; #[pymodule_export] use crate::tokenizer::Tokenizer; - #[pymodule_export] - use litellm_host_python::{ForkedAfterNativeRuntimeStarted, ProcessReservedForForking}; - use pyo3::{prelude::*, types::PyModule}; #[pymodule_init] fn init(module: &Bound<'_, PyModule>) -> PyResult<()> { @@ -101,8 +100,6 @@ mod tests { "atranscription", "messages", "amessages", - "chat_completions", - "achat_completions", "completion", "acompletion", "responses", diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index df15a9591a1..4a7efeffdd2 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -19,13 +19,6 @@ pub(crate) struct RouteOptions { pub(crate) timeout: Option, } -pub(crate) fn messages_argument(value: &Bound<'_, PyAny>) -> PyResult> { - match from_py_argument(value)? { - Value::Array(values) => Ok(values), - _ => Err(PyValueError::new_err("messages must be a list")), - } -} - fn required_object(name: &'static str, value: Value) -> PyResult> { match value { Value::Object(values) => Ok(values), @@ -449,18 +442,6 @@ module.__getattr__ = fail fn argument_converters_keep_nested_values_and_accept_explicit_none() { Python::initialize(); Python::attach(|py| { - let messages = py - .eval( - c"[{'role': 'user', 'content': [{'type': 'text', 'text': 'hi'}]}]", - None, - None, - ) - .unwrap(); - assert_eq!( - Value::Array(messages_argument(&messages).unwrap()), - json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]) - ); - let params = py.eval(c"{'temperature': 0.2}", None, None).unwrap(); assert_eq!( optional_object("optional_params", ¶ms).unwrap(), diff --git a/litellm-rust/crates/python-bridge/src/python_settings.rs b/litellm-rust/crates/python-bridge/src/python_settings.rs index 8f49444f730..99afc772b99 100644 --- a/litellm-rust/crates/python-bridge/src/python_settings.rs +++ b/litellm-rust/crates/python-bridge/src/python_settings.rs @@ -6,13 +6,16 @@ use crate::coercion::{FieldSpec, ProjectionError}; const MODULE: &str = "litellm.rust_bridge.settings"; #[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq)] -#[strum(serialize_all = "snake_case")] pub(crate) enum PythonSettings { #[strum(serialize = "http_settings")] Http, + #[strum(serialize = "url_policy")] UrlPolicy, + #[strum(serialize = "provider_defaults")] ProviderDefaults, + #[strum(serialize = "secret_manager")] SecretManager, + #[strum(serialize = "secret_manager_binding")] SecretManagerBinding, } @@ -23,17 +26,16 @@ pub(crate) struct Snapshot<'py> { impl Snapshot<'_> { pub(crate) fn read(&self, spec: &FieldSpec) -> Result { - spec.read(&self.value, self.group.name()) + spec.read(&self.value, self.group.into()) } } impl PythonSettings { - pub(crate) fn name(self) -> &'static str { - self.into() - } - pub(crate) fn read(self, py: Python<'_>) -> PyResult> { - let value = py.import(MODULE)?.getattr(self.name())?.call0()?; + let value = py + .import(MODULE)? + .getattr(<&'static str>::from(self))? + .call0()?; Ok(Snapshot { group: self, value }) } diff --git a/litellm-rust/crates/python-bridge/src/routes/AGENTS.md b/litellm-rust/crates/python-bridge/src/routes/AGENTS.md index c2afed49b45..948c67eb647 100644 --- a/litellm-rust/crates/python-bridge/src/routes/AGENTS.md +++ b/litellm-rust/crates/python-bridge/src/routes/AGENTS.md @@ -11,3 +11,7 @@ The host driver owns sequencing and terminal events; the bridge supplies fallibl Use the shared `run_public_call` boundary with hooks supplied by bridge composition. `callbacks-legacy-python` owns legacy argument sharing and `Logging` dispatch behind `PublicCall` and `LegacyLogging`. Route bindings supply `callbacks-legacy-python::LoggingOperation` when composing legacy logging and may retain the request needed for projection, but must not duplicate the legacy callback contract Regression tests must observe that an unstarted call does no setup, hook and preflight rewrites affect resource configuration, setup failures reach the selected failure handler once, and provider work is not replayed. Retain existing read-point and object-identity guarantees while changing setup timing + +## Layout + +`messages/` is the reference shape. Each lifecycle route is a folder: `mod.rs` holds `run_` and the thin sync and async `#[pyfunction]` entrypoints that take `NativeCall` and call it, `host.rs` holds the `PythonBinding` host, and anything else route-specific (projection, error mapping, extra pyclasses) gets its own sibling file. Scaffolded routes without a lifecycle (embeddings, audio transcription) stay a single `.rs` diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs index fc7bfb08b75..e62673d4d4f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -5,6 +5,8 @@ use litellm_inference_transcription::{ use pyo3::prelude::*; use serde_json::{Map, Value}; +use super::NativeCall; + use crate::{ errors::route_error_to_pyerr, marshal::{RouteOptions, optional_object_field, required_field, value_route_options}, @@ -40,13 +42,12 @@ async fn execute( } #[pyfunction] -pub(crate) fn transcription(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; +pub(crate) fn transcription(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + let arguments = call.resolved()?; let audio: Value = - litellm_host_python::from_py_argument(&required_field(&call.bound, "audio")?)?; - let options = value_route_options(&call.bound)?; - let optional_params = - optional_object_field(&call.bound, "optional_params")?.unwrap_or_default(); + litellm_host_python::from_py_argument(&required_field(&arguments, "audio")?)?; + let options = value_route_options(&arguments)?; + let optional_params = optional_object_field(&arguments, "optional_params")?.unwrap_or_default(); let http = crate::http::provider_client(py, &call.kwargs, false)?; let secrets = crate::secrets::source(py)?; run_sync( @@ -59,14 +60,13 @@ pub(crate) fn transcription(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult< #[pyfunction] pub(crate) fn atranscription<'py>( py: Python<'py>, - call: Bound<'py, PyAny>, + call: NativeCall<'py>, ) -> PyResult> { - let call = super::NativeCall::extract(&call)?; + let arguments = call.resolved()?; let audio: Value = - litellm_host_python::from_py_argument(&required_field(&call.bound, "audio")?)?; - let options = value_route_options(&call.bound)?; - let optional_params = - optional_object_field(&call.bound, "optional_params")?.unwrap_or_default(); + litellm_host_python::from_py_argument(&required_field(&arguments, "audio")?)?; + let options = value_route_options(&arguments)?; + let optional_params = optional_object_field(&arguments, "optional_params")?.unwrap_or_default(); let http = crate::http::provider_client(py, &call.kwargs, true)?; let secrets = crate::secrets::source(py)?; run_async( diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs deleted file mode 100644 index 27f076b7cce..00000000000 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ /dev/null @@ -1,150 +0,0 @@ -mod host; - -use pyo3::types::{PyDict, PyTuple}; - -use crate::execution::{run_async, run_sync}; -use litellm_inference_chat::{ChatCompletionsRoute, Error, types::ChatCompletionsRequest}; -use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse; -use pyo3::prelude::*; -use serde_json::{Map, Value}; - -use crate::{ - errors::route_error_to_pyerr, - marshal::{ - RouteOptions, messages_argument, optional_object_field, required_field, value_route_options, - }, -}; - -async fn execute( - http: Result, - secrets: std::sync::Arc, - messages: Vec, - optional_params: Map, - options: RouteOptions, -) -> Result { - let RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout, - } = options; - ChatCompletionsRoute::new(http?, crate::http::resources().auth.clone(), secrets) - .execute( - ChatCompletionsRequest { - model: &model, - messages: Value::Array(messages), - optional_params, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - timeout, - }, - &(), - None, - ) - .await -} - -#[pyfunction] -pub(crate) fn chat_completions(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; - let messages: Vec = messages_argument(&required_field(&call.bound, "messages")?)?; - let optional_params = - optional_object_field(&call.bound, "optional_params")?.unwrap_or_default(); - let options = value_route_options(&call.bound)?; - let http = crate::http::provider_client(py, &call.kwargs, false)?; - let secrets = crate::secrets::source(py)?; - run_sync( - py, - execute(http, secrets, messages, optional_params, options), - route_error_to_pyerr, - ) -} - -#[pyfunction] -pub(crate) fn achat_completions<'py>( - py: Python<'py>, - call: Bound<'py, PyAny>, -) -> PyResult> { - let call = super::NativeCall::extract(&call)?; - let messages: Vec = messages_argument(&required_field(&call.bound, "messages")?)?; - let optional_params = - optional_object_field(&call.bound, "optional_params")?.unwrap_or_default(); - let options = value_route_options(&call.bound)?; - let http = crate::http::provider_client(py, &call.kwargs, true)?; - let secrets = crate::secrets::source(py)?; - run_async( - py, - execute(http, secrets, messages, optional_params, options), - route_error_to_pyerr, - ) -} - -fn run_public( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, - asynchronous: bool, -) -> PyResult> { - use super::inference::InferenceHost; - use litellm_callbacks_legacy_python::LoggingOperation; - let host = InferenceHost::new( - request.clone().unbind(), - "litellm.rust_bridge.chat_completions.route_host", - ); - let cache_call_type = if asynchronous { - "acompletion" - } else { - "completion" - }; - crate::cache::admit_native(py, &kwargs, cache_call_type)?; - let (arguments, hooks) = crate::routes::call_hooks( - py, - LoggingOperation::Completion, - &request, - &args, - &kwargs, - asynchronous, - )?; - crate::routes::run_public_call( - py, - arguments, - move |py, arguments, request| { - let route = ChatCompletionsRoute::new( - crate::http::provider_client(py, arguments, asynchronous)? - .map_err(crate::http::client_error)?, - crate::http::resources().auth.clone(), - crate::secrets::source(py)?, - ); - let (cache, cache_options) = - crate::cache::configured_native(py, arguments, cache_call_type)?; - let route = match cache { - Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new( - cache, - litellm_cache_response::CacheScope::Shared, - )), - None => route, - }; - Ok(route.machine(request, cache_options.policy)) - }, - host::ChatCompletionsPythonHost(host), - hooks, - asynchronous, - ) -} - -#[pyfunction] -pub(crate) fn completion(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; - run_public(py, call.bound.into_any(), call.args, call.kwargs, false) -} - -#[pyfunction] -pub(crate) fn acompletion(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; - run_public(py, call.bound.into_any(), call.args, call.kwargs, true) -} diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs index e7525cd3a91..1547a4daa93 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs @@ -1,6 +1,6 @@ use std::convert::Infallible; -use super::super::inference::InferenceHost; +use crate::routes::inference::InferenceHost; use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned}; use litellm_inference_chat::{Error, route::ChatCompletions, types::ChatCompletionsCall}; use pyo3::{ diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs new file mode 100644 index 00000000000..2b69d0e6321 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs @@ -0,0 +1,71 @@ +mod host; + +use std::sync::Arc; + +use host::ChatCompletionsPythonHost; +use litellm_auth::AuthServices; +use litellm_cache_response::{CachePolicy, ScopedCache}; +use litellm_callbacks_legacy_python::LoggingOperation; +use litellm_host::{call::HostedMachine, protocol::Protocol}; +use litellm_inference_chat::{ChatCompletionsRoute, route::ChatCompletions}; +use litellm_secrets::source::SecretSource; +use pyo3::prelude::*; + +use super::{ + NativeCall, + inference::{InferenceHost, InferenceRoute, run_inference}, +}; + +fn run_chat_completions( + py: Python<'_>, + call: NativeCall<'_>, + asynchronous: bool, +) -> PyResult> { + let host = InferenceHost::new( + call.resolved()?.unbind(), + "litellm.rust_bridge.chat_completions.route_host", + ); + run_inference::( + py, + call, + asynchronous, + ChatCompletionsPythonHost(host), + ) +} + +#[pyfunction] +pub(crate) fn completion(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_chat_completions(py, call, false) +} + +#[pyfunction] +pub(crate) fn acompletion(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_chat_completions(py, call, true) +} + +impl InferenceRoute for ChatCompletionsRoute { + type Protocol = ChatCompletions; + const OPERATION: LoggingOperation = LoggingOperation::Completion; + const SYNC_CALL_TYPE: &'static str = "completion"; + const ASYNC_CALL_TYPE: &'static str = "acompletion"; + + fn new( + http: litellm_http::Client, + auth: Arc, + secrets: Arc, + ) -> Self { + Self::new(http, auth, secrets) + } + + fn with_cache(self, cache: ScopedCache) -> Self { + self.with_cache(cache) + } + + fn machine( + self, + call: ::Request, + policy: CachePolicy, + ) -> HostedMachine { + self.machine(call, policy) + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/embeddings.rs b/litellm-rust/crates/python-bridge/src/routes/embeddings.rs index 1d34e3c21ff..0f73eefedd5 100644 --- a/litellm-rust/crates/python-bridge/src/routes/embeddings.rs +++ b/litellm-rust/crates/python-bridge/src/routes/embeddings.rs @@ -1,17 +1,18 @@ use pyo3::prelude::*; +use super::NativeCall; use crate::errors::RustBridgeDeclined; #[pyfunction] -pub(crate) fn embedding(call: Bound<'_, PyAny>) -> PyResult> { - drop(super::NativeCall::extract(&call)?); +pub(crate) fn embedding(call: NativeCall<'_>) -> PyResult> { + drop(call); Err(RustBridgeDeclined::new_err( "native embeddings route is not implemented", )) } #[pyfunction] -pub(crate) fn aembedding(call: Bound<'_, PyAny>) -> PyResult> { +pub(crate) fn aembedding(call: NativeCall<'_>) -> PyResult> { embedding(call) } @@ -31,12 +32,13 @@ mod tests { let locals = PyDict::new(py); py.run( c"from types import SimpleNamespace -call = SimpleNamespace(args=(), kwargs={}, bound={'model':'test-model','input':'hello'})", +call = SimpleNamespace(args=(), kwargs={'model':'test-model','input':'hello'}, base={})", Some(&locals), Some(&locals), ) .unwrap(); - let call = locals.get_item("call").unwrap().unwrap(); + let call: super::NativeCall<'_> = + locals.get_item("call").unwrap().unwrap().extract().unwrap(); let error = if asynchronous { super::aembedding(call) } else { diff --git a/litellm-rust/crates/python-bridge/src/routes/inference.rs b/litellm-rust/crates/python-bridge/src/routes/inference.rs index 28025b80b0f..0333aaeb04a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/inference.rs +++ b/litellm-rust/crates/python-bridge/src/routes/inference.rs @@ -1,11 +1,19 @@ +use std::sync::Arc; + +use litellm_auth::AuthServices; +use litellm_cache_response::{CachePolicy, CacheScope, ScopedCache}; +use litellm_callbacks_legacy_python::LoggingOperation; use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; -use litellm_host_python::{from_py, lookup}; +use litellm_host::{call::HostedMachine, protocol::Protocol}; +use litellm_host_python::{PythonBinding, PythonHostCalls, from_py, present}; use litellm_http::transport::Error as TransportError; use litellm_inference::RouteError; +use litellm_secrets::source::SecretSource; use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; use serde::Serialize; use serde_json::{Map, Value}; +use super::NativeCall; use crate::{ errors::{RustUpstreamError, route_error_to_pyerr}, marshal::{ @@ -15,7 +23,7 @@ use crate::{ }; pub(super) struct InferenceHost { - pub request: Py, + pub request: Py, module: &'static str, } @@ -26,7 +34,7 @@ pub(super) struct ProjectedCall { } impl InferenceHost { - pub fn new(request: Py, module: &'static str) -> Self { + pub fn new(request: Py, module: &'static str) -> Self { Self { request, module } } @@ -85,21 +93,7 @@ impl InferenceHost { arguments: &Bound<'py, PyDict>, name: &str, ) -> PyResult>> { - let request = self.request.bind(py); - if let Some(value) = lookup(arguments, request, name)? { - return Ok((!value.is_none()).then_some(value)); - } - if request.is_instance_of::() { - return Ok(None); - } - let parameter = request - .getattr("parameters")? - .call_method1("get", (name,))?; - if !parameter.is_none() { - return Ok(Some(parameter)); - } - let extra = request.getattr("kwargs")?.call_method1("get", (name,))?; - Ok((!extra.is_none()).then_some(extra)) + present(arguments, self.request.bind(py), name) } pub fn parameters( @@ -140,3 +134,62 @@ impl InferenceHost { Ok(PyErr::from_value(mapped)) } } + +pub(super) trait InferenceRoute: Sized + 'static { + type Protocol: Protocol; + const OPERATION: LoggingOperation; + const SYNC_CALL_TYPE: &'static str; + const ASYNC_CALL_TYPE: &'static str; + + fn new( + http: litellm_http::Client, + auth: Arc, + secrets: Arc, + ) -> Self; + fn with_cache(self, cache: ScopedCache) -> Self; + fn machine( + self, + call: ::Request, + policy: CachePolicy, + ) -> HostedMachine; +} + +pub(super) fn run_inference( + py: Python<'_>, + call: NativeCall<'_>, + asynchronous: bool, + host: H, +) -> PyResult> +where + R: InferenceRoute, + H: PythonBinding + PythonHostCalls + 'static, +{ + let call_type = if asynchronous { + R::ASYNC_CALL_TYPE + } else { + R::SYNC_CALL_TYPE + }; + crate::cache::admit_native(py, &call.kwargs, call_type)?; + let (arguments, hooks) = super::call_hooks(py, R::OPERATION, &call, asynchronous)?; + super::run_public_call( + py, + arguments, + move |py, arguments, request| { + let route = R::new( + crate::http::provider_client(py, arguments, asynchronous)? + .map_err(crate::http::client_error)?, + crate::http::resources().auth.clone(), + crate::secrets::source(py)?, + ); + let (cache, cache_options) = crate::cache::configured_native(py, arguments, call_type)?; + let route = match cache { + Some(cache) => route.with_cache(ScopedCache::new(cache, CacheScope::Shared)), + None => route, + }; + Ok(route.machine(request, cache_options.policy)) + }, + host, + hooks, + asynchronous, + ) +} diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index bde1cc7269c..841ef3081e3 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -2,7 +2,7 @@ use crate::cache::{CacheCall, Cached, PythonCache, Selection}; use litellm_host_python::{PythonHostCalls, PythonOwned}; use bytes::Bytes; -use litellm_host_python::{InvokeError, PythonBinding, from_py, lookup, to_py}; +use litellm_host_python::{InvokeError, PythonBinding, from_py, present, to_py}; use litellm_http::transport::Error as TransportError; use litellm_inference_messages::{ Error, MessagesCall, MessagesSettings, MessagesShaping, messages_body, @@ -89,12 +89,12 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult { /// The Python side of the Messages route: projects the prepared arguments and builds the /// public response, chunks and exceptions. pub(super) struct MessagesPythonHost { - request: Py, + request: Py, cache: PythonCache, } impl MessagesPythonHost { - pub(super) fn new(request: Py, asynchronous: bool) -> Self { + pub(super) fn new(request: Py, asynchronous: bool) -> Self { Self { request, cache: PythonCache::new(asynchronous), @@ -107,9 +107,7 @@ impl MessagesPythonHost { arguments: &Bound<'_, PyDict>, ) -> PyResult> { let request = self.request.bind(py); - let argument = |name: &str| -> PyResult>> { - Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none())) - }; + let argument = |name: &str| present(arguments, request, name); let string = |name: &str| -> PyResult> { argument(name)?.map(|value| value.extract()).transpose() }; @@ -153,8 +151,7 @@ impl MessagesPythonHost { ) -> PyResult>> { let request = self.request.bind(py); let mapping = |name: &str| -> PyResult>> { - lookup(arguments, request, name)? - .filter(|value| !value.is_none()) + present(arguments, request, name)? .map(|value| from_py(&value)) .transpose() }; @@ -169,8 +166,7 @@ impl MessagesPythonHost { py: Python<'_>, arguments: &Bound<'_, PyDict>, ) -> PyResult> { - lookup(arguments, self.request.bind(py), "provider_specific_header")? - .filter(|value| !value.is_none()) + present(arguments, self.request.bind(py), "provider_specific_header")? .map(|value| from_py(&value)) .transpose() } @@ -200,9 +196,9 @@ impl MessagesPythonHost { self.request .bind(py) .get_item("custom_llm_provider") - .and_then(|value| value.extract::>()) .ok() .flatten() + .and_then(|value| value.extract::>().ok().flatten()) .unwrap_or_else(|| "anthropic".into()) } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index 318aae27121..5f69cf078b7 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -2,26 +2,13 @@ mod host; use host::MessagesPythonHost; use litellm_callbacks_legacy_python::LoggingOperation; -use pyo3::{ - prelude::*, - types::{PyDict, PyTuple}, -}; +use pyo3::prelude::*; -fn run_messages( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, - asynchronous: bool, -) -> PyResult> { - let (arguments, hooks) = crate::routes::call_hooks( - py, - LoggingOperation::Messages, - &request, - &args, - &kwargs, - asynchronous, - )?; +use super::NativeCall; + +fn run_messages(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult> { + let (arguments, hooks) = + crate::routes::call_hooks(py, LoggingOperation::Messages, &call, asynchronous)?; crate::routes::run_public_call( py, arguments, @@ -57,20 +44,18 @@ fn run_messages( }, )) }, - MessagesPythonHost::new(request.unbind(), asynchronous), + MessagesPythonHost::new(call.resolved()?.unbind(), asynchronous), hooks, asynchronous, ) } #[pyfunction] -pub(crate) fn messages(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; - run_messages(py, call.bound.into_any(), call.args, call.kwargs, false) +pub(crate) fn messages(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_messages(py, call, false) } #[pyfunction] -pub(crate) fn amessages(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; - run_messages(py, call.bound.into_any(), call.args, call.kwargs, true) +pub(crate) fn amessages(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_messages(py, call, true) } diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index bf52fe3bfe8..803214656e9 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -8,28 +8,32 @@ pub(crate) mod responses; pub(crate) mod token_counter; pub(crate) mod traces; -use litellm_callbacks_legacy_python::LoggingOperation; -use litellm_callbacks_legacy_python::{LegacyLogging, PublicCall}; +use litellm_callbacks_legacy_python::{LegacyLogging, LoggingOperation, PublicCall}; use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol}; -use litellm_host_python::{HookChain, PythonBinding, PythonCallHooks, PythonHostCalls}; +use litellm_host_python::{ + HookChain, PythonBinding, PythonCallHooks, PythonHostCalls, effective_py_args, +}; use pyo3::{ prelude::*, types::{PyDict, PyMapping, PyTuple}, }; -struct NativeCall<'py> { +/// The public call as Python bound it: `base` holds the positional arguments by name plus +/// the signature defaults, `kwargs` the caller's keyword dict. +#[derive(FromPyObject)] +pub(crate) struct NativeCall<'py> { + #[pyo3(attribute)] args: Bound<'py, PyTuple>, + #[pyo3(attribute, from_py_with = mapping_dict)] kwargs: Bound<'py, PyDict>, - bound: Bound<'py, PyDict>, + #[pyo3(attribute, from_py_with = mapping_dict)] + base: Bound<'py, PyDict>, } impl<'py> NativeCall<'py> { - fn extract(call: &Bound<'py, PyAny>) -> PyResult { - Ok(Self { - args: call.getattr("args")?.cast_into()?, - kwargs: mapping_dict(&call.getattr("kwargs")?)?, - bound: mapping_dict(&call.getattr("bound")?)?, - }) + /// The call before any hook ran, for the reads that admit or decline it. + fn resolved(&self) -> PyResult> { + effective_py_args(&self.base, &self.kwargs) } } @@ -46,12 +50,10 @@ fn mapping_dict<'py>(value: &Bound<'py, PyAny>) -> PyResult> fn call_hooks( py: Python<'_>, operation: LoggingOperation, - request: &Bound<'_, PyAny>, - args: &Bound<'_, PyTuple>, - kwargs: &Bound<'_, PyDict>, + call: &NativeCall<'_>, asynchronous: bool, ) -> PyResult<(Py, impl PythonCallHooks + use<>)> { - let call = PublicCall::capture(request, args, kwargs)?; + let call = PublicCall::capture(&call.base, &call.args, &call.kwargs)?; let arguments = call.arguments(py); Ok(( arguments, @@ -99,6 +101,7 @@ mod tests { prelude::*, types::{PyDict, PyList}, }; + use rstest::rstest; fn value_call<'py>( py: Python<'py>, @@ -117,7 +120,7 @@ mod tests { .set_item("args", pyo3::types::PyTuple::empty(py)) .unwrap(); attributes.set_item("kwargs", &fields).unwrap(); - attributes.set_item("bound", &fields).unwrap(); + attributes.set_item("base", PyDict::new(py)).unwrap(); py.import("types") .unwrap() .getattr("SimpleNamespace") @@ -126,8 +129,10 @@ mod tests { .unwrap() } - #[test] - fn route_arguments_that_fail_to_convert_raise_value_error() { + #[rstest] + #[case::sync("transcription")] + #[case::asynchronous("atranscription")] + fn route_arguments_that_fail_to_convert_raise_value_error(#[case] name: &str) { Python::initialize(); Python::attach(|py| { let module = crate::native_module(py); @@ -151,48 +156,24 @@ value = Broken() .expect("locals should be readable") .expect("helper value should exist"); - for name in ["chat_completions", "achat_completions"] { - let error = module - .getattr(name) - .and_then(|function| { - function.call1((value_call(py, "messages", &broken, None),)) - }) - .expect_err("route should reject a value it cannot convert"); + let error = module + .getattr(name) + .and_then(|function| function.call1((value_call(py, "audio", &broken, None),))) + .expect_err("route should reject a value it cannot convert"); - assert!( - error.is_instance_of::(py), - "{name} surfaced {error} instead of ValueError" - ); - } + assert!( + error.is_instance_of::(py), + "{name} surfaced {error} instead of ValueError" + ); }); } - #[test] + #[rstest] fn sync_and_async_routes_apply_the_same_input_validation() { Python::initialize(); Python::attach(|py| { let module = crate::native_module(py); - let invalid_messages = PyDict::new(py); - let sync_chat_error = module - .getattr("chat_completions") - .and_then(|function| { - function.call1((value_call(py, "messages", &invalid_messages, None),)) - }) - .expect_err("sync chat should reject a non-list messages value"); - let async_chat_error = module - .getattr("achat_completions") - .and_then(|function| { - function.call1((value_call(py, "messages", &invalid_messages, None),)) - }) - .expect_err("async chat should reject a non-list messages value"); - - assert_eq!( - sync_chat_error.to_string(), - "ValueError: messages must be a list" - ); - assert_eq!(async_chat_error.to_string(), sync_chat_error.to_string()); - let invalid_headers = PyList::empty(py); let kwargs = PyDict::new(py); kwargs @@ -221,51 +202,43 @@ value = Broken() }); } - #[test] + #[rstest] + fn missing_and_explicit_none_optional_params_share_the_next_error() { + Python::initialize(); + Python::attach(|py| { + let module = crate::native_module(py); + let audio = PyDict::new(py); + let transcribe = |optional_params: Option>| { + let kwargs = PyDict::new(py); + if let Some(optional_params) = optional_params { + kwargs.set_item("optional_params", optional_params).unwrap(); + } + module + .getattr("transcription") + .and_then(|function| { + function.call1((value_call(py, "audio", &audio, Some(&kwargs)),)) + }) + .expect_err("model 'model' has no provider") + .to_string() + }; + + let omitted = transcribe(None); + assert_eq!(transcribe(Some(py.None().into_bound(py))), omitted); + assert_ne!(omitted, "ValueError: optional_params must be a dict"); + assert_eq!( + transcribe(Some(PyList::empty(py).into_any())), + "ValueError: optional_params must be a dict" + ); + }); + } + + #[rstest] fn route_input_validation_preserves_left_to_right_order() { Python::initialize(); Python::attach(|py| { let module = crate::native_module(py); let invalid = PyList::empty(py); - let chat_kwargs = PyDict::new(py); - chat_kwargs - .set_item("optional_params", &invalid) - .expect("kwargs should accept optional_params"); - chat_kwargs - .set_item("extra_headers", &invalid) - .expect("kwargs should accept extra_headers"); - let invalid_messages = PyDict::new(py); - let error = module - .getattr("chat_completions") - .and_then(|function| { - function.call1((value_call( - py, - "messages", - &invalid_messages, - Some(&chat_kwargs), - ),)) - }) - .expect_err("messages should be validated first"); - assert_eq!(error.to_string(), "ValueError: messages must be a list"); - - let valid_messages = PyList::empty(py); - let error = module - .getattr("chat_completions") - .and_then(|function| { - function.call1((value_call( - py, - "messages", - &valid_messages, - Some(&chat_kwargs), - ),)) - }) - .expect_err("optional_params should be validated before headers"); - assert_eq!( - error.to_string(), - "ValueError: optional_params must be a dict" - ); - let headers_kwargs = PyDict::new(py); headers_kwargs .set_item("extra_headers", &invalid) @@ -286,43 +259,4 @@ value = Broken() assert!(!error.to_string().contains("extra_headers")); }); } - - #[test] - fn missing_and_explicit_none_optional_params_share_the_next_error() { - Python::initialize(); - Python::attach(|py| { - let module = crate::native_module(py); - let messages = PyList::empty(py); - let headers = PyList::empty(py); - let omitted = PyDict::new(py); - omitted - .set_item("extra_headers", &headers) - .expect("kwargs should accept extra_headers"); - let explicit = PyDict::new(py); - explicit - .set_item("optional_params", py.None()) - .expect("kwargs should accept optional_params"); - explicit - .set_item("extra_headers", &headers) - .expect("kwargs should accept extra_headers"); - - let omitted_error = module - .getattr("chat_completions") - .and_then(|function| { - function.call1((value_call(py, "messages", &messages, Some(&omitted)),)) - }) - .expect_err("omitted optional_params should reach header validation"); - let explicit_error = module - .getattr("chat_completions") - .and_then(|function| { - function.call1((value_call(py, "messages", &messages, Some(&explicit)),)) - }) - .expect_err("None optional_params should reach header validation"); - assert_eq!( - omitted_error.to_string(), - "ValueError: extra_headers must be a dict" - ); - assert_eq!(explicit_error.to_string(), omitted_error.to_string()); - }); - } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 7ddaba70937..ca683600d71 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -27,12 +27,12 @@ enum OcrHostData { /// document as it goes), acquires Azure AD tokens, and builds the public response and /// exception. pub(super) struct OcrPythonHost { - request: Py, + request: Py, data: OcrHostData, } impl OcrPythonHost { - pub(super) fn new(request: Py) -> Self { + pub(super) fn new(request: Py) -> Self { Self { request, data: OcrHostData::Unprojected, @@ -209,7 +209,7 @@ del provider .unwrap() .cast_into::() .unwrap(); - let mut host = OcrPythonHost::new(py.None()); + let mut host = OcrPythonHost::new(PyDict::new(py).unbind()); assert!(host.decode_request(py, &kwargs).unwrap().caller_token); locals.del_item("kwargs").unwrap(); drop(kwargs); diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index 71b1054fc3c..854cdb25f7b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -9,10 +9,9 @@ use litellm_core_utils::settings::ProcessEnvironment; use litellm_host_python::to_py; use litellm_inference_ocr::provider_config; use litellm_llms::base_llm::ocr::settings::OcrSettings; -use pyo3::{ - prelude::*, - types::{PyDict, PyTuple}, -}; +use pyo3::prelude::*; + +use super::NativeCall; use crate::{ coercion::FieldSpec, @@ -29,21 +28,9 @@ const ENABLE_AZURE_AD_TOKEN_REFRESH: FieldSpec = Ok(field.exact_true()) }); -fn run_ocr( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, - asynchronous: bool, -) -> PyResult> { - let (arguments, hooks) = crate::routes::call_hooks( - py, - LoggingOperation::Ocr, - &request, - &args, - &kwargs, - asynchronous, - )?; +fn run_ocr(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult> { + let (arguments, hooks) = + crate::routes::call_hooks(py, LoggingOperation::Ocr, &call, asynchronous)?; crate::routes::run_public_call( py, arguments, @@ -61,7 +48,7 @@ fn run_ocr( let route = litellm_inference_ocr::OcrRoute::new(client); Ok(route.machine(request, None)) }, - OcrPythonHost::new(request.unbind()), + OcrPythonHost::new(call.resolved()?.unbind()), hooks, asynchronous, ) @@ -81,15 +68,13 @@ fn project_provider_defaults(snapshot: &Snapshot<'_>) -> PyResult { } #[pyfunction] -pub(crate) fn ocr(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; - run_ocr(py, call.bound.into_any(), call.args, call.kwargs, false) +pub(crate) fn ocr(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_ocr(py, call, false) } #[pyfunction] -pub(crate) fn aocr(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; - run_ocr(py, call.bound.into_any(), call.args, call.kwargs, true) +pub(crate) fn aocr(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_ocr(py, call, true) } #[pyfunction] diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index ee7f6238510..7ae5a1f3ce2 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -1,5 +1,5 @@ use litellm_auth::SecretValue; -use litellm_host_python::from_py; +use litellm_host_python::{from_py, present}; use litellm_inference_ocr::{ types::{LiteLLMOcrRequest, OcrDocumentInput}, wire::{OcrWireRequest, consumed_optional_params, decode_document, decode_request_input}, @@ -22,13 +22,13 @@ pub(super) struct OcrHostHandles { } struct OcrArguments<'a, 'py> { - request: &'a Bound<'py, PyAny>, + bound: &'a Bound<'py, PyDict>, kwargs: &'a Bound<'py, PyDict>, } impl<'py> OcrArguments<'_, 'py> { fn lookup(&self, name: &str) -> PyResult> { - litellm_host_python::lookup(self.kwargs, self.request, name)? + litellm_host_python::lookup(self.kwargs, self.bound, name)? .ok_or_else(|| PyValueError::new_err(format!("missing argument: {name}"))) } @@ -58,7 +58,7 @@ impl<'py> OcrArguments<'_, 'py> { fn extra_headers(&self) -> PyResult>> { self.lookup("extra_headers")? .extract::>>()? - .map(|value| from_py(value.bind(self.request.py()))) + .map(|value| from_py(value.bind(self.bound.py()))) .transpose() } @@ -66,7 +66,7 @@ impl<'py> OcrArguments<'_, 'py> { Ok(self .lookup("timeout")? .extract::>>()? - .map(|value| python_timeout_seconds(self.request.py(), value)) + .map(|value| python_timeout_seconds(self.bound.py(), value)) .transpose()? .flatten()) } @@ -110,10 +110,10 @@ impl ProjectedDocument { } pub(super) fn project_request( - request: &Bound<'_, PyAny>, + bound: &Bound<'_, PyDict>, kwargs: &Bound<'_, PyDict>, ) -> PyResult<(LiteLLMOcrRequest, OcrHostHandles)> { - let arguments = OcrArguments { request, kwargs }; + let arguments = OcrArguments { bound, kwargs }; let model = arguments.model()?; let custom_llm_provider = arguments.custom_llm_provider()?; let document = ProjectedDocument::project(&arguments.document()?)?; @@ -122,7 +122,7 @@ pub(super) fn project_request( .map_err(ocr_error_to_pyerr)?; let names = specs.iter().map(|spec| spec.name).collect::>(); let optional_params = - project_optional_fields(names.iter().copied(), |name| kwargs.get_item(name))?; + project_optional_fields(names.iter().copied(), |name| present(kwargs, bound, name))?; let input_sources = request_input_sources( kwargs, names @@ -136,7 +136,7 @@ pub(super) fn project_request( let timeout_seconds = arguments.timeout_seconds()?; let wire = OcrWireRequest { model, - document: document.resolve(request.py())?, + document: document.resolve(bound.py())?, api_key, api_base, custom_llm_provider, @@ -169,11 +169,20 @@ mod tests { locals } + fn dict<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyDict> { + locals + .get_item(name) + .unwrap() + .unwrap() + .cast_into::() + .unwrap() + } + fn arguments<'a, 'py>( - request: &'a Bound<'py, PyAny>, + bound: &'a Bound<'py, PyDict>, kwargs: &'a Bound<'py, PyDict>, ) -> OcrArguments<'a, 'py> { - OcrArguments { request, kwargs } + OcrArguments { bound, kwargs } } fn project_document(document: &Bound<'_, PyAny>) -> PyResult { @@ -204,141 +213,26 @@ sys.modules['litellm.rust_bridge.timeouts'] = timeouts } #[test] - fn kwargs_override_request_attributes_including_explicit_none() { + fn kwargs_override_bound_values_including_explicit_none() { Python::initialize(); Python::attach(|py| { let locals = eval( py, c" -class Request: - def __init__(self): - self.accesses = [] - def __getattribute__(self, name): - if name != 'accesses': - object.__getattribute__(self, 'accesses').append(name) - return object.__getattribute__(self, name) -request = Request() -request.model = 'from-request' -request.custom_llm_provider = 'mistral' +bound = {'model': 'from-bound', 'custom_llm_provider': 'mistral'} kwargs = {'model': 'from-kwargs', 'custom_llm_provider': None} ", ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let arguments = arguments(&request, &kwargs); + let (bound, kwargs) = (dict(&locals, "bound"), dict(&locals, "kwargs")); + let arguments = arguments(&bound, &kwargs); assert_eq!(arguments.model().unwrap(), "from-kwargs"); assert_eq!(arguments.custom_llm_provider().unwrap(), None); - let accesses: Vec = request.getattr("accesses").unwrap().extract().unwrap(); - assert_eq!(accesses, Vec::::new()); }); } - #[test] - fn missing_kwargs_read_the_request_property_once() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -class Request: - def __init__(self): - self.reads = 0 - @property - def model(self): - self.reads += 1 - return 'mistral-ocr-latest' -request = Request() -kwargs = {} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - assert_eq!( - arguments(&request, &kwargs).model().unwrap(), - "mistral-ocr-latest" - ); - assert_eq!( - request.getattr("reads").unwrap().extract::().unwrap(), - 1 - ); - }); - } - - #[test] - fn request_property_exceptions_keep_their_identity() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -failure = LookupError('model failed') -class Request: - @property - def model(self): - raise failure -request = Request() -kwargs = {} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let error = arguments(&request, &kwargs).model().unwrap_err(); - assert!( - error - .value(py) - .is(locals.get_item("failure").unwrap().unwrap()) - ); - }); - } - - #[test] - fn unused_raising_property_is_never_inspected() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -class Request: - @property - def unused(self): - raise RuntimeError('unused') - model = 'mistral-ocr-latest' - custom_llm_provider = None -request = Request() -kwargs = {} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let arguments = arguments(&request, &kwargs); - assert_eq!(arguments.model().unwrap(), "mistral-ocr-latest"); - assert_eq!(arguments.custom_llm_provider().unwrap(), None); - }); - } - - /// A reader that rewrites the request while it runs shows which arguments projection - /// read before it and which after: every other argument is read first, and the read - /// happens exactly once. + /// A reader that rewrites the bound arguments while it runs shows which arguments + /// projection read before it and which after: every other argument is read first, and + /// the read happens exactly once. #[test] fn document_readers_are_read_once_after_every_other_argument() { Python::initialize(); @@ -347,37 +241,28 @@ kwargs = {} let locals = eval( py, c" -class Request: - model = 'mistral/mistral-ocr-latest' - custom_llm_provider = None - api_key = None - api_base = 'https://original.example.com' - extra_headers = {'x-source': 'original'} - timeout = 1 - @property - def document(self): - return document class Reader: reads = 0 def read(self): Reader.reads += 1 - Request.api_base = 'https://mutated.example.com' - Request.extra_headers = {'x-source': 'mutated'} - Request.timeout = 9 + bound['api_base'] = 'https://mutated.example.com' + bound['extra_headers'] = {'x-source': 'mutated'} + bound['timeout'] = 9 return b'abc' -document = {'type': 'file', 'file': Reader(), 'mime_type': 'application/pdf'} -request = Request() +bound = { + 'model': 'mistral/mistral-ocr-latest', + 'custom_llm_provider': None, + 'api_key': None, + 'api_base': 'https://original.example.com', + 'extra_headers': {'x-source': 'original'}, + 'timeout': 1, + 'document': {'type': 'file', 'file': Reader(), 'mime_type': 'application/pdf'}, +} kwargs = {} ", ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let (projected, _) = project_request(&request, &kwargs).unwrap(); + let (projected, _) = + project_request(&dict(&locals, "bound"), &dict(&locals, "kwargs")).unwrap(); assert_eq!( py.eval(c"Reader.reads", Some(&locals), Some(&locals)) .unwrap() @@ -524,34 +409,26 @@ document = Document() }); } - fn request_and_kwargs<'py>( + fn bound_and_kwargs<'py>( py: Python<'py>, kwargs: &std::ffi::CStr, - ) -> (Bound<'py, PyAny>, Bound<'py, PyDict>) { + ) -> (Bound<'py, PyDict>, Bound<'py, PyDict>) { let locals = eval( py, c" -class Request: - model = 'mistral/mistral-ocr-latest' - custom_llm_provider = 'mistral' - document = {'type': 'document_url', 'document_url': 'https://example.com/request.pdf'} - api_key = None - api_base = 'https://request.example.com' - extra_headers = {'x-source': 'request'} - timeout = 1 -request = Request() +bound = { + 'model': 'mistral/mistral-ocr-latest', + 'custom_llm_provider': 'mistral', + 'document': {'type': 'document_url', 'document_url': 'https://example.com/bound.pdf'}, + 'api_key': None, + 'api_base': 'https://bound.example.com', + 'extra_headers': {'x-source': 'bound'}, + 'timeout': 1, +} ", ); py.run(kwargs, Some(&locals), Some(&locals)).unwrap(); - ( - locals.get_item("request").unwrap().unwrap(), - locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(), - ) + (dict(&locals, "bound"), dict(&locals, "kwargs")) } #[test] @@ -559,7 +436,7 @@ request = Request() Python::initialize(); Python::attach(|py| { stub_timeout_conversion(py); - let (request, kwargs) = request_and_kwargs( + let (bound, kwargs) = bound_and_kwargs( py, c" kwargs = { @@ -575,7 +452,7 @@ kwargs = { } ", ); - let (projected, _) = project_request(&request, &kwargs).unwrap(); + let (projected, _) = project_request(&bound, &kwargs).unwrap(); assert_eq!( projected.optional_params.keys().collect::>(), ["pages"] @@ -584,12 +461,29 @@ kwargs = { }); } + #[rstest::rstest] + #[case::explicit_none_is_unset(c"bound['pages'] = [1]\nkwargs = {'pages': None}", None)] + #[case::bound_fallback(c"bound['pages'] = [1]\nkwargs = {}", Some(serde_json::json!([1])))] + #[case::keyword_wins(c"bound['pages'] = [1]\nkwargs = {'pages': [0]}", Some(serde_json::json!([0])))] + fn optional_params_read_through_bound_and_drop_none( + #[case] script: &std::ffi::CStr, + #[case] expected: Option, + ) { + Python::initialize(); + Python::attach(|py| { + stub_timeout_conversion(py); + let (bound, kwargs) = bound_and_kwargs(py, script); + let (projected, _) = project_request(&bound, &kwargs).unwrap(); + assert_eq!(projected.optional_params.get("pages").cloned(), expected); + }); + } + #[test] fn replacement_kwargs_project_provider_connection_and_timeout() { Python::initialize(); Python::attach(|py| { stub_timeout_conversion(py); - let (request, kwargs) = request_and_kwargs( + let (bound, kwargs) = bound_and_kwargs( py, c" kwargs = { @@ -602,7 +496,7 @@ kwargs = { } ", ); - let (projected, handles) = project_request(&request, &kwargs).unwrap(); + let (projected, handles) = project_request(&bound, &kwargs).unwrap(); assert_eq!(handles.provider, "azure_ai"); assert_eq!(projected.model, "mistral-ocr-latest"); assert_eq!( diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/host.rs b/litellm-rust/crates/python-bridge/src/routes/responses/host.rs index 25d3fc07332..b22bbc6192b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses/host.rs @@ -1,6 +1,6 @@ use std::convert::Infallible; -use super::super::inference::InferenceHost; +use crate::routes::inference::InferenceHost; use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned}; use litellm_inference_responses::{Error, route::Responses, types::ResponsesCall}; use pyo3::{ diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs b/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs new file mode 100644 index 00000000000..50e5790a532 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs @@ -0,0 +1,103 @@ +mod host; +mod websocket; + +use std::sync::Arc; + +use host::ResponsesPythonHost; +use litellm_auth::AuthServices; +use litellm_cache_response::{CachePolicy, ScopedCache}; +use litellm_callbacks_legacy_python::LoggingOperation; +use litellm_host::{call::HostedMachine, protocol::Protocol}; +use litellm_host_python::present; +use litellm_inference_responses::{ResponsesRoute, route::Responses}; +use litellm_secrets::source::SecretSource; +use pyo3::prelude::*; +use serde_json::Value; +pub(crate) use websocket::ResponsesWebSocketConnection; + +use super::{ + NativeCall, + inference::{InferenceHost, InferenceRoute, run_inference}, +}; +use crate::errors::RustBridgeDeclined; + +const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.responses.route_host"; + +fn run_responses(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult> { + let resolved = call.resolved()?; + if let Some(reason) = py + .import(ROUTE_HOST_MODULE)? + .getattr("decline_reason")? + .call1((&resolved,))? + .extract::>()? + { + return Err(RustBridgeDeclined::new_err(reason)); + } + let argument = |name: &str| present(&call.kwargs, &call.base, name); + let model = argument("model")? + .ok_or_else(|| pyo3::exceptions::PyValueError::new_err("model is required"))? + .extract::()?; + let provider = argument("custom_llm_provider")? + .map(|value| value.extract::()) + .transpose()?; + if provider + .as_deref() + .is_some_and(|provider| provider != "openai") + || model + .strip_prefix("openai/") + .unwrap_or(&model) + .contains('/') + { + return Err(RustBridgeDeclined::new_err( + "native HTTP responses provider", + )); + } + if argument("stream")? + .map(|value| litellm_host_python::from_py::(&value)) + .transpose()? + .is_some_and(|value| value == Value::Bool(true)) + { + return Err(RustBridgeDeclined::new_err( + "native Python responses streaming", + )); + } + let host = InferenceHost::new(resolved.unbind(), ROUTE_HOST_MODULE); + run_inference::(py, call, asynchronous, ResponsesPythonHost(host)) +} + +#[pyfunction] +pub(crate) fn responses(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_responses(py, call, false) +} + +#[pyfunction] +pub(crate) fn aresponses(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_responses(py, call, true) +} + +impl InferenceRoute for ResponsesRoute { + type Protocol = Responses; + const OPERATION: LoggingOperation = LoggingOperation::Responses; + const SYNC_CALL_TYPE: &'static str = "responses"; + const ASYNC_CALL_TYPE: &'static str = "aresponses"; + + fn new( + http: litellm_http::Client, + auth: Arc, + secrets: Arc, + ) -> Self { + Self::new(http, auth, secrets) + } + + fn with_cache(self, cache: ScopedCache) -> Self { + self.with_cache(cache) + } + + fn machine( + self, + call: ::Request, + policy: CachePolicy, + ) -> HostedMachine { + self.machine(call, policy) + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/responses.rs b/litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs similarity index 57% rename from litellm-rust/crates/python-bridge/src/routes/responses.rs rename to litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs index 780f2e6929b..615a838a597 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs @@ -1,121 +1,12 @@ -mod host; - use litellm_inference_responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; -use pyo3::{ - prelude::*, - types::{PyDict, PyTuple}, -}; +use pyo3::prelude::*; use serde_json::Value; use crate::{ - errors::{RustBridgeDeclined, route_error_to_pyerr}, + errors::route_error_to_pyerr, marshal::{marshal_headers, optional_timeout}, }; -fn run_public( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, - asynchronous: bool, -) -> PyResult> { - use super::inference::InferenceHost; - use litellm_callbacks_legacy_python::LoggingOperation; - let host = InferenceHost::new( - request.clone().unbind(), - "litellm.rust_bridge.responses.route_host", - ); - if let Some(reason) = py - .import("litellm.rust_bridge.responses.route_host")? - .getattr("decline_reason")? - .call1((&request,))? - .extract::>()? - { - return Err(RustBridgeDeclined::new_err(reason)); - } - let model = host - .argument(py, &kwargs, "model")? - .ok_or_else(|| pyo3::exceptions::PyValueError::new_err("model is required"))? - .extract::()?; - let provider = host - .argument(py, &kwargs, "custom_llm_provider")? - .map(|value| value.extract::()) - .transpose()?; - if provider - .as_deref() - .is_some_and(|provider| provider != "openai") - || model - .strip_prefix("openai/") - .unwrap_or(&model) - .contains('/') - { - return Err(RustBridgeDeclined::new_err( - "native HTTP responses provider", - )); - } - if host - .argument(py, &kwargs, "stream")? - .map(|value| litellm_host_python::from_py::(&value)) - .transpose()? - .is_some_and(|value| value == Value::Bool(true)) - { - return Err(RustBridgeDeclined::new_err( - "native Python responses streaming", - )); - } - let cache_call_type = if asynchronous { - "aresponses" - } else { - "responses" - }; - crate::cache::admit_native(py, &kwargs, cache_call_type)?; - let (arguments, hooks) = crate::routes::call_hooks( - py, - LoggingOperation::Responses, - &request, - &args, - &kwargs, - asynchronous, - )?; - crate::routes::run_public_call( - py, - arguments, - move |py, arguments, request| { - let route = litellm_inference_responses::ResponsesRoute::new( - crate::http::provider_client(py, arguments, asynchronous)? - .map_err(crate::http::client_error)?, - crate::http::resources().auth.clone(), - crate::secrets::source(py)?, - ); - let (cache, cache_options) = - crate::cache::configured_native(py, arguments, cache_call_type)?; - let route = match cache { - Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new( - cache, - litellm_cache_response::CacheScope::Shared, - )), - None => route, - }; - Ok(route.machine(request, cache_options.policy)) - }, - host::ResponsesPythonHost(host), - hooks, - asynchronous, - ) -} - -#[pyfunction] -pub(crate) fn responses(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; - run_public(py, call.bound.into_any(), call.args, call.kwargs, false) -} - -#[pyfunction] -pub(crate) fn aresponses(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; - run_public(py, call.bound.into_any(), call.args, call.kwargs, true) -} - #[pyclass] pub(crate) struct ResponsesWebSocketConnection { inner: RustResponsesWebSocketConnection, diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 61066e322a7..413245c1ec2 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -172,7 +172,7 @@ impl NativeTraceStorage { BTreeMap, >, ) -> PyResult> { - let table = InsertTable::parse(table).map_err(map_error)?; + let table = table.parse::().map_err(map_error)?; let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; let connection = self.config.storage().writer().clone(); let database = self.config.storage().database().to_owned(); diff --git a/litellm-rust/crates/secrets-types/src/config.rs b/litellm-rust/crates/secrets-types/src/config.rs index 82d48f7b2e1..9b6e82c3a6b 100644 --- a/litellm-rust/crates/secrets-types/src/config.rs +++ b/litellm-rust/crates/secrets-types/src/config.rs @@ -5,18 +5,37 @@ use strum::IntoStaticStr; use crate::SecretValue; -#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, IntoStaticStr, PartialEq, Serialize)] +#[derive( + Clone, + Copy, + Debug, + Deserialize, + Eq, + Hash, + IntoStaticStr, + PartialEq, + Serialize, + strum::VariantArray, +)] #[serde(rename_all = "snake_case")] -#[strum(serialize_all = "snake_case")] pub enum KeyManagementSystem { + #[strum(serialize = "google_kms")] GoogleKms, + #[strum(serialize = "azure_key_vault")] AzureKeyVault, + #[strum(serialize = "aws_secret_manager")] AwsSecretManager, + #[strum(serialize = "google_secret_manager")] GoogleSecretManager, + #[strum(serialize = "hashicorp_vault")] HashicorpVault, + #[strum(serialize = "cyberark")] Cyberark, + #[strum(serialize = "local")] Local, + #[strum(serialize = "aws_kms")] AwsKms, + #[strum(serialize = "custom")] Custom, } diff --git a/litellm-rust/crates/secrets-types/tests/context.rs b/litellm-rust/crates/secrets-types/tests/context.rs index a167e87954f..8f8eeabdb62 100644 --- a/litellm-rust/crates/secrets-types/tests/context.rs +++ b/litellm-rust/crates/secrets-types/tests/context.rs @@ -128,17 +128,21 @@ fn rotation_write_context_preserves_the_operation_context(aws_context: SecretOpe fn provider_context_accepts_only_its_owner( #[case] owner: KeyManagementSystem, #[case] context: SecretOperationContext, -) { - for system in [ - KeyManagementSystem::AwsSecretManager, + #[values( + KeyManagementSystem::GoogleKms, KeyManagementSystem::AzureKeyVault, + KeyManagementSystem::AwsSecretManager, KeyManagementSystem::GoogleSecretManager, KeyManagementSystem::HashicorpVault, KeyManagementSystem::Cyberark, - ] { - assert_eq!(context.validate_for(system).is_ok(), system == owner); - assert!(SecretOperationContext::Default.validate_for(system).is_ok()); - } + KeyManagementSystem::Local, + KeyManagementSystem::AwsKms, + KeyManagementSystem::Custom + )] + system: KeyManagementSystem, +) { + assert_eq!(context.validate_for(system).is_ok(), system == owner); + assert!(SecretOperationContext::Default.validate_for(system).is_ok()); if matches!( owner, KeyManagementSystem::AzureKeyVault | KeyManagementSystem::GoogleSecretManager diff --git a/litellm-rust/crates/secrets/src/oidc.rs b/litellm-rust/crates/secrets/src/oidc.rs index b6e8dbc123b..adeb6c3db00 100644 --- a/litellm-rust/crates/secrets/src/oidc.rs +++ b/litellm-rust/crates/secrets/src/oidc.rs @@ -23,17 +23,22 @@ const OIDC_ALLOWED_CREDENTIAL_DIRS: &str = "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS const DEFAULT_CREDENTIAL_DIRS: &str = "/var/run/secrets,/run/secrets"; #[derive(Clone, Copy, Debug, Eq, Hash, PartialEq, strum::EnumString, strum::AsRefStr)] -#[strum(serialize_all = "snake_case")] pub enum OidcProvider { + #[strum(serialize = "google")] Google, #[strum(serialize = "circleci")] CircleCi, #[strum(serialize = "circleci_v2")] CircleCiV2, + #[strum(serialize = "github")] Github, + #[strum(serialize = "azure")] Azure, + #[strum(serialize = "file")] File, + #[strum(serialize = "env")] Env, + #[strum(serialize = "env_path")] EnvPath, } diff --git a/litellm-rust/crates/testkit/Cargo.toml b/litellm-rust/crates/testkit/Cargo.toml index 98a36a1e87f..97efb753384 100644 --- a/litellm-rust/crates/testkit/Cargo.toml +++ b/litellm-rust/crates/testkit/Cargo.toml @@ -13,6 +13,7 @@ serde.workspace = true semver.workspace = true serde_json.workspace = true sha2.workspace = true +strum.workspace = true tar.workspace = true target-lexicon.workspace = true thiserror.workspace = true diff --git a/litellm-rust/crates/testkit/src/agent/configure.rs b/litellm-rust/crates/testkit/src/agent/configure.rs index a095aacc92f..a41bec9379a 100644 --- a/litellm-rust/crates/testkit/src/agent/configure.rs +++ b/litellm-rust/crates/testkit/src/agent/configure.rs @@ -5,7 +5,7 @@ use semver::Version; use crate::Error; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, PartialEq, Eq, strum::VariantArray)] pub enum Wire { ChatCompletions, Messages, diff --git a/litellm-rust/crates/testkit/tests/configure.rs b/litellm-rust/crates/testkit/tests/configure.rs index ca3587c3474..87f53d2f9c0 100644 --- a/litellm-rust/crates/testkit/tests/configure.rs +++ b/litellm-rust/crates/testkit/tests/configure.rs @@ -2,6 +2,7 @@ use std::path::Path; use litellm_testkit::{ClaudeCode, Codex, Configure, Error, Opencode, Settings, Version, Wire}; use rstest::rstest; +use strum::VariantArray; fn settings(wire: Wire) -> Settings { Settings { @@ -121,7 +122,10 @@ fn opencode_uses_a_different_provider_package_for_every_wire() { .unwrap() .to_owned() }; - let packages = [Wire::ChatCompletions, Wire::Responses, Wire::Messages].map(package); + let packages = Wire::VARIANTS + .iter() + .map(|wire| package(*wire)) + .collect::>(); assert_eq!( packages diff --git a/litellm-rust/crates/traces-cache/src/cache.rs b/litellm-rust/crates/traces-cache/src/cache.rs index eda054e6c7d..e326c0c34d4 100644 --- a/litellm-rust/crates/traces-cache/src/cache.rs +++ b/litellm-rust/crates/traces-cache/src/cache.rs @@ -59,7 +59,7 @@ impl SnapshotKey { } /// Native sessions can resume without a terminal record, so their reads retain `LIVE_TTL`. -/// Other traces with known spend settle after `SETTLED_AFTER_MS` of inactivity. +/// Other traces settle after `SETTLED_AFTER_MS` without activity or recoverable missing spend. #[derive(Clone, Copy, Debug, Eq, PartialEq)] pub enum Freshness { Live, @@ -67,7 +67,7 @@ pub enum Freshness { } impl Freshness { - pub fn of(rows: &[TraceSpansRow], spend_known: bool, snapshot_ms: u64) -> Self { + pub fn of(rows: &[TraceSpansRow], trace: &Trace, snapshot_ms: u64) -> Self { if rows .iter() .any(|row| matches!(row.framework.as_str(), "claude-code" | "claude-agent-sdk")) @@ -82,7 +82,7 @@ impl Freshness { let quiet_ms = i64::try_from(snapshot_ms) .unwrap_or(i64::MAX) .saturating_sub(last_end_ms); - if spend_known && quiet_ms >= SETTLED_AFTER_MS as i64 { + if !trace.gateway_spend_pending && quiet_ms >= SETTLED_AFTER_MS as i64 { Self::Settled } else { Self::Live @@ -388,21 +388,54 @@ mod tests { const LAST_END_MS: u64 = 1_790_742_989_010; #[rstest] - #[case::just_ended(LAST_END_MS, true, Freshness::Live)] - #[case::quiet_just_under(LAST_END_MS + SETTLED_AFTER_MS - 1, true, Freshness::Live)] - #[case::quiet_long_enough(LAST_END_MS + SETTLED_AFTER_MS, true, Freshness::Settled)] - #[case::spend_unknown(LAST_END_MS + SETTLED_AFTER_MS, false, Freshness::Live)] - fn freshness_settles_once_spans_stop_and_spend_is_known( + #[case::just_ended(LAST_END_MS, Freshness::Live)] + #[case::quiet_just_under(LAST_END_MS + SETTLED_AFTER_MS - 1, Freshness::Live)] + #[case::quiet_long_enough(LAST_END_MS + SETTLED_AFTER_MS, Freshness::Settled)] + fn freshness_settles_traces_without_model_calls_once_spans_stop( #[case] snapshot_ms: u64, - #[case] spend_known: bool, #[case] expected: Freshness, ) { assert_eq!( - Freshness::of(&[row("root")], spend_known, snapshot_ms), + Freshness::of(&[row("root")], &trace("root"), snapshot_ms), expected ); } + #[rstest] + #[case::missing_log(None)] + #[case::catalog_estimate(Some(0.25))] + fn estimated_amounts_do_not_settle_pending_gateway_spend(#[case] amount: Option) { + let model = TraceSpansRow { + kind: litellm_traces::ObservationType::Llm, + call_keys: vec![litellm_traces::CallKey::ProviderResponse("response".into())], + call_evidence: Some(litellm_traces::CallEvidenceKind::Complete), + ..row("model") + }; + let rows = [model]; + let original = resolve_trace("trace", "ref", &rows, &[]).unwrap(); + let trace = Trace { + summary: TraceSummary { + llm_calls: 1, + spend: amount, + priced_calls: u64::from(amount.is_some()), + ..original.summary + }, + spans: original + .spans + .into_iter() + .map(|span| litellm_traces::Span { + spend: amount, + ..span + }) + .collect(), + ..original + }; + assert_eq!( + Freshness::of(&rows, &trace, LAST_END_MS + SETTLED_AFTER_MS), + Freshness::Live + ); + } + #[rstest] #[case::claude_code("claude-code")] #[case::claude_agent_sdk("claude-agent-sdk")] @@ -412,7 +445,11 @@ mod tests { ..row("native") }; assert_eq!( - Freshness::of(&[row("root"), native], true, LAST_END_MS + SETTLED_AFTER_MS), + Freshness::of( + &[row("root"), native], + &trace("root"), + LAST_END_MS + SETTLED_AFTER_MS, + ), Freshness::Live ); } diff --git a/litellm-rust/crates/traces-cache/src/list.rs b/litellm-rust/crates/traces-cache/src/list.rs index 1b1dcbbefc3..288f6e19c63 100644 --- a/litellm-rust/crates/traces-cache/src/list.rs +++ b/litellm-rust/crates/traces-cache/src/list.rs @@ -162,10 +162,8 @@ async fn resolve_runs( let spend = spend_window(spans).map_or(&[][..], |window| spend_within(&spend_rows, window)); resolve_trace(&row.trace_id, &row.trace_ref, spans, spend).map(|trace| { - ListedRun::Resolved( - Box::new(trace.summary), - Freshness::of(spans, true, snapshot_ms), - ) + let freshness = Freshness::of(spans, &trace, snapshot_ms); + ListedRun::Resolved(Box::new(trace.summary), freshness) }) }) .collect()) diff --git a/litellm-rust/crates/traces-cache/src/reader.rs b/litellm-rust/crates/traces-cache/src/reader.rs index b0f83a89118..00abbd78601 100644 --- a/litellm-rust/crates/traces-cache/src/reader.rs +++ b/litellm-rust/crates/traces-cache/src/reader.rs @@ -211,14 +211,16 @@ impl TraceReader { .await .map_err(|error| Miss::Read(map_store_error(error)))?; let spend_rows = spend(store, access, &rows).await; - let freshness = Freshness::of(&rows, spend_rows.is_some(), snapshot_ms); resolve_trace( trace_id, trace_ref, &rows, spend_rows.as_deref().unwrap_or_default(), ) - .map(|trace| (trace, freshness)) + .map(|trace| { + let freshness = Freshness::of(&rows, &trace, snapshot_ms); + (trace, freshness) + }) .ok_or(Miss::Absent) }) .await @@ -325,6 +327,7 @@ fn page( let create_page = |count: usize| { let end = position.offset.saturating_add(count).min(spans.len()); Trace { + gateway_spend_pending: snapshot.trace().gateway_spend_pending, summary: snapshot.trace().summary.clone(), agents: snapshot.trace().agents.clone(), spans: spans[position.offset..end].to_vec(), diff --git a/litellm-rust/crates/traces-cache/tests/read.rs b/litellm-rust/crates/traces-cache/tests/read.rs index 22293db98e8..b206f63e58e 100644 --- a/litellm-rust/crates/traces-cache/tests/read.rs +++ b/litellm-rust/crates/traces-cache/tests/read.rs @@ -539,6 +539,110 @@ fn spend_row(response_id: &str, cost: f64) -> SpendByResponseIdsRow { } } +#[rstest] +#[case::missing(&[], None, 0, false, 0.5, None)] +#[case::partial(&[Some(0.25)], Some(0.25), 1, false, 0.5, None)] +#[case::delayed_zero(&[Some(0.25)], Some(0.25), 1, false, 0.0, None)] +#[case::null_amount(&[Some(0.25), None], Some(0.25), 1, false, 0.5, None)] +#[case::complete_zero(&[Some(0.25), Some(0.0)], Some(0.25), 2, true, 0.5, None)] +#[case::no_call_id(&[Some(0.25)], Some(0.25), 1, true, 0.5, Some(CallEvidenceKind::Unknown))] +#[case::incomplete_identity(&[Some(0.25)], Some(0.25), 1, true, 0.5, Some(CallEvidenceKind::Partial))] +#[tokio::test] +async fn gateway_cost_refreshes_until_every_model_call_is_priced( + #[case] initial_costs: &[Option], + #[case] initial_total: Option, + #[case] initial_priced: u64, + #[case] settled: bool, + #[case] final_second: f64, + #[case] terminal: Option, +) { + let rows: Vec<_> = std::iter::once(span(0)) + .chain((1..=2).map(|index| TraceSpansRow { + kind: ObservationType::Llm, + call_keys: vec![CallKey::ProviderResponse(format!("response-{index}"))], + call_evidence: Some(if index == 2 { + terminal.unwrap_or(CallEvidenceKind::Complete) + } else { + CallEvidenceKind::Complete + }), + ..span(index) + })) + .collect(); + let store = FakeStore::with_spans("ref", rows.clone()); + store.set_list_runs(vec![run("trace", "ref")]); + store.set_run_spans(rows); + store.state.lock().unwrap().spend = initial_costs + .iter() + .enumerate() + .map(|(index, cost)| SpendByResponseIdsRow { + spend: *cost, + ..spend_row(&format!("response-{}", index + 1), 0.0) + }) + .collect(); + let reader = TraceReader::new(usize::MAX); + let access = access(); + let detail = reader + .get_trace_page(&store, &access, "trace", "ref", None, 1) + .await + .unwrap() + .unwrap(); + let list = reader + .list_traces(&store, &access, 0, i64::MAX, None, 2) + .await + .unwrap(); + assert_eq!(detail.summary.spend, initial_total); + assert_eq!(detail.summary.priced_calls, initial_priced); + assert_eq!(list.data[0].spend, initial_total); + assert_eq!(list.data[0].priced_calls, initial_priced); + + store.state.lock().unwrap().spend = vec![ + spend_row("response-1", 0.25), + spend_row("response-2", final_second), + ]; + tokio::time::sleep(LIVE_TTL + Duration::from_millis(200)).await; + let refreshed = reader + .get_trace(&store, &access, "trace", "ref") + .await + .unwrap() + .unwrap(); + let listed = reader + .list_traces(&store, &access, 0, i64::MAX, None, 2) + .await + .unwrap(); + let expected = if settled { + initial_total + } else { + Some(0.25 + final_second) + }; + assert_eq!(refreshed.summary.spend, expected); + assert_eq!(listed.data[0].spend, expected); + let expected_priced = if settled { initial_priced } else { 2 }; + assert_eq!(refreshed.summary.priced_calls, expected_priced); + assert_eq!(listed.data[0].priced_calls, expected_priced); + assert_eq!( + store.calls(Operation::TraceSpans), + if settled { 1 } else { 2 } + ); + assert_eq!( + store.calls(Operation::RunSpans), + if settled { 1 } else { 2 } + ); + let pinned = reader + .get_trace_page( + &store, + &access, + "trace", + "ref", + detail.next_cursor.as_deref(), + 1, + ) + .await + .unwrap() + .unwrap(); + assert_eq!(pinned.summary.spend, initial_total); + assert_eq!(pinned.summary.priced_calls, initial_priced); +} + #[rstest] #[tokio::test] async fn failed_batch_spend_lookup_falls_back_to_each_run_instead_of_losing_every_cost() { diff --git a/litellm-rust/crates/traces-clickhouse/src/insert.rs b/litellm-rust/crates/traces-clickhouse/src/insert.rs index ed8db188c14..18c077977cc 100644 --- a/litellm-rust/crates/traces-clickhouse/src/insert.rs +++ b/litellm-rust/crates/traces-clickhouse/src/insert.rs @@ -30,29 +30,19 @@ fn max_insert_bytes() -> Result { pub type InsertRow = BTreeMap>; +#[derive(Clone, Copy, Debug, PartialEq, Eq, strum::EnumString, strum::IntoStaticStr)] +#[strum(parse_err_ty = Error, parse_err_fn = invalid_table)] pub enum InsertTable { + #[strum(serialize = "otel_traces")] OtelTraces, + #[strum(serialize = "spend_logs")] SpendLogs, + #[strum(serialize = "lens_feedback")] LensFeedback, } -impl InsertTable { - pub fn parse(value: &str) -> Result { - match value { - "otel_traces" => Ok(Self::OtelTraces), - "spend_logs" => Ok(Self::SpendLogs), - "lens_feedback" => Ok(Self::LensFeedback), - _ => Err(Error::InvalidTable), - } - } - - fn name(&self) -> &'static str { - match self { - Self::OtelTraces => "otel_traces", - Self::SpendLogs => "spend_logs", - Self::LensFeedback => "lens_feedback", - } - } +fn invalid_table(_name: &str) -> Error { + Error::InvalidTable } pub async fn insert_rows( @@ -81,7 +71,7 @@ pub async fn insert_shared_rows( client, connection, database, - table.name(), + <&'static str>::from(table), &token, body, ) @@ -242,7 +232,25 @@ mod tests { use rstest::rstest; use serde_json::json; - use super::{Error, shared_rows, write_rows}; + use super::{Error, InsertTable, shared_rows, write_rows}; + + #[rstest] + #[case::otel_traces("otel_traces", InsertTable::OtelTraces)] + #[case::spend_logs("spend_logs", InsertTable::SpendLogs)] + #[case::lens_feedback("lens_feedback", InsertTable::LensFeedback)] + fn insert_table_parses_each_table_name(#[case] name: &str, #[case] expected: InsertTable) { + assert_eq!(name.parse::().unwrap(), expected); + } + + #[rstest] + #[case::unknown("events")] + #[case::case_sensitive("OTEL_TRACES")] + fn insert_table_rejects_unknown_names(#[case] name: &str) { + assert!(matches!( + name.parse::(), + Err(Error::InvalidTable) + )); + } #[rstest] fn encoded_limit_counts_utf8_bytes_across_rows() { diff --git a/litellm-rust/crates/traces-clickhouse/src/query.rs b/litellm-rust/crates/traces-clickhouse/src/query.rs index e564dcd31d2..cd1a96eef31 100644 --- a/litellm-rust/crates/traces-clickhouse/src/query.rs +++ b/litellm-rust/crates/traces-clickhouse/src/query.rs @@ -40,12 +40,12 @@ struct MetadataRow { metadata: String, } -#[macro_rules_attribute::apply(request_type)] +#[macro_rules_attribute::apply(crate::request_type)] struct AttributeRow { key: String, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Debug, Eq, Ord, PartialEq, PartialOrd)] #[serde(untagged)] enum PathPart { @@ -53,18 +53,24 @@ enum PathPart { Index(usize), } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Copy, Debug, Eq, Ord, PartialEq, PartialOrd, strum::Display)] #[serde(rename_all = "lowercase")] -#[strum(serialize_all = "lowercase")] #[cfg_attr(feature = "schema", schemars(rename = "MetadataValueType"))] enum JsonKind { + #[strum(serialize = "array")] Array, + #[strum(serialize = "boolean")] Boolean, + #[strum(serialize = "integer")] Integer, + #[strum(serialize = "null")] Null, + #[strum(serialize = "number")] Number, + #[strum(serialize = "object")] Object, + #[strum(serialize = "string")] String, } @@ -82,13 +88,13 @@ impl JsonKind { } } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Copy, Debug, strum::Display)] enum MapValueType { String, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[cfg_attr(feature = "schema", schemars(deny_unknown_fields))] #[cfg_attr(feature = "schema", schemars(rename = "TraceQueryMetadataField"))] struct MetadataField { @@ -97,7 +103,7 @@ struct MetadataField { expression: String, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[cfg_attr(feature = "schema", schemars(rename = "TraceQueryColumn"))] struct ColumnSchema { name: String, @@ -107,7 +113,7 @@ struct ColumnSchema { details: BTreeMap, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[cfg_attr(feature = "schema", schemars(deny_unknown_fields))] #[cfg_attr(feature = "schema", schemars(rename = "TraceQueryTable"))] struct TableSchema { @@ -175,7 +181,7 @@ impl Serialize for Discovery { } } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] struct MetadataSample { fields: Vec, sampled_rows: usize, @@ -194,7 +200,7 @@ impl Unobserved for MetadataSample { } } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[cfg_attr(feature = "schema", schemars(deny_unknown_fields))] #[cfg_attr(feature = "schema", schemars(rename = "TraceQueryMetadata"))] struct MetadataCatalog { @@ -206,7 +212,7 @@ struct MetadataCatalog { scope: &'static str, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[cfg_attr(feature = "schema", schemars(deny_unknown_fields))] #[cfg_attr(feature = "schema", schemars(rename = "TraceQueryAttributeField"))] struct AttributeField { @@ -216,7 +222,7 @@ struct AttributeField { expression: String, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] struct AttributeSample { fields: Vec, truncated: bool, @@ -231,7 +237,7 @@ impl Unobserved for AttributeSample { } } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[cfg_attr(feature = "schema", schemars(deny_unknown_fields))] #[cfg_attr(feature = "schema", schemars(rename = "TraceQueryAttributes"))] struct AttributeCatalog { @@ -243,7 +249,7 @@ struct AttributeCatalog { scope: &'static str, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[cfg_attr(feature = "schema", schemars(deny_unknown_fields))] #[cfg_attr(feature = "schema", schemars(rename = "TraceQueryNormalizedField"))] struct NormalizedField { @@ -267,7 +273,7 @@ impl From<&NormalizedFieldDefinition> for NormalizedField { } } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[cfg_attr(feature = "schema", schemars(deny_unknown_fields))] #[cfg_attr(feature = "schema", schemars(rename = "TraceQueryRelationship"))] struct Relationship { @@ -284,7 +290,7 @@ const RELATIONSHIPS: [Relationship; 1] = [Relationship { meaning: "LiteLLMRequestId contains the first normalized request or provider response ID. This relationship matches response IDs only; CallKeys retains all typed identifiers. Cached requests can share response_id; joins may return multiple spend rows", }]; -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[cfg_attr(feature = "schema", schemars(deny_unknown_fields))] #[cfg_attr(feature = "schema", schemars(rename = "TraceQueryHelp"))] pub struct QueryHelp { diff --git a/litellm-rust/crates/traces-clickhouse/src/query/lens.rs b/litellm-rust/crates/traces-clickhouse/src/query/lens.rs index 97b74de4c20..2af254488bd 100644 --- a/litellm-rust/crates/traces-clickhouse/src/query/lens.rs +++ b/litellm-rust/crates/traces-clickhouse/src/query/lens.rs @@ -18,7 +18,7 @@ pub const LENS_QUERIES: [litellm_traces::ReadQuery; 9] = [ litellm_traces::ReadQuery::FeedbackSummary, ]; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[serde(rename_all = "lowercase")] pub enum ExecutionSource { @@ -27,7 +27,7 @@ pub enum ExecutionSource { Both, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[serde(rename_all = "lowercase")] pub enum ContentSource { @@ -35,7 +35,7 @@ pub enum ContentSource { Requests, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[cfg_attr(feature = "schema", schemars(deny_unknown_fields))] pub struct LensAccessParams { @@ -54,7 +54,7 @@ pub struct LensAccessParams { pub struct LensAvailability; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[serde(deny_unknown_fields)] pub struct LensAvailabilityParams { @@ -62,7 +62,7 @@ pub struct LensAvailabilityParams { pub access: LensAccessParams, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[cfg_attr(feature = "schema", schemars(rename = "ActivityAvailability"))] pub struct LensAvailabilityRow { @@ -89,7 +89,7 @@ impl Query for LensAvailability { pub struct LensAgents; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[serde(deny_unknown_fields)] pub struct LensAgentsParams { @@ -97,7 +97,7 @@ pub struct LensAgentsParams { pub access: LensAccessParams, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[cfg_attr(feature = "schema", schemars(rename = "AgentRow"))] pub struct LensAgentsRow { @@ -114,7 +114,7 @@ impl Query for LensAgents { pub struct TraceAgents; /// Same access shape as `list_traces`: every team, the caller's own traces, or their teams' traces. -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[serde(deny_unknown_fields)] #[cfg_attr(feature = "schema", schemars(deny_unknown_fields))] @@ -138,7 +138,7 @@ pub struct TraceAgentsParams { pub limit: u32, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[cfg_attr(feature = "schema", schemars(rename = "TraceAgentRow"))] pub struct TraceAgentsRow { @@ -174,7 +174,7 @@ impl Query for TraceAgents { pub struct LensSample; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[serde(deny_unknown_fields)] pub struct LensSampleParams { @@ -209,7 +209,7 @@ pub struct LensSampleParams { pub offset: u64, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[cfg_attr(feature = "schema", schemars(rename = "ExecutionRow"))] pub struct LensSampleRow { @@ -265,7 +265,7 @@ impl Query for LensSample { pub struct LensContent; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[serde(deny_unknown_fields)] pub struct LensContentParams { @@ -281,7 +281,7 @@ pub struct LensContentParams { pub offset: u32, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[cfg_attr(feature = "schema", schemars(rename = "PartRow"))] pub struct LensContentRow { @@ -309,7 +309,7 @@ impl Query for LensContent { pub struct LensEvidence; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[serde(deny_unknown_fields)] pub struct LensEvidenceParams { @@ -324,7 +324,7 @@ pub struct LensEvidenceParams { pub quote: String, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[cfg_attr(feature = "schema", schemars(rename = "CountRow"))] pub struct LensEvidenceRow { @@ -345,7 +345,7 @@ impl Query for LensEvidence { pub struct LensFeedbackTarget; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[serde(deny_unknown_fields)] pub struct LensFeedbackTargetParams { @@ -355,7 +355,7 @@ pub struct LensFeedbackTargetParams { pub trace_ref: String, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[cfg_attr(feature = "schema", schemars(rename = "FeedbackTargetRow"))] pub struct LensFeedbackTargetRow { @@ -373,7 +373,7 @@ impl Query for LensFeedbackTarget { pub struct LensFeedback; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[serde(deny_unknown_fields)] pub struct LensFeedbackParams { @@ -383,7 +383,7 @@ pub struct LensFeedbackParams { pub trace_ref: String, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[cfg_attr(feature = "schema", schemars(rename = "FeedbackRow"))] pub struct LensFeedbackRow { @@ -410,7 +410,7 @@ impl Query for LensFeedback { pub struct LensFeedbackSummary; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[serde(deny_unknown_fields)] pub struct LensFeedbackSummaryParams { @@ -419,7 +419,7 @@ pub struct LensFeedbackSummaryParams { pub trace_ids: Vec, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Debug)] #[cfg_attr(feature = "schema", schemars(rename = "FeedbackSummaryRow"))] pub struct LensFeedbackSummaryRow { diff --git a/litellm-rust/crates/traces-clickhouse/src/table.rs b/litellm-rust/crates/traces-clickhouse/src/table.rs index c74cf6d4de1..c9039ad8bed 100644 --- a/litellm-rust/crates/traces-clickhouse/src/table.rs +++ b/litellm-rust/crates/traces-clickhouse/src/table.rs @@ -1,12 +1,14 @@ -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[cfg_attr(feature = "schema", schemars(rename = "TraceTableName"))] #[derive( Clone, Copy, Debug, strum::Display, strum::AsRefStr, strum::EnumIter, strum::IntoStaticStr, )] #[serde(rename_all = "snake_case")] -#[strum(serialize_all = "snake_case")] pub enum TraceTable { + #[strum(serialize = "otel_traces")] OtelTraces, + #[strum(serialize = "agent_traces_by_key")] AgentTracesByKey, + #[strum(serialize = "spend_logs")] SpendLogs, } diff --git a/litellm-rust/crates/traces-clickhouse/tests/queries.rs b/litellm-rust/crates/traces-clickhouse/tests/queries.rs index 6e63adfc347..5cb24dc5548 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/queries.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/queries.rs @@ -126,10 +126,12 @@ async fn lens_sample_keeps_spans_before_window_start_and_excludes_old_only_trace } #[derive(Clone, Copy, strum::AsRefStr)] -#[strum(serialize_all = "snake_case")] enum ScopeCase { + #[strum(serialize = "admin")] Admin, + #[strum(serialize = "team")] Team, + #[strum(serialize = "other_team")] OtherTeam, } diff --git a/litellm-rust/crates/traces/src/normalize/format/claude_code.rs b/litellm-rust/crates/traces/src/normalize/format/claude_code.rs index 37fb2d01330..d849d0fea87 100644 --- a/litellm-rust/crates/traces/src/normalize/format/claude_code.rs +++ b/litellm-rust/crates/traces/src/normalize/format/claude_code.rs @@ -15,14 +15,23 @@ use crate::{ /// Claude Code's built-in tracing, identified by its instrumentation scope. pub(crate) struct ClaudeCode; +#[derive(Debug, PartialEq, Eq, strum::EnumString)] enum SpanType { + #[strum(serialize = "assistant_response")] AssistantResponse, + #[strum(serialize = "tool_result")] ToolResult, + #[strum(serialize = "api_request_body")] ApiRequestBody, + #[strum(serialize = "compaction")] Compaction, + #[strum(serialize = "interaction")] Interaction, + #[strum(serialize = "llm_request")] LlmRequest, + #[strum(serialize = "tool")] Tool, + #[strum(disabled)] Other, } @@ -33,16 +42,7 @@ fn span_type(name: &str, attributes: &BTreeMap) -> SpanType { } else { kind }; - match kind { - "assistant_response" => SpanType::AssistantResponse, - "tool_result" => SpanType::ToolResult, - "api_request_body" => SpanType::ApiRequestBody, - "compaction" => SpanType::Compaction, - "interaction" => SpanType::Interaction, - "llm_request" => SpanType::LlmRequest, - "tool" => SpanType::Tool, - _ => SpanType::Other, - } + kind.parse().unwrap_or(SpanType::Other) } /// `agent:custom:search_agent` -> `search_agent`: the subagent a request ran for. @@ -318,7 +318,7 @@ mod tests { use rstest::rstest; use serde_json::Value; - use super::CLAUDE_CODE_SCOPE; + use super::{CLAUDE_CODE_SCOPE, SpanType, span_type}; use crate::{ Error, normalize::{Normalization, NormalizedSpan, ObservationType}, @@ -355,6 +355,22 @@ mod tests { .collect() } + #[rstest] + #[case::assistant_response("assistant_response", SpanType::AssistantResponse)] + #[case::tool_result("tool_result", SpanType::ToolResult)] + #[case::api_request_body("api_request_body", SpanType::ApiRequestBody)] + #[case::compaction("compaction", SpanType::Compaction)] + #[case::interaction("interaction", SpanType::Interaction)] + #[case::llm_request("llm_request", SpanType::LlmRequest)] + #[case::tool("tool", SpanType::Tool)] + #[case::unknown("surprise", SpanType::Other)] + fn span_type_maps_each_recorded_kind(#[case] kind: &str, #[case] expected: SpanType) { + assert_eq!( + span_type("anything", &attributes(&[("span.type", kind)])), + expected + ); + } + #[rstest] fn notification_prompts_keep_user_provenance_and_compaction_is_system() { let prompt_text = diff --git a/litellm-rust/crates/traces/src/normalize/format/genai.rs b/litellm-rust/crates/traces/src/normalize/format/genai.rs index 5ba9b449737..2bfd6ce5370 100644 --- a/litellm-rust/crates/traces/src/normalize/format/genai.rs +++ b/litellm-rust/crates/traces/src/normalize/format/genai.rs @@ -11,18 +11,24 @@ use crate::{ pub(crate) struct GenAi; #[derive(strum::EnumString)] -#[strum(serialize_all = "snake_case")] pub(crate) enum Operation { + #[strum(serialize = "create_agent")] CreateAgent, + #[strum(serialize = "invoke_agent")] InvokeAgent, + #[strum(serialize = "invoke_workflow")] InvokeWorkflow, + #[strum(serialize = "chat")] Chat, #[strum(serialize = "text_completion", serialize = "completion")] TextCompletion, + #[strum(serialize = "generate_content")] GenerateContent, + #[strum(serialize = "execute_tool")] ExecuteTool, #[strum(serialize = "embeddings", serialize = "embedding")] Embeddings, + #[strum(serialize = "retrieval")] Retrieval, } diff --git a/litellm-rust/crates/traces/src/normalize/metadata.rs b/litellm-rust/crates/traces/src/normalize/metadata.rs index 100008714c5..600bfce4287 100644 --- a/litellm-rust/crates/traces/src/normalize/metadata.rs +++ b/litellm-rust/crates/traces/src/normalize/metadata.rs @@ -24,32 +24,56 @@ pub enum AgentType { serde_with::DeserializeFromStr, serde_with::SerializeDisplay, )] -#[strum(serialize_all = "kebab-case")] pub enum Integration { + #[strum(serialize = "claude-code")] ClaudeCode, + #[strum(serialize = "claude-agent-sdk")] ClaudeAgentSdk, + #[strum(serialize = "openai-codex")] OpenaiCodex, + #[strum(serialize = "deepagents-code")] DeepagentsCode, + #[strum(serialize = "cursor")] Cursor, + #[strum(serialize = "pi")] Pi, + #[strum(serialize = "opencode")] Opencode, + #[strum(serialize = "copilot")] Copilot, + #[strum(serialize = "langchain")] Langchain, + #[strum(serialize = "langgraph")] Langgraph, + #[strum(serialize = "deepagents")] Deepagents, + #[strum(serialize = "autogen")] Autogen, + #[strum(serialize = "crewai")] Crewai, + #[strum(serialize = "google-adk")] GoogleAdk, + #[strum(serialize = "llama-index")] LlamaIndex, + #[strum(serialize = "mastra")] Mastra, + #[strum(serialize = "microsoft-agent-framework")] MicrosoftAgentFramework, + #[strum(serialize = "openai-agents")] OpenaiAgents, + #[strum(serialize = "pydantic-ai")] PydanticAi, + #[strum(serialize = "semantic-kernel")] SemanticKernel, + #[strum(serialize = "strands")] Strands, + #[strum(serialize = "vercel-ai-sdk")] VercelAiSdk, + #[strum(serialize = "instructor")] Instructor, + #[strum(serialize = "n8n")] N8n, + #[strum(serialize = "temporal")] Temporal, #[strum(default)] Other(String), @@ -138,23 +162,36 @@ impl AgentMetadata { } #[derive(strum::EnumString, strum::IntoStaticStr)] -#[strum(serialize_all = "snake_case")] enum MetadataField { + #[strum(serialize = "lc_agent_name")] LcAgentName, + #[strum(serialize = "ls_integration")] LsIntegration, + #[strum(serialize = "ls_agent_type")] LsAgentType, + #[strum(serialize = "ls_agent_purpose")] LsAgentPurpose, + #[strum(serialize = "ls_agent_runtime")] LsAgentRuntime, #[strum(serialize = "ls_agent_runtime_version", to_string = "ls_agent_version")] LsAgentVersion, + #[strum(serialize = "ls_trace_schema_version")] LsTraceSchemaVersion, + #[strum(serialize = "thread_id")] ThreadId, + #[strum(serialize = "ls_subagent_id")] LsSubagentId, + #[strum(serialize = "ls_subagent_type")] LsSubagentType, + #[strum(serialize = "ls_tool_name")] LsToolName, + #[strum(serialize = "ls_model_name")] LsModelName, + #[strum(serialize = "ls_provider")] LsProvider, + #[strum(serialize = "git_branch")] GitBranch, + #[strum(serialize = "git_commit_sha")] GitCommitSha, #[strum(serialize = "repository_url", to_string = "git_repo_url")] GitRepoUrl, diff --git a/litellm-rust/crates/traces/src/normalize/mod.rs b/litellm-rust/crates/traces/src/normalize/mod.rs index 53d5e521799..894fcd1bb03 100644 --- a/litellm-rust/crates/traces/src/normalize/mod.rs +++ b/litellm-rust/crates/traces/src/normalize/mod.rs @@ -29,23 +29,35 @@ use instrumentation::Instrumentation; pub(crate) use messages::{HIDDEN_BLOCK_TYPES, MessagePayload, encode}; pub use metadata::{AgentMetadata, AgentType, Integration}; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Clone, Copy, Debug, Eq, PartialEq, strum::EnumString)] #[serde(rename_all = "lowercase")] -#[strum(serialize_all = "lowercase", ascii_case_insensitive)] +#[strum(ascii_case_insensitive)] #[cfg_attr(feature = "schema", schemars(rename = "SpanType"))] pub enum ObservationType { + #[strum(serialize = "agent")] Agent, + #[strum(serialize = "llm")] Llm, + #[strum(serialize = "tool")] Tool, + #[strum(serialize = "chain")] Chain, + #[strum(serialize = "framework")] Framework, + #[strum(serialize = "retriever")] Retriever, + #[strum(serialize = "embedding")] Embedding, + #[strum(serialize = "reranker")] Reranker, + #[strum(serialize = "guardrail")] Guardrail, + #[strum(serialize = "evaluator")] Evaluator, + #[strum(serialize = "prompt")] Prompt, + #[strum(serialize = "decision")] Decision, } diff --git a/litellm-rust/crates/traces/src/query.rs b/litellm-rust/crates/traces/src/query.rs index 26eaa33d5ac..33dcbb31abf 100644 --- a/litellm-rust/crates/traces/src/query.rs +++ b/litellm-rust/crates/traces/src/query.rs @@ -2,23 +2,38 @@ pub mod guide; pub mod named; #[derive(Clone, Copy, Debug, Eq, PartialEq, strum::EnumString, strum::Display, strum::AsRefStr)] -#[strum(serialize_all = "snake_case")] pub enum ReadQuery { + #[strum(serialize = "list_traces")] ListTraces, + #[strum(serialize = "trace_agents")] TraceAgents, + #[strum(serialize = "trace_spans")] TraceSpans, + #[strum(serialize = "trace_page_spans")] TracePageSpans, + #[strum(serialize = "trace_identity")] TraceIdentity, + #[strum(serialize = "span_detail")] SpanDetail, + #[strum(serialize = "span_error")] SpanError, + #[strum(serialize = "spend_by_response_ids")] SpendByResponseIds, + #[strum(serialize = "availability")] Availability, + #[strum(serialize = "agents")] Agents, + #[strum(serialize = "sample")] Sample, + #[strum(serialize = "content")] Content, + #[strum(serialize = "evidence")] Evidence, + #[strum(serialize = "feedback_target")] FeedbackTarget, + #[strum(serialize = "feedback")] Feedback, + #[strum(serialize = "feedback_summary")] FeedbackSummary, } diff --git a/litellm-rust/crates/traces/src/query/guide.rs b/litellm-rust/crates/traces/src/query/guide.rs index b97f518a7ae..8fbde481df0 100644 --- a/litellm-rust/crates/traces/src/query/guide.rs +++ b/litellm-rust/crates/traces/src/query/guide.rs @@ -1,6 +1,6 @@ use askama::Template; -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[cfg_attr(feature = "schema", schemars(rename = "TraceQueryExample"))] pub struct Example { pub name: String, diff --git a/litellm-rust/crates/traces/src/query/named.rs b/litellm-rust/crates/traces/src/query/named.rs index dfa7ac2f2cd..cd8caf4da48 100644 --- a/litellm-rust/crates/traces/src/query/named.rs +++ b/litellm-rust/crates/traces/src/query/named.rs @@ -1,7 +1,7 @@ use serde::{Deserialize, Serialize}; use std::collections::BTreeMap; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Clone, Debug)] #[cfg_attr(feature = "schema", schemars(rename = "TraceScope"))] pub struct ReadAccessParams { diff --git a/litellm-rust/crates/traces/src/query_access.rs b/litellm-rust/crates/traces/src/query_access.rs index c543ddb5808..12cd551b09d 100644 --- a/litellm-rust/crates/traces/src/query_access.rs +++ b/litellm-rust/crates/traces/src/query_access.rs @@ -1,6 +1,6 @@ use crate::InvalidScope; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Clone, Debug)] #[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)] pub enum QueryScope { diff --git a/litellm-rust/crates/traces/src/request.rs b/litellm-rust/crates/traces/src/request.rs index 0ade6112f02..e9ac54972a6 100644 --- a/litellm-rust/crates/traces/src/request.rs +++ b/litellm-rust/crates/traces/src/request.rs @@ -1,7 +1,7 @@ pub const TRACE_PAGE_SIZE_MIN: u16 = 1; pub const TRACE_PAGE_SIZE_MAX: u16 = 500; -#[macro_rules_attribute::apply(request_type)] +#[macro_rules_attribute::apply(crate::request_type)] #[derive(Clone, Debug)] pub struct TraceListRequest { /// Window start, unix ms. Default: 24h ago @@ -15,7 +15,7 @@ pub struct TraceListRequest { pub cursor: Option, } -#[macro_rules_attribute::apply(request_type)] +#[macro_rules_attribute::apply(crate::request_type)] #[derive(Clone, Debug)] pub struct TraceDetailRequest { #[serde(default)] @@ -31,14 +31,14 @@ pub struct TraceDetailRequest { pub page_size: Option, } -#[macro_rules_attribute::apply(request_type)] +#[macro_rules_attribute::apply(crate::request_type)] #[derive(Clone, Debug)] pub struct TraceSpanRequest { #[serde(default)] pub trace_ref: String, } -#[macro_rules_attribute::apply(request_type)] +#[macro_rules_attribute::apply(crate::request_type)] #[derive(Clone, Debug)] pub struct TraceErrorPageRequest { #[serde(default)] @@ -48,7 +48,7 @@ pub struct TraceErrorPageRequest { pub cursor: Option, } -#[macro_rules_attribute::apply(request_type)] +#[macro_rules_attribute::apply(crate::request_type)] #[derive(Clone, Debug)] #[serde(deny_unknown_fields)] pub struct TraceQueryRequest { diff --git a/litellm-rust/crates/traces/src/resolve/resolution.rs b/litellm-rust/crates/traces/src/resolve/resolution.rs index e19a98a1ede..007fc10edf8 100644 --- a/litellm-rust/crates/traces/src/resolve/resolution.rs +++ b/litellm-rust/crates/traces/src/resolve/resolution.rs @@ -31,7 +31,11 @@ pub(super) struct Resolution<'a> { call_matches: HashMap>, } -pub(super) type CallMatch<'a> = (Option>, SpendMatch); +pub(super) struct CallMatch<'a> { + pub(super) requests: Option>, + pub(super) state: SpendMatch, + pub(super) spend_pending: bool, +} impl<'a> Resolution<'a> { pub(super) fn new(rows: &'a [TraceSpansRow], spend: &'a [SpendRow]) -> Self { @@ -138,18 +142,16 @@ impl<'a> Resolution<'a> { pub(super) fn call_requests(&self, call: usize) -> Option> { self.call_match(call) - .and_then(|(requests, _)| requests.clone()) + .and_then(|matched| matched.requests.clone()) + } + + pub(super) fn gateway_spend_pending(&self) -> bool { + self.call_matches + .values() + .any(|matched| matched.spend_pending) } fn resolve_call_match(&self, call: usize) -> CallMatch<'a> { - if let Some(requests) = self.resolve_call_requests(call) { - return (Some(requests), SpendMatch::Matched); - } - let evidence = self.requests(call); - (None, evidence.unmatched_reason()) - } - - fn resolve_call_requests(&self, call: usize) -> Option> { let wrappers = self.graph.ancestors(call).into_iter().filter(|ancestor| { self.kind(*ancestor) == ObservationType::Llm && self @@ -178,7 +180,7 @@ impl<'a> Resolution<'a> { .collect() }) .flatten(); - let selected: Requests<'a> = transport_requests + let selected = transport_requests .map(|requests| requests.into_iter().flatten().collect()) .into_iter() .chain(sources.iter().filter_map(SpendEvidence::complete_requests)) @@ -187,15 +189,24 @@ impl<'a> Resolution<'a> { .iter() .chain(&transports) .all(|source| source.agrees_with(selected)) - })?; - Some( - selected - .into_iter() - .map(|request| (request.identity(), request)) - .collect::>() - .into_values() - .collect(), - ) + }); + if let Some(selected) = selected { + let requests = spend::unique(selected); + return CallMatch { + spend_pending: spend::request_cost(&requests).is_none(), + requests: Some(requests), + state: SpendMatch::Matched, + }; + } + // A leaf without complete identifiers may still be priced from a wrapper or transport + // once its spend arrives. Only the absence of every complete source is terminal. + let spend_pending = sources.iter().any(SpendEvidence::has_complete_keys) + || (!transports.is_empty() && transports.iter().all(SpendEvidence::has_complete_keys)); + CallMatch { + requests: None, + state: sources[0].unmatched_reason(), + spend_pending, + } } fn transports(&self, call: usize) -> Vec { diff --git a/litellm-rust/crates/traces/src/resolve/spend.rs b/litellm-rust/crates/traces/src/resolve/spend.rs index a6b2fc0dbc5..e74cc3704af 100644 --- a/litellm-rust/crates/traces/src/resolve/spend.rs +++ b/litellm-rust/crates/traces/src/resolve/spend.rs @@ -162,6 +162,10 @@ pub(super) enum SpendEvidence<'a> { } impl<'a> SpendEvidence<'a> { + pub(super) fn has_complete_keys(&self) -> bool { + matches!(self, Self::Complete(keys) if !keys.is_empty()) + } + pub(super) fn unmatched_reason(&self) -> SpendMatch { match self { Self::Unknown => SpendMatch::NoCallId, diff --git a/litellm-rust/crates/traces/src/resolve/view.rs b/litellm-rust/crates/traces/src/resolve/view.rs index 62c27676605..1b1ae7a783e 100644 --- a/litellm-rust/crates/traces/src/resolve/view.rs +++ b/litellm-rust/crates/traces/src/resolve/view.rs @@ -25,8 +25,8 @@ fn optional(value: &str) -> Option { fn span(resolution: &Resolution<'_>, index: usize, trace_start_ns: i64) -> Span { let row = resolution.row(index); let status = resolution.status_source(index); - let (requests, spend_match) = if let Some((requests, matched)) = resolution.call_match(index) { - (requests.clone(), Some(*matched)) + let (requests, spend_match) = if let Some(matched) = resolution.call_match(index) { + (matched.requests.clone(), Some(matched.state)) } else { (resolution.requests(index).complete_requests(), None) }; @@ -249,6 +249,7 @@ pub fn resolve_trace( }), }; Some(Trace { + gateway_spend_pending: resolution.gateway_spend_pending(), summary, agents, spans, diff --git a/litellm-rust/crates/traces/src/response.rs b/litellm-rust/crates/traces/src/response.rs index cd8538ed602..af5173d6ee9 100644 --- a/litellm-rust/crates/traces/src/response.rs +++ b/litellm-rust/crates/traces/src/response.rs @@ -1,4 +1,4 @@ -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Debug)] #[serde(deny_unknown_fields)] pub struct TraceSQLResponse { diff --git a/litellm-rust/crates/traces/src/tenant.rs b/litellm-rust/crates/traces/src/tenant.rs index a097dd947c1..bd519150a24 100644 --- a/litellm-rust/crates/traces/src/tenant.rs +++ b/litellm-rust/crates/traces/src/tenant.rs @@ -1,6 +1,6 @@ /// Who sent a batch of spans. Always taken from the caller's authentication, never from span /// attributes. -#[macro_rules_attribute::apply(request_type)] +#[macro_rules_attribute::apply(crate::request_type)] #[derive(Clone, Debug, Default, Eq, PartialEq)] pub struct Tenant { pub team_id: String, diff --git a/litellm-rust/crates/traces/src/ui.rs b/litellm-rust/crates/traces/src/ui.rs index ef849d2b440..03f1b7e7e9c 100644 --- a/litellm-rust/crates/traces/src/ui.rs +++ b/litellm-rust/crates/traces/src/ui.rs @@ -6,7 +6,7 @@ use serde_json::Value; use crate::normalize::{HIDDEN_BLOCK_TYPES, MessagePayload, encode}; -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Copy, Debug, PartialEq)] #[serde(rename_all = "lowercase")] pub enum ChatRole { @@ -16,7 +16,7 @@ pub enum ChatRole { Tool, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Debug, PartialEq)] #[serde(tag = "kind", rename_all = "snake_case")] #[cfg_attr(feature = "schema", schemars(rename = "UIContent"))] @@ -29,7 +29,7 @@ pub enum UiContent { Text { text: String }, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Debug, PartialEq)] #[cfg_attr(feature = "schema", schemars(rename = "UIMessage"))] pub struct UiMessage { @@ -41,7 +41,7 @@ pub struct UiMessage { pub tool_calls: Option>, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Debug, PartialEq)] #[cfg_attr(feature = "schema", schemars(rename = "UIToolCall"))] pub struct UiToolCall { @@ -49,7 +49,7 @@ pub struct UiToolCall { pub arguments: String, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Debug, PartialEq)] #[cfg_attr(feature = "schema", schemars(rename = "UIField"))] pub struct UiField { diff --git a/litellm-rust/crates/traces/src/view.rs b/litellm-rust/crates/traces/src/view.rs index 729ab5ef1c3..34bf6da07cd 100644 --- a/litellm-rust/crates/traces/src/view.rs +++ b/litellm-rust/crates/traces/src/view.rs @@ -4,7 +4,7 @@ use std::collections::BTreeMap; use crate::ui::UiContent; -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Clone, Copy, Debug, Eq, PartialEq)] #[serde(rename_all = "lowercase")] pub enum SpanStatus { @@ -16,7 +16,7 @@ pub enum SpanStatus { Unset, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Debug, PartialEq)] pub struct Span { pub span_id: String, @@ -41,7 +41,7 @@ pub struct Span { pub spend_match: Option, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Copy, Debug, Eq, PartialEq)] #[serde(rename_all = "snake_case")] pub enum SpendMatch { @@ -53,7 +53,7 @@ pub enum SpendMatch { } /// One distinct agent in a trace: 200 invocations of `researcher` are one node. -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Debug, PartialEq)] pub struct AgentNode { pub name: String, @@ -66,7 +66,7 @@ pub struct AgentNode { pub priced_calls: u64, } -#[macro_rules_attribute::apply(wire_type)] +#[macro_rules_attribute::apply(crate::wire_type)] #[derive(Clone, Copy, Debug, Eq, PartialEq)] #[serde(rename_all = "snake_case")] pub enum RunSourceType { @@ -81,7 +81,7 @@ pub enum RunSourceType { } /// The conversation that started the run, from the `agent.source.*` span attributes. -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Debug, PartialEq)] pub struct RunSource { #[serde(rename = "type")] @@ -94,7 +94,7 @@ pub struct RunSource { pub user: String, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Debug, PartialEq)] pub struct TraceSummary { #[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))] @@ -128,9 +128,12 @@ pub struct TraceSummary { pub source: Option, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Clone, Debug, PartialEq)] pub struct Trace { + /// Read-cache metadata from gateway resolution; display estimates must not clear it. + #[serde(skip)] + pub gateway_spend_pending: bool, pub summary: TraceSummary, pub agents: Vec, pub spans: Vec, @@ -138,14 +141,14 @@ pub struct Trace { pub next_cursor: Option, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Debug, PartialEq)] pub struct TracePage { pub data: Vec, pub next_cursor: Option, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Debug, PartialEq)] pub struct SpanDetail { pub span_id: String, @@ -156,7 +159,7 @@ pub struct SpanDetail { pub attributes: BTreeMap, } -#[macro_rules_attribute::apply(response_type)] +#[macro_rules_attribute::apply(crate::response_type)] #[derive(Debug, PartialEq)] pub struct SpanErrorPage { pub span_id: String, diff --git a/litellm-rust/crates/traces/tests/resolve.rs b/litellm-rust/crates/traces/tests/resolve.rs index 73f691cb2a6..d81aed59a76 100644 --- a/litellm-rust/crates/traces/tests/resolve.rs +++ b/litellm-rust/crates/traces/tests/resolve.rs @@ -494,21 +494,70 @@ fn model_calls_link_the_spend_log_they_were_priced_from() { } #[rstest] -fn model_call_span_cost_agrees_with_the_run_total_when_priced_from_a_wrapper() { +#[case::wrapper_without_leaf_id(false, litellm_traces::CallEvidenceKind::Unknown)] +#[case::wrapper_with_partial_leaf(false, litellm_traces::CallEvidenceKind::Partial)] +#[case::transport_without_leaf_id(true, litellm_traces::CallEvidenceKind::Unknown)] +#[case::transport_with_partial_leaf(true, litellm_traces::CallEvidenceKind::Partial)] +fn model_call_span_cost_agrees_with_the_run_total_when_priced_from_related_evidence( + #[case] transport: bool, + #[case] leaf_evidence: litellm_traces::CallEvidenceKind, +) { let rows = [ TraceSpansRow { - call_keys: vec![litellm_traces::CallKey::LiteLlmRequest("gateway".into())], + trace_id: "trace".into(), + parent_span_id: if transport { "call" } else { "" }.into(), + kind: if transport { + litellm_traces::ObservationType::Chain + } else { + litellm_traces::ObservationType::Llm + }, + call_keys: vec![if transport { + litellm_traces::CallKey::Transport + } else { + litellm_traces::CallKey::LiteLlmRequest("gateway".into()) + }], call_evidence: Some(litellm_traces::CallEvidenceKind::Complete), - ..llm("wrapper", "", "agent", "") + ..llm("evidence", "", "agent", "") + }, + TraceSpansRow { + trace_id: "trace".into(), + call_evidence: Some(leaf_evidence), + ..llm( + "call", + if transport { "" } else { "evidence" }, + "agent", + "chatcmpl-request", + ) }, - llm("call", "wrapper", "agent", ""), ]; + let pending = resolve_trace("trace", "ref", &rows, &[]).unwrap(); + assert_eq!( + pending.spans[1].spend_match, + Some( + if leaf_evidence == litellm_traces::CallEvidenceKind::Partial { + SpendMatch::IncompleteEvidence + } else { + SpendMatch::NoCallId + } + ) + ); + assert!(pending.gateway_spend_pending); + assert!( + serde_json::to_value(&pending) + .unwrap() + .get("gateway_spend_pending") + .is_none() + ); + assert_eq!(pending.summary.spend, None); let logs = [SpendByResponseIdsRow { + trace_id: "trace".into(), + span_id: "evidence".into(), litellm_call_id: "gateway".into(), ..spend("request", "chatcmpl-request", 0.25) }]; let trace = resolve_trace("trace", "ref", &rows, &logs).unwrap(); assert_eq!(trace.summary.spend, Some(0.25)); + assert!(!trace.gateway_spend_pending); assert_eq!( ( trace.spans[1].spend, @@ -1568,7 +1617,12 @@ fn complete_wrapper_reconciles_ambiguous_response( owned_spend("request-a", "response", "team", "", "key", 0.25), owned_spend("request-b", "response", "team", "", "key", 0.5), owned_spend("request-c", "other-response", "team", "", "key", 0.75), + owned_spend("request-d", "response", "team", "", "key", 0.0), ]; + let pending = resolve_trace("trace", "ref", &rows, &logs[1..]).unwrap(); + assert_eq!(pending.spans[1].spend_match, Some(SpendMatch::Ambiguous)); + assert!(pending.gateway_spend_pending); + assert_eq!(pending.summary.spend, None); let trace = resolve_trace("trace", "ref", &rows, &logs).unwrap(); assert_eq!(trace.summary.spend, expected); assert_eq!(trace.agents[0].spend, expected); diff --git a/litellm-rust/docs/adrs/AGENTS.md b/litellm-rust/docs/adrs/AGENTS.md new file mode 100644 index 00000000000..040318c623d --- /dev/null +++ b/litellm-rust/docs/adrs/AGENTS.md @@ -0,0 +1,38 @@ +# ADR 000: Record architecture decisions + +Status: Accepted + +## Context + +Decisions about the Rust core get made in various places, and in synchronous conversations or on whiteboards. +Knowing the context of when and how the decision was made helps greatly to contextualize code for a reader, as well as informing agents about how things should be done. +Especially with the rust core, there are a great many decisions that need to be made in a large solution space, and the reasons for such decisions may not be obvious. + +## Decision + +We write a short ADR for any decision that is hard to reverse or that a new contributor would reasonably question. Each ADR lives in this folder as `adr_NNN_short_title.md`, numbered in order, and has four sections: Context, Decision, Alternatives Considered, and Consequences. If writing one takes more than a few minutes, it is too long. + +ADRs are never edited after they are accepted, apart from the status line. To change course, write a new ADR and mark the old one `Superseded by ADR NNN`. + +ADRs should be hand-written as much as possible, with the intention of being +extremely intention-dense. Agents have a habit of making statements stronger +than they otherwise should be, which makes it challenging to interpret +reasoning, which is the entire point of ADRs. + +ADR filenames should ideally make it clear what decision was made. Their +contents should explain (briefly) alternatives considered, as well as +consequences of the decision in pro/con format. + +## Alternatives Considered + +**Not using ADRs**: This would make us able to move a bit faster, but it makes decisionmaking less clear. Additionally, given the lack of other code documentation, there is no living document other than the code. + +## Consequences + +Benefits: + +- Reviewers and future contributors can find reasoning behind a choice easily + +Costs: + +- A few minutes of thought and writing per significant decision diff --git a/litellm-rust/docs/adrs/adr_001_only_type_inference_inside_rust.md b/litellm-rust/docs/adrs/adr_001_only_type_inference_inside_rust.md new file mode 100644 index 00000000000..0973b13369b --- /dev/null +++ b/litellm-rust/docs/adrs/adr_001_only_type_inference_inside_rust.md @@ -0,0 +1,15 @@ +# ADR 001: Only type requests inside Rust + +Status: Accepted + +## Context + +The python SDK had typed interfaces for requests living in the bridge in Python. However, there were multiple ways for requests to get there, and so the typing is very shallow and ultimately is just a shim placed on top of untyped args and kwargs. + +## Decision + +For the time being, avoid trying to restrict or type requests and their possible arguments inside python. Just pass the raw object over into rust, and have a translation layer at the boundary which handles it. + +## Consequences + +Some python interfaces are a lot less clear about what they expect as input. The typing of the rust module is entirely internal to it, and must be inferred from other channels. \ No newline at end of file diff --git a/litellm/__init__.py b/litellm/__init__.py index 3f93166ecc8..a44af3ceac8 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1421,7 +1421,6 @@ from .exceptions import ( ModelNotMappedError as ModelNotMappedError, ) from .budget_manager import BudgetManager -from .proxy.proxy_cli import run_server from .router import Router from .assistants.main import * from .batches.main import * diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index abefab95e9f..dd41e49d320 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -22,6 +22,7 @@ LITELLM_LOGGING_NAMES: Final = ( # Utils names that support lazy loading via _lazy_import_utils UTILS_NAMES: Final = ( + "run_server", "exception_type", "get_optional_params", "get_response_string", @@ -459,6 +460,7 @@ UTILS_MODULE_NAMES: Final = ( # Import maps for registry pattern - reduces repetition _UTILS_IMPORT_MAP: Final = { + "run_server": ("litellm.proxy.proxy_cli", "run_server"), "exception_type": (".utils", "exception_type"), "get_optional_params": (".utils", "get_optional_params"), "get_response_string": (".utils", "get_response_string"), diff --git a/litellm/_version.py b/litellm/_version.py index 2034cc4f332..140adcb66cf 100644 --- a/litellm/_version.py +++ b/litellm/_version.py @@ -1,6 +1,22 @@ +from typing import Final + import importlib_metadata -try: - version = importlib_metadata.version("litellm") -except Exception: - version = "unknown" + +def _installed_version(distribution: str) -> str | None: + try: + return importlib_metadata.version(distribution) + except Exception: + return None + + +_legacy_version: Final = _installed_version("litellm") +_core_version: Final = _installed_version("litellm-core") + +if _legacy_version is not None and _core_version is not None: + raise RuntimeError( + "litellm and litellm-core are both installed and share the litellm namespace. " + "Install them in separate environments." + ) + +version = _legacy_version or _core_version or "unknown" diff --git a/litellm/anthropic_beta_headers_config.json b/litellm/anthropic_beta_headers_config.json index 332d9b3ad9d..7690372a214 100644 --- a/litellm/anthropic_beta_headers_config.json +++ b/litellm/anthropic_beta_headers_config.json @@ -36,7 +36,8 @@ "web-search-2025-03-05": "web-search-2025-03-05", "mid-conversation-output-config-2026-07-01": "mid-conversation-output-config-2026-07-01", "thinking-display-updates-2026-08-18": "thinking-display-updates-2026-08-18", - "mid-conversation-tool-changes-2026-07-01": "mid-conversation-tool-changes-2026-07-01" + "mid-conversation-tool-changes-2026-07-01": "mid-conversation-tool-changes-2026-07-01", + "inline-tools-2026-09-15": "inline-tools-2026-09-15" }, "azure_ai": { "advisor-tool-2026-03-01": null, @@ -214,7 +215,8 @@ "token-efficient-tools-2025-02-19": null, "tool-search-tool-2025-10-19": "tool-search-tool-2025-10-19", "web-fetch-2025-09-10": null, - "web-search-2025-03-05": "web-search-2025-03-05" + "web-search-2025-03-05": "web-search-2025-03-05", + "inline-tools-2026-09-15": "inline-tools-2026-09-15" }, "databricks": { "advisor-tool-2026-03-01": null, diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 85ef5a93937..b928f33e961 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -71,15 +71,25 @@ class CacheMode(str, Enum): #### LiteLLM.Completion / Embedding Cache #### +def _is_conversation_item(item: object) -> bool: + if isinstance(item, BaseModel): + return True + if not isinstance(item, Mapping): + return False + block: Final = cast(Mapping[str, object], item) # cast-ok: isinstance leaves the key and value types unknown + return block.get("type") != "file" + + def _request_message_count(kwargs: Mapping[str, object]) -> int: - """Chat and Messages API `messages`, else Responses API `input` items; embedding `input` strings count as none""" + """Chat and Messages API `messages`, else Responses API `input` items; embedding strings and file blocks count as none""" messages: Final = kwargs.get("messages") if isinstance(messages, list): return len(messages) input_items: Final = kwargs.get("input") if not isinstance(input_items, list): return 0 - return sum(1 for item in input_items if isinstance(item, (Mapping, BaseModel))) + items: Final = cast(list[object], input_items) # cast-ok: isinstance leaves the element type unknown + return sum(1 for item in items if _is_conversation_item(item)) class Cache: diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index f9215864d1e..9ffc95c6260 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -21,7 +21,7 @@ import time from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Generator, Mapping from typing import TYPE_CHECKING, Any, Final, Optional, TypeVar -from pydantic import ConfigDict, ValidationError +from pydantic import ConfigDict, SkipValidation, ValidationError import litellm from litellm._internal_context import post_response_phase @@ -39,7 +39,7 @@ from litellm.litellm_core_utils.logging_utils import ( from litellm.types.caching import CACHED_STREAM_EVENTS_KEY, EMBEDDING_CACHE_FORMAT_VERSION, CachedEmbedding from litellm.types.integrations.custom_logger import converted_stream_requested from litellm.types.llms.base import LiteLLMBaseModel -from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.llms.openai import ChatCompletionFileObject, ResponsesAPIResponse from litellm.types.rerank import RerankResponse from litellm.types.utils import ( CachingDetails, @@ -71,6 +71,8 @@ from litellm.litellm_core_utils.core_helpers import ( ) from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +EmbeddingCacheInputElement = str | list[int] | ChatCompletionFileObject + class CachingHandlerResponse(LiteLLMBaseModel): """ @@ -82,7 +84,7 @@ class CachingHandlerResponse(LiteLLMBaseModel): cached_result: object | None = None final_embedding_cached_response: EmbeddingResponse | None = None embedding_all_elements_cache_hit: bool = False # this is set to True when all elements in the list have a cache hit in the embedding cache, if true return the final_embedding_cached_response no need to make an API call - embedding_uncached_input: list[str | list[int]] | None = None + embedding_uncached_input: SkipValidation[list[EmbeddingCacheInputElement]] | None = None in_memory_cache_obj: Final = InMemoryCache() @@ -508,7 +510,7 @@ class LLMCachingHandler: _sync_get_cache = sync_get_cache - def handle_kwargs_input_list_or_str(self, kwargs: dict[str, object]) -> list[str]: + def handle_kwargs_input_list_or_str(self, kwargs: dict[str, object]) -> list[EmbeddingCacheInputElement]: """ Handles the input of kwargs['input'] being a list or a string """ @@ -517,7 +519,11 @@ class LLMCachingHandler: elif isinstance(kwargs["input"], list): return kwargs["input"] else: - raise ValueError("input must be a string or a list") + raise litellm.BadRequestError( + message="input must be a string or a list of strings and content blocks", + model=str(kwargs.get("model")), + llm_provider=str(kwargs.get("custom_llm_provider")), + ) def _extract_model_from_cached_results(self, non_null_list: list[tuple[int, CachedEmbedding]]) -> str | None: """ @@ -851,10 +857,7 @@ class LLMCachingHandler: self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs) cached_result: object | None = None if call_type == CallTypes.aembedding.value: - if isinstance(new_kwargs["input"], str): - new_kwargs["input"] = [new_kwargs["input"]] - elif not isinstance(new_kwargs["input"], list): - raise ValueError("input must be a string or a list") + new_kwargs["input"] = self.handle_kwargs_input_list_or_str(new_kwargs) tasks: Final[list[Awaitable[object]]] = [] for idx, i in enumerate(new_kwargs["input"]): preset_cache_key = litellm.cache.get_cache_key(**{**new_kwargs, "input": i}) diff --git a/litellm/chat_completions/dispatch.py b/litellm/chat_completions/dispatch.py index 2eee37d41ed..99900177820 100644 --- a/litellm/chat_completions/dispatch.py +++ b/litellm/chat_completions/dispatch.py @@ -58,14 +58,14 @@ def _public_request( messages: Final = optional_sequence(fields.get("messages")) if not isinstance(model, str) or messages is None: return None - return native_call(args, kwargs, fields) + return native_call(legacy, args, kwargs) def _context(request: NativeCall) -> RouteContext: return RouteContext( Route.CHAT_COMPLETIONS, - provider=optional_str(request.bound.get("custom_llm_provider")), - model=str(request.bound["model"]), + provider=optional_str(request.resolved.get("custom_llm_provider")), + model=str(request.resolved["model"]), ) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 142d624ef71..c5811fcf66a 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -5,6 +5,7 @@ Handler for transforming /chat/completions api requests to litellm.responses req import json import os from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence +from itertools import accumulate, chain from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, TypeVar, Union, cast @@ -114,6 +115,83 @@ def _strip_prompt_cache_breakpoints(input_items: list[object]) -> list[object]: return [_strip_prompt_cache_breakpoints_from_item(item) for item in input_items] +def _is_audio_input_part(part: object) -> bool: + if not isinstance(part, dict): + return False + content_part: Final = cast(dict[str, object], part) # cast-ok: isinstance confirms a content block mapping + return content_part.get("type") == "input_audio" + + +def _breakpoint_of(part: object) -> object: + if not isinstance(part, dict): + return None + content_part: Final = cast(dict[str, object], part) # cast-ok: isinstance confirms a content block mapping + return content_part.get("prompt_cache_breakpoint") + + +def _pending_audio_breakpoint(pending: object, part: object) -> object: + if not _is_audio_input_part(part): + return None + return pending if pending is not None else _breakpoint_of(part) + + +def _carried_breakpoints(content: Sequence[object]) -> tuple[object, ...]: + seeded_from_the_end: Final = chain((None,), reversed(content)) + pending_after_each: Final = tuple(accumulate(seeded_from_the_end, _pending_audio_breakpoint)) + return tuple(reversed(pending_after_each[:-1])) + + +def _with_carried_breakpoint(part: object, marker: object) -> object: + if marker is None or not isinstance(part, dict) or _breakpoint_of(part) is not None: + return part + kept_part: Final = cast(dict[str, object], part) # cast-ok: isinstance confirms a content block mapping + return {**kept_part, "prompt_cache_breakpoint": marker} + + +def _without_audio_input_parts_in_content(value: object) -> object: + if not isinstance(value, list): + return value + content: Final = cast(list[object], value) # cast-ok: isinstance confirms a list of content blocks + if not any(_is_audio_input_part(part) for part in content): + return value + return [ + _with_carried_breakpoint(part, marker) + for part, marker in zip(content, _carried_breakpoints(content)) + if not _is_audio_input_part(part) + ] + + +def _without_audio_input_parts_in_item(value: object) -> object: + if not isinstance(value, dict): + return value + input_item: Final = cast(dict[str, object], value) # cast-ok: isinstance confirms a Responses input item mapping + return { + key: _without_audio_input_parts_in_content(item) if key in ("content", "output") else item + for key, item in input_item.items() + } + + +def _without_audio_input_parts(input_items: list[object]) -> list[object]: + return [_without_audio_input_parts_in_item(item) for item in input_items] + + +def _supports_audio_input(model: str, litellm_params: Mapping[str, object]) -> bool: + custom_llm_provider: Final = litellm_params.get("custom_llm_provider") + provider: Final = custom_llm_provider if isinstance(custom_llm_provider, str) else None + base_model: Final = litellm_params.get("base_model") + return litellm.supports_audio_input(model=model, custom_llm_provider=provider) or ( + isinstance(base_model, str) + and bool(base_model) + and litellm.supports_audio_input(model=base_model, custom_llm_provider=provider) + ) + + +def _drops_audio_input(model: str, litellm_params: Mapping[str, object]) -> bool: + if not (litellm_params.get("drop_params") or litellm.drop_params): + return False + return not _supports_audio_input(model, litellm_params) + + def _provider_metadata(response_fields: Mapping[str, object] | None) -> Mapping[str, object]: return MappingProxyType( { @@ -678,6 +756,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # of instructions, mirroring how non-string system content is already # handled in convert_chat_completion_messages_to_responses_api. is_system_only_request: Final = not converted_input_items and converted_instructions is not None + target_input_items: Final = ( + _without_audio_input_parts(converted_input_items) + if _drops_audio_input(model, litellm_params) + else converted_input_items + ) input_items: Final = ( [ { @@ -687,7 +770,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): } ] if is_system_only_request - else converted_input_items + else target_input_items ) instructions: Final = None if is_system_only_request else converted_instructions @@ -1194,6 +1277,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ) result.append(converted) verbose_logger.debug("Chat provider: image_url -> %s", converted) + elif original_type == "input_audio": + converted = with_prompt_cache_breakpoint( + {"type": "input_audio", "input_audio": item.get("input_audio")}, + _prompt_cache_breakpoint_for_wire(item.get("prompt_cache_breakpoint"), drop_params), + ) + result.append(converted) + verbose_logger.debug("Chat provider: input_audio -> %s", converted) else: # Try to map other types to responses API format item_type = original_type or "input_text" diff --git a/litellm/constants.py b/litellm/constants.py index c4adc0de22f..74273b9ecb9 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1,10 +1,20 @@ import os import sys +from enum import Enum from types import MappingProxyType from typing import Final, Literal from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_in_range, get_env_int_or_none +SERVER_STREAMING_CLASSIFICATION_KEY: Final = "litellm_server_streaming_classification" + + +class ServerStreamingClassification(str, Enum): + MARKER = "litellm-server-streaming" + + +SERVER_STREAMING_CLASSIFICATION_MARKER: Final = ServerStreamingClassification.MARKER + DEFER_PYDANTIC_BUILD: Final = os.getenv("DEFER_PYDANTIC_BUILD", "true") in ("true", "1", "on") DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm")) AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview")) @@ -1835,6 +1845,8 @@ RESET_BUDGET_JOB_LOCK_TTL_SECONDS: Final[int] = 900 PROXY_BATCH_POLLING_INTERVAL: Final = int(os.getenv("PROXY_BATCH_POLLING_INTERVAL", 3600)) MAX_OBJECTS_PER_POLL_CYCLE: Final = max(1, int(os.getenv("MAX_OBJECTS_PER_POLL_CYCLE", 50))) MANAGED_OBJECT_STALENESS_CUTOFF_DAYS: Final = max(1, int(os.getenv("MANAGED_OBJECT_STALENESS_CUTOFF_DAYS", 7))) +BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS: Final = 10.0 +BATCH_OUTPUT_FILE_FALLBACK_RETRY_AFTER_SECONDS: Final = 60 STALE_OBJECT_CLEANUP_BATCH_SIZE: Final = max(1, int(os.getenv("STALE_OBJECT_CLEANUP_BATCH_SIZE", 1000))) # Set PROXY_BATCH_POLLING_ENABLED=false to disable the CheckBatchCost and # CheckResponsesCost background polling jobs entirely (e.g. to avoid DB load on diff --git a/litellm/embeddings/dispatch.py b/litellm/embeddings/dispatch.py index 54408691910..3fd8893464d 100644 --- a/litellm/embeddings/dispatch.py +++ b/litellm/embeddings/dispatch.py @@ -35,14 +35,14 @@ def _public_request( model: Final = fields.get("model") if not isinstance(model, str): return None - return native_call(args, kwargs, fields) + return native_call(legacy, args, kwargs) def _context(request: NativeCall) -> RouteContext: return RouteContext( Route.EMBEDDINGS, - provider=optional_str(request.bound.get("custom_llm_provider")), - model=str(request.bound["model"]), + provider=optional_str(request.resolved.get("custom_llm_provider")), + model=str(request.resolved["model"]), ) diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 22d76242f2f..6047c625665 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -12,7 +12,7 @@ Supported for both `v1/chat/completions` (via the prompt-management hook) and import copy import os import re -from collections.abc import Iterable, Mapping, Sequence +from collections.abc import Callable, Iterable, Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, cast from urllib.parse import urlparse @@ -90,6 +90,14 @@ def _validated_object_list(value: object) -> list[object] | None: return None +def configured_injection_points(value: object) -> Sequence[CacheControlInjectionPoint]: + if not isinstance(value, (list, tuple)): + return () + if all(isinstance(entry, dict) for entry in value): + return cast(Sequence[CacheControlInjectionPoint], value) + return tuple(cast(CacheControlInjectionPoint, entry) for entry in value if isinstance(entry, dict)) + + def supports_openai_prompt_cache_breakpoint(model: str) -> bool: model_map_flag: Final = _model_map_prompt_cache_breakpoint_flag(model) if model_map_flag is not None: @@ -106,7 +114,60 @@ def _model_map_prompt_cache_breakpoint_flag(model: str) -> bool | None: entries: Final = (litellm.model_cost.get(key) for key in (model, model.rsplit("/", 1)[-1])) flags: Final = (entry.get("supports_prompt_cache_breakpoint") for entry in entries if isinstance(entry, dict)) - return next((bool(flag) for flag in flags if flag is not None), None) + return next((flag is True for flag in flags if flag is not None), None) + + +def _hosted_entry_flag(entry: Mapping[str, object], resolve_provider: Callable[[], str | None]) -> bool | None: + flag: Final = entry.get("supports_prompt_cache_breakpoint") + entry_provider: Final = entry.get("litellm_provider") + if flag is None or entry_provider is None or entry_provider == "openai": + return None + return (flag is True) if entry_provider == resolve_provider() else None + + +def _hosted_openai_dialect_flag( + model: str, custom_llm_provider: str | None, resolve_provider: Callable[[str], str | None] +) -> bool | None: + """ + Explicit opt-in for an OpenAI-shaped deployment served by another provider. + + ``model_cost`` is keyed per deployment string, so a flag on the deployment's own + entry states the dialect directly, which a provider name cannot express. A routing + form the map does not key verbatim (``bedrock_mantle/us-east-1/openai.gpt-5.6-sol``) + is read through the candidate keys ``get_model_info`` resolves it with, an entry + keyed with the region outranking the region-free one. A bare name the map does not + key has no such candidates, so it costs no provider lookup. Entries for the openai + provider are left to the caller's api_base check, so an OpenAI-compatible + third-party host is still not assumed to speak the dialect. A bare deployment name + another provider serves (``gpt-6-astra`` on azure_ai) may collide with the openai + row of the same name, so an exact entry only speaks for a deployment when it is + keyed for that deployment's provider. + """ + import litellm + + exact_entry: Final = litellm.model_cost.get(model) + exact_provider: Final = exact_entry.get("litellm_provider") if isinstance(exact_entry, dict) else None + if custom_llm_provider is None and (exact_provider == "openai" or ("/" not in model and exact_provider is None)): + return None + provider: Final = custom_llm_provider or resolve_provider(model) + if provider is None or provider == "openai": + return None + if isinstance(exact_entry, dict) and exact_provider == provider: + return _hosted_entry_flag(exact_entry, lambda: provider) + from litellm.utils import get_potential_model_names + + names: Final = get_potential_model_names(model, provider) + candidates: Final = ( + names["combined_model_name"], + names["region_free_combined_model_name"], + names["split_model"], + names["combined_stripped_model_name"], + names["stripped_model_name"], + names["provider_prefixed_model_name"], + ) + entries: Final = (litellm.model_cost.get(candidate) for candidate in candidates) + flags: Final = (_hosted_entry_flag(entry, lambda: provider) for entry in entries if isinstance(entry, dict)) + return next((flag for flag in flags if flag is not None), None) def targets_openai_api(api_base: object) -> bool: @@ -227,8 +288,8 @@ class AnthropicCacheControlHook(CustomPromptManagement): """ # Extract cache control injection points carry_unmatched: Final = bool(non_default_params.pop(CARRY_UNMATCHED_MESSAGE_POINTS, False)) - injection_points: Final[list[CacheControlInjectionPoint]] = non_default_params.pop( - "cache_control_injection_points", [] + injection_points: Final = configured_injection_points( + non_default_params.pop("cache_control_injection_points", None) ) if not injection_points: return model, messages, non_default_params @@ -282,8 +343,9 @@ class AnthropicCacheControlHook(CustomPromptManagement): if ( openai_dialect and AnthropicCacheControlHook.count_request_cache_breakpoints(processed_messages) > breakpoints_before + and non_default_params.get("prompt_cache_options") is None ): - non_default_params.setdefault("prompt_cache_options", PromptCacheOptions(mode="implicit")) + non_default_params["prompt_cache_options"] = PromptCacheOptions(mode="implicit") # Points this pass did not place: non-message ones for the provider transform, and # the deferred role-targeted ones. Deferring is what reaches the Responses API's @@ -310,7 +372,14 @@ class AnthropicCacheControlHook(CustomPromptManagement): api_base: object = None, prompt_cache_options: object = None, ) -> bool: - if model is None or not supports_openai_prompt_cache_breakpoint(model): + if model is None: + return False + hosted_flag: Final = _hosted_openai_dialect_flag( + model, custom_llm_provider, AnthropicCacheControlHook._resolve_provider + ) + if hosted_flag is not None: + return hosted_flag + if not supports_openai_prompt_cache_breakpoint(model): return False if (custom_llm_provider or AnthropicCacheControlHook._resolve_provider(model)) != "openai": return False @@ -701,7 +770,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): api_base: object, prompt_cache_options: object, ) -> Sequence[Mapping[str, object]]: - if not supports_openai_prompt_cache_breakpoint(model): + if not AnthropicCacheControlHook._may_target_openai_prompt_cache_breakpoint(model, custom_llm_provider): return points return AnthropicCacheControlHook._stamped( points, @@ -711,6 +780,20 @@ class AnthropicCacheControlHook(CustomPromptManagement): ), ) + @staticmethod + def _may_target_openai_prompt_cache_breakpoint(model: str, custom_llm_provider: str | None) -> bool: + """Cheap gate before the dialect is resolved: the model's own row or version, or, when the + serving provider is already known, a row keyed for that provider (``openai.gpt-5.6-sol`` + served by ``bedrock_mantle``), which costs no provider lookup.""" + if supports_openai_prompt_cache_breakpoint(model): + return True + if custom_llm_provider is None: + return False + return ( + _hosted_openai_dialect_flag(model, custom_llm_provider, AnthropicCacheControlHook._resolve_provider) + is not None + ) + @staticmethod def _stamped(points: Sequence[Mapping[str, object]], key: str, value: object) -> Sequence[Mapping[str, object]]: return [{**point, key: value} for point in points] @@ -859,7 +942,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): """ import litellm - configured: Final = non_default_params.get("cache_control_injection_points") + configured: Final = configured_injection_points(non_default_params.get("cache_control_injection_points")) if configured: tools_keeping_marks: Final = tuple( tool for tool in tools or () if not _chat_transform_drops_tool_cache_control(tool) @@ -990,9 +1073,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): bool | None, kwargs.pop("enable_prompt_caching", None) ) cache_control: Final = kwargs.get("cache_control") - configured: Final = 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) - ) + configured: Final = configured_injection_points(kwargs.pop("cache_control_injection_points", None)) injection_points: Final[Sequence[CacheControlInjectionPoint]] = configured or ( AnthropicCacheControlHook.get_default_injection_points( messages=typed_messages, @@ -1028,8 +1109,8 @@ class AnthropicCacheControlHook(CustomPromptManagement): AnthropicCacheControlHook.count_request_cache_breakpoints(messages, system) - breakpoints_before ) AnthropicCacheControlHook.record_gateway_injection(kwargs, breakpoints_added) - if openai_dialect and breakpoints_added > 0: - kwargs.setdefault("prompt_cache_options", PromptCacheOptions(mode="implicit")) + if openai_dialect and breakpoints_added > 0 and kwargs.get("prompt_cache_options") is None: + kwargs["prompt_cache_options"] = PromptCacheOptions(mode="implicit") if remaining: kwargs["cache_control_injection_points"] = remaining return messages, system diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 6af82d9546b..0cf71431199 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -6,7 +6,7 @@ import secrets from collections.abc import Mapping, Sequence from datetime import datetime from types import MappingProxyType -from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, get_args +from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, cast, get_args import httpx @@ -21,11 +21,14 @@ from litellm.litellm_core_utils.core_helpers import ( ) from litellm.secret_managers.main import str_to_bool from litellm.types.guardrails import ( + DEFAULT_GUARDRAIL_STREAM_SCOPE, DynamicGuardrailParams, GuardrailEventHooks, + GuardrailStreamScope, LitellmParams, LoggingOnlyScope, Mode, + runtime_stream_scope, ) from litellm.types.llms.openai import AllMessageValues from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel @@ -49,6 +52,8 @@ from litellm.constants import ( GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS, LOGS_GUARDRAIL_INFORMATION_MARKER, PRE_CALL_EXECUTED_GUARDRAILS_KEY, + SERVER_STREAMING_CLASSIFICATION_KEY, + SERVER_STREAMING_CLASSIFICATION_MARKER, ) from litellm.exceptions import ( BlockedPiiEntityError, @@ -173,6 +178,42 @@ def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None return None +_REALTIME_STREAMING_HOOKS: Final = frozenset({GuardrailEventHooks.realtime_input_transcription}) + + +def without_server_streaming_classification(data: Mapping[str, object]) -> dict[str, object]: + return { + key: value + for key, value in data.items() + if key != SERVER_STREAMING_CLASSIFICATION_KEY or value != SERVER_STREAMING_CLASSIFICATION_MARKER + } + + +def guardrail_request_data_with_streaming( + data: Mapping[str, object], + *, + is_streaming: bool, +) -> dict[str, object]: + data_without_server_classification: Final = without_server_streaming_classification(data) + if not is_streaming: + return data_without_server_classification + return { + **data_without_server_classification, + SERVER_STREAMING_CLASSIFICATION_KEY: SERVER_STREAMING_CLASSIFICATION_MARKER, + } + + +def _request_is_streaming(data: object, event_type: GuardrailEventHooks | None = None) -> bool: + if event_type in _REALTIME_STREAMING_HOOKS: + return True + if not isinstance(data, Mapping): + return False + return ( + data.get("stream") is True + or data.get(SERVER_STREAMING_CLASSIFICATION_KEY) is SERVER_STREAMING_CLASSIFICATION_MARKER + ) + + class CustomGuardrail(CustomLogger): # If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path. use_native_during_call_hook: ClassVar[bool] = False @@ -183,6 +224,9 @@ class CustomGuardrail(CustomLogger): records_own_guardrail_information: ClassVar[bool] = False logging_only_scope: LoggingOnlyScope | None + stream_scope_default: GuardrailStreamScope = DEFAULT_GUARDRAIL_STREAM_SCOPE + stream_scope_by_hook: tuple[tuple[str, GuardrailStreamScope], ...] = () + timeout: float | httpx.Timeout | None = None def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks @@ -258,6 +302,8 @@ class CustomGuardrail(CustomLogger): self.run_in_parallel: bool = run_in_parallel self.scan_raw_request: bool = scan_raw_request self.only_scan_new_messages: bool = only_scan_new_messages + stream_scope_arg: Final[object] = cast(object, kwargs.pop("stream_scope", None)) # cast-ok: config + self.apply_stream_scope(stream_scope_arg) self.logging_only_scope = None if timeout is not None: self.timeout = timeout @@ -1099,6 +1145,23 @@ class CustomGuardrail(CustomLogger): return name in suppressed_compression_guardrails() + def apply_stream_scope(self, stream_scope: object) -> None: + default, by_hook = runtime_stream_scope(stream_scope) + self.stream_scope_default = default + self.stream_scope_by_hook = tuple(by_hook.items()) + + def stream_scope_allows(self, data: object, event_type: GuardrailEventHooks) -> bool: + scope: Final = next( + (scope for hook, scope in self.stream_scope_by_hook if hook == event_type.value), + self.stream_scope_default, + ) + if scope == "both": + return True + is_streaming: Final = _request_is_streaming(data, event_type) + if scope == "streaming": + return is_streaming + return not is_streaming + def should_run_guardrail( self, data, @@ -1142,8 +1205,10 @@ class CustomGuardrail(CustomLogger): data, self.event_hook, event_type ) if result is not None: - return result - return True + tagged_result: Final[bool] = bool(cast(object, result)) # cast-ok: helper return + data_obj: Final[object] = cast(object, data) # cast-ok: data param + return tagged_result and self.stream_scope_allows(data_obj, event_type) + return self.stream_scope_allows(cast(object, data), event_type) # cast-ok: data param return False if ( @@ -1167,8 +1232,9 @@ class CustomGuardrail(CustomLogger): ) result = EnterpriseCustomGuardrailHelper._should_run_if_mode_by_tag(data, self.event_hook, event_type) if result is not None: - return result - return True + mode_tag_result: Final[bool] = bool(cast(object, result)) # cast-ok: helper return + return mode_tag_result and self.stream_scope_allows(cast(object, data), event_type) # cast-ok: data + return self.stream_scope_allows(cast(object, data), event_type) # cast-ok: data param def _event_hook_is_event_type(self, event_type: GuardrailEventHooks) -> bool: """ diff --git a/litellm/integrations/focus/destinations/s3_destination.py b/litellm/integrations/focus/destinations/s3_destination.py index 661cf1933ff..0535523dd99 100644 --- a/litellm/integrations/focus/destinations/s3_destination.py +++ b/litellm/integrations/focus/destinations/s3_destination.py @@ -7,7 +7,6 @@ from collections.abc import Mapping from datetime import timezone from typing import Final, TypedDict -import boto3 from typing_extensions import ReadOnly from .base import FocusDestination, FocusTimeWindow @@ -75,6 +74,8 @@ class FocusS3Destination(FocusDestination): } def _upload(self, content: bytes, object_key: str) -> None: + import boto3 + s3_client: Final = boto3.client("s3", **self._client_kwargs()) s3_client.put_object( Bucket=self.bucket_name, diff --git a/litellm/integrations/otel/emitter.py b/litellm/integrations/otel/emitter.py index e9441ee2a9a..a4cb0a84541 100644 --- a/litellm/integrations/otel/emitter.py +++ b/litellm/integrations/otel/emitter.py @@ -91,13 +91,32 @@ def span_attribute_limit(span: Span) -> int | None: return span._limits.max_span_attributes # pyright: ignore[reportPrivateUsage] # SDK has no public getter -def attribute_budget(span: Span, reserved: int) -> int | None: - """How many mapped attributes fit on ``span`` next to what it already carries and ``reserved`` more.""" +def _carried_keys( + attributes: Mapping[str, AttrValue] | tuple[tuple[str, AttrValue], ...], +) -> frozenset[str]: + """The keys a span already carries: a live span exposes a ``Mapping``, an ended one a tuple of pairs.""" + if isinstance(attributes, Mapping): + return frozenset(attributes) + return frozenset(key for key, _value in attributes) + + +def attribute_budget(span: Span, reserved: int, overwrites: frozenset[str] = frozenset()) -> int | None: + """How many mapped attributes fit on ``span`` next to what it already carries and ``reserved`` more. + + ``overwrites`` are the mapped keys already present on ``span``: setting one + replaces the value in place and consumes no slot against the limit, so only + the genuinely new pre-existing keys reduce the budget. Counting the + overwritten ones too reserves slots the fit can never spend and sheds + indexed message attributes for nothing. + """ limit: Final = span_attribute_limit(span) if limit is None: return None - on_span: Final = len(span.attributes or ()) if isinstance(span, ReadableSpan) else 0 - return limit - on_span - reserved + if not isinstance(span, ReadableSpan): + return limit - reserved + existing_keys: Final = _carried_keys(span.attributes or ()) + fresh: Final = len(existing_keys - overwrites) if overwrites else len(existing_keys) + return limit - fresh - reserved def stamp_error( @@ -284,7 +303,7 @@ class SpanEmitter: ) stamped_later: Final = error_attributes(error) if error else _NO_ATTRIBUTES reserved: Final = len(stamped_later.keys() - mapped.keys()) - for key, value in fit_indexed_messages(mapped, attribute_budget(span, reserved)).items(): + for key, value in fit_indexed_messages(mapped, attribute_budget(span, reserved, frozenset(mapped))).items(): span.set_attribute(key, value) if error: stamped: Final = stamp_error(span, error) diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index 888c90b7661..498ebe655da 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -575,6 +575,7 @@ class OpenTelemetryV2(CustomLogger): request_purpose=call.purpose, trace=call.trace, session_id=call.session_id, + metadata_keys=tuple(self.config.baggage_metadata_keys), ) end_time_ns: Final = to_ns(end_time) if carrier is not None and carrier.span is not None: diff --git a/litellm/integrations/otel/mappers/openinference.py b/litellm/integrations/otel/mappers/openinference.py index a064c2c7e61..8d7b4d7cf4a 100644 --- a/litellm/integrations/otel/mappers/openinference.py +++ b/litellm/integrations/otel/mappers/openinference.py @@ -7,7 +7,7 @@ Phoenix + any other OpenInference-aware backend simultaneously. """ import json -from collections.abc import Callable, Mapping, Sequence +from collections.abc import Callable, Iterator, Mapping, Sequence from itertools import accumulate, chain, groupby from types import MappingProxyType from typing import Final @@ -15,10 +15,12 @@ from typing import Final from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, SpanData from litellm.integrations.otel.mappers.utils import ( MAX_TOOL_DEFINITION_ATTRS_PER_SPAN, + MessageToolCall, collect, - drop_none, + drop_none_pairs, json_if, message_content, + message_tool_calls, output_messages, tool_definition_attrs, ) @@ -31,49 +33,146 @@ from litellm.integrations.otel.model.payloads import ( _INPUT_MESSAGES: Final = "llm.input_messages" _OUTPUT_MESSAGES: Final = "llm.output_messages" _MESSAGE_FAMILIES: Final = (_INPUT_MESSAGES, _OUTPUT_MESSAGES) +_MESSAGE_BASE: Final = -1 + +_ParsedMessage = tuple[object, str | None, tuple[MessageToolCall, ...]] -def _message_key_groups(attrs: Mapping[str, AttrValue]) -> Mapping[tuple[str, int], tuple[str, ...]]: - """Per-index message keys in ``attrs`` grouped by ``(family, index)``.""" - tagged: Final = sorted( - (family, int(key.split(".")[2]), key) - for key in attrs - for family in _MESSAGE_FAMILIES - if key.startswith(f"{family}.") +def _parse_message(message: object) -> _ParsedMessage: + role: Final = message.get("role") if isinstance(message, dict) else None + return role, message_content(message), message_tool_calls(message) + + +def _tool_call_attribute_pairs( + prefix: str, idx: int, tool_calls: tuple[MessageToolCall, ...] +) -> Iterator[tuple[str, str | None]]: + for tool_idx, tool_call in enumerate(tool_calls): + yield f"{prefix}.{idx}.message.tool_calls.{tool_idx}.tool_call.id", tool_call.id + yield f"{prefix}.{idx}.message.tool_calls.{tool_idx}.tool_call.function.name", tool_call.name + yield f"{prefix}.{idx}.message.tool_calls.{tool_idx}.tool_call.function.arguments", tool_call.arguments + + +def _message_attribute_pairs( + prefix: str, + messages: Sequence[_ParsedMessage], + *, + with_tool_call_attrs: bool, +) -> Iterator[tuple[str, str | None]]: + for idx, (role, content, tool_calls) in enumerate(messages): + yield f"{prefix}.{idx}.message.role", role if isinstance(role, str) else None + yield f"{prefix}.{idx}.message.content", content + if with_tool_call_attrs: + yield from _tool_call_attribute_pairs(prefix, idx, tool_calls) + + +def _message_value(messages: Sequence[_ParsedMessage]) -> str: + return json.dumps( + [ + { + "role": role, + "content": content, + **({"tool_calls": [tool_call.to_openai_dict() for tool_call in tool_calls]} if tool_calls else {}), + } + for role, content, tool_calls in messages + ] + ) + + +def _message_key_group(key: str) -> tuple[str, int, int, str] | None: + family: Final = next( + (family for family in _MESSAGE_FAMILIES if key.startswith(f"{family}.")), + None, + ) + if family is None: + return None + parts: Final = key.split(".") + message_idx: Final = int(parts[2]) + tool_idx: Final = int(parts[5]) if parts[4] == "tool_calls" else _MESSAGE_BASE + return family, message_idx, tool_idx, key + + +def _message_key_groups(attrs: Mapping[str, AttrValue]) -> Mapping[tuple[str, int, int], tuple[str, ...]]: + """Message and tool-call keys in ``attrs`` grouped by family, message index, and tool index.""" + tagged: Final = tuple(tag for key in attrs if (tag := _message_key_group(key)) is not None) + return MappingProxyType( + {group: tuple(key for _, _, _, key in keys) for group, keys in groupby(sorted(tagged), key=lambda tag: tag[:3])} + ) + + +def _tool_call_groups_by_message( + groups: Mapping[tuple[str, int, int], tuple[str, ...]], +) -> Mapping[tuple[str, int], tuple[tuple[str, int, int], ...]]: + """Tool-call groups indexed by ``(family, message index)``, each tuple highest tool index first. + + Indexing once keeps the shed order linear in the group count: rescanning the + full group map per message made attribute fitting quadratic on long prompts. + """ + ordered: Final = sorted( + (group for group in groups if group[2] != _MESSAGE_BASE), + key=lambda group: (group[0], group[1], group[2]), ) return MappingProxyType( - {group: tuple(key for _, _, key in keys) for group, keys in groupby(tagged, key=lambda tag: tag[:2])} + { + message: tuple(reversed(tuple(message_tool_groups))) + for message, message_tool_groups in groupby(ordered, key=lambda group: (group[0], group[1])) + } ) -def _shed_order(groups: Mapping[tuple[str, int], tuple[str, ...]]) -> tuple[tuple[str, int], ...]: - """Message groups least valuable first: middle prompt turns, extra choices, then the opener, the newest turn - and the first choice.""" - inputs: Final = sorted(idx for family, idx in groups if family == _INPUT_MESSAGES) - outputs: Final = sorted(idx for family, idx in groups if family == _OUTPUT_MESSAGES) +def _message_shed_groups( + groups: Mapping[tuple[str, int, int], tuple[str, ...]], + tool_call_groups: Mapping[tuple[str, int], tuple[tuple[str, int, int], ...]], + family: str, + message_idx: int, +) -> Iterator[tuple[str, int, int]]: + yield from tool_call_groups.get((family, message_idx), ()) + base_group: Final = (family, message_idx, _MESSAGE_BASE) + if base_group in groups: + yield base_group + + +def _shed_order(groups: Mapping[tuple[str, int, int], tuple[str, ...]]) -> tuple[tuple[str, int, int], ...]: + """Middle inputs, extra choices, pinned inputs, then the first choice, with tool calls before message keys.""" + tool_call_groups: Final = _tool_call_groups_by_message(groups) + inputs: Final = sorted(frozenset(idx for family, idx, _ in groups if family == _INPUT_MESSAGES)) + outputs: Final = sorted(frozenset(idx for family, idx, _ in groups if family == _OUTPUT_MESSAGES)) pinned_inputs: Final = tuple(dict.fromkeys((*inputs[:1], *inputs[-1:]))) - return ( + message_order: Final = ( *((_INPUT_MESSAGES, idx) for idx in inputs[1:-1]), *((_OUTPUT_MESSAGES, idx) for idx in reversed(outputs[1:])), *((_INPUT_MESSAGES, idx) for idx in pinned_inputs), *((_OUTPUT_MESSAGES, idx) for idx in outputs[:1]), ) + return tuple( + chain.from_iterable( + _message_shed_groups(groups, tool_call_groups, family, message_idx) for family, message_idx in message_order + ) + ) + + +_METADATA_KEY: Final = "metadata" def fit_indexed_messages(attrs: Mapping[str, AttrValue], budget: int | None) -> Mapping[str, AttrValue]: - """``attrs`` with whole per-index messages shed, least valuable first, until at most ``budget`` keys remain. + """``attrs`` with indexed message attributes shed, least valuable first, until at most ``budget`` keys remain. ``None`` means the span has no attribute count limit. Every message still rides the ``input.value`` and - ``output.value`` blobs, so shedding a per-index pair loses no content. + ``output.value`` blobs, so shedding a per-index pair loses no content. The ``metadata`` blob sheds only + after every indexed message attribute: message attributes are the indexed, queryable view (the blobs + carry no per-index keys), so a count-squeezed span keeps them and the single metadata key absorbs only + the residual shortfall. Shedding happens here, before ``span.set_attribute``, so the fit is exact and + the SDK's dropped-attributes counter never silently masks the choice. """ if budget is None or len(attrs) <= budget: return attrs groups: Final = _message_key_groups(attrs) - order: Final = _shed_order(groups) - running: Final = tuple(accumulate(len(groups[group]) for group in order)) + sheddable: Final = (*(groups[group] for group in _shed_order(groups)),) + flex: Final = (_METADATA_KEY,) if _METADATA_KEY in attrs else () + candidates: Final = (*sheddable, flex) if flex else sheddable + running: Final = tuple(accumulate(len(candidate) for candidate in candidates)) excess: Final = len(attrs) - budget - shed_count: Final = next((n + 1 for n, total in enumerate(running) if total >= excess), len(order)) - shed: Final = frozenset(chain.from_iterable(groups[group] for group in order[:shed_count])) + shed_count: Final = next((n + 1 for n, total in enumerate(running) if total >= excess), len(candidates)) + shed: Final = frozenset(chain.from_iterable(candidates[:shed_count])) return MappingProxyType({key: value for key, value in attrs.items() if key not in shed}) @@ -85,7 +184,9 @@ class OpenInferenceMapper: - ``llm.model_name`` / ``llm.provider`` / ``llm.invocation_parameters`` - ``llm.input_messages.{i}.message.role`` / ``...content`` - ``llm.output_messages.{i}.message.role`` / ``...content`` + - ``llm.output_messages.{i}.message.tool_calls.{j}.tool_call.*`` - ``llm.token_count.prompt`` / ``...completion`` / ``...total`` + - ``metadata`` — JSON object of allowlisted promoted request metadata - ``input.value`` / ``output.value`` — JSON-serialized request / response """ @@ -121,6 +222,7 @@ class OpenInferenceMapper: "llm.invocation_parameters": lambda d: json_if( collect(OpenInferenceMapper._INVOCATION_PARAMS, d.request_params) ), + "metadata": lambda d: json_if(dict(sorted(d.promoted_metadata.items()))), } def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None: @@ -137,27 +239,36 @@ class OpenInferenceMapper: return { **collect(self._LLM_CALL_ATTRS, data), **collect(self._BLOB_ATTRS, data), - **self._messages(_INPUT_MESSAGES, "input.value", data.messages_in), - **self._messages(_OUTPUT_MESSAGES, "output.value", output_messages(data)), + **self._messages( + _INPUT_MESSAGES, + "input.value", + data.messages_in, + with_tool_call_attrs=False, + ), + **self._messages( + _OUTPUT_MESSAGES, + "output.value", + output_messages(data), + with_tool_call_attrs=True, + ), **self._tools(data), } @staticmethod - def _messages(prefix: str, value_key: str, messages: Sequence[object]) -> AttributeMap: + def _messages( + prefix: str, + value_key: str, + messages: Sequence[object], + *, + with_tool_call_attrs: bool, + ) -> AttributeMap: """``{prefix}.{idx}.message.*`` keys for every message + the ``value_key`` blob of all of them.""" - parsed: Final = [(m.get("role") if isinstance(m, dict) else None, message_content(m)) for m in messages] - attrs: Final = drop_none( - { - key: value - for idx, (role, content) in enumerate(parsed) - for key, value in ( - (f"{prefix}.{idx}.message.role", role if isinstance(role, str) else None), - (f"{prefix}.{idx}.message.content", content), - ) - } + parsed: Final = tuple(_parse_message(message) for message in messages) + attrs: Final = drop_none_pairs( + _message_attribute_pairs(prefix, parsed, with_tool_call_attrs=with_tool_call_attrs) ) if parsed: - attrs[value_key] = json.dumps([{"role": role, "content": content} for role, content in parsed]) + attrs[value_key] = _message_value(parsed) return attrs def _tools(self, data: LLMCallSpanData) -> AttributeMap: diff --git a/litellm/integrations/otel/mappers/utils.py b/litellm/integrations/otel/mappers/utils.py index 5582734585f..80ef2dd553a 100644 --- a/litellm/integrations/otel/mappers/utils.py +++ b/litellm/integrations/otel/mappers/utils.py @@ -7,11 +7,28 @@ they live in one place. import json from collections.abc import Callable, Iterable, Mapping, Sequence +from dataclasses import dataclass from typing import Final from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue from litellm.integrations.otel.model.payloads import LLMCallSpanData, ToolDefinition + +@dataclass(frozen=True, slots=True) +class MessageToolCall: + id: str | None + type: str + name: str | None + arguments: str | None + + def to_openai_dict(self) -> Mapping[str, object]: + return { + "id": self.id, + "type": self.type, + "function": {"name": self.name, "arguments": self.arguments}, + } + + DEFAULT_SPAN_ATTRIBUTE_LIMIT: Final = 128 """The OTel SDK's default per-span attribute count limit.""" @@ -121,3 +138,50 @@ def message_content(message: object) -> str | None: def output_messages(data: LLMCallSpanData) -> list: """The ``message`` payload of each response choice.""" return [c.get("message") for c in data.choices_out if isinstance(c, dict)] + + +def message_tool_calls(message: object) -> tuple[MessageToolCall, ...]: + if not isinstance(message, dict): + return () + tool_calls: Final = message.get("tool_calls") + if not isinstance(tool_calls, (list, tuple)) or not tool_calls: + return () + return tuple(tool_call for value in tool_calls if (tool_call := _message_tool_call(value)) is not None) + + +def _message_tool_call(value: object) -> MessageToolCall | None: + if not isinstance(value, dict): + return None + function: Final = value.get("function") + function_data: Final = function if isinstance(function, dict) else {} + raw_type: Final = value.get("type") + raw_arguments: Final = function_data.get("arguments") + arguments: Final = ( + raw_arguments + if isinstance(raw_arguments, str) + else _stringify_tool_arguments(raw_arguments) + if raw_arguments is not None + else None + ) + name: Final = function_data.get("name") + identifier: Final = value.get("id") + return MessageToolCall( + id=identifier if isinstance(identifier, str) else None, + type=raw_type if isinstance(raw_type, str) else "function", + name=name if isinstance(name, str) else None, + arguments=arguments, + ) + + +def _stringify_tool_arguments(value: object) -> str: + """Serialize non-string tool-call arguments, falling back to ``repr``. + + Arguments normally arrive as already-JSON strings, but provider adapters and + ``model_construct`` responses hand over raw Python objects. ``json.dumps`` + raises on those (tuple-keyed dicts, cycles) and the escaping exception would + lose the whole span, so keep a readable ``repr`` instead. + """ + try: + return json.dumps(value, default=str) + except (TypeError, ValueError): + return repr(value) diff --git a/litellm/integrations/otel/model/baggage.py b/litellm/integrations/otel/model/baggage.py index 131848e1380..606fd79f218 100644 --- a/litellm/integrations/otel/model/baggage.py +++ b/litellm/integrations/otel/model/baggage.py @@ -18,7 +18,7 @@ from collections.abc import Callable, Mapping from types import MappingProxyType from typing import Final -from litellm.integrations.otel.model.metadata import REQUESTER_METADATA_PATH, RequestIdentity +from litellm.integrations.otel.model.metadata import RequestIdentity, allowlisted_metadata from litellm.integrations.otel.model.semconv import GenAI, LiteLLM # Attribute key -> value extractor over (identity, request_model, @@ -92,9 +92,8 @@ def promoted_metadata(metadata: Mapping[str, str], metadata_keys: tuple[str, ... """Allowlisted entries of a flattened metadata mapping under ``litellm.metadata.*``.""" return MappingProxyType( { - f"{LiteLLM.METADATA_PREFIX}{meta_key.removeprefix(REQUESTER_METADATA_PATH)}": value - for meta_key in metadata_keys - if (value := metadata.get(meta_key)) + f"{LiteLLM.METADATA_PREFIX}{meta_key}": value + for meta_key, value in allowlisted_metadata(metadata, metadata_keys).items() } ) diff --git a/litellm/integrations/otel/model/metadata.py b/litellm/integrations/otel/model/metadata.py index 809fc794461..8524df71b44 100644 --- a/litellm/integrations/otel/model/metadata.py +++ b/litellm/integrations/otel/model/metadata.py @@ -61,6 +61,16 @@ REQUESTER_METADATA_KEY: Final = "requester_metadata" REQUESTER_METADATA_PATH: Final = f"{REQUESTER_METADATA_KEY}." +def allowlisted_metadata(metadata: Mapping[str, str], metadata_keys: tuple[str, ...]) -> Mapping[str, str]: + return MappingProxyType( + { + meta_key.removeprefix(REQUESTER_METADATA_PATH): value + for meta_key in metadata_keys + if (value := metadata.get(meta_key)) + } + ) + + @dataclass(frozen=True) class RequestIdentity: call_id: str | None = None diff --git a/litellm/integrations/otel/model/payloads.py b/litellm/integrations/otel/model/payloads.py index ccec93920c9..101e6285946 100644 --- a/litellm/integrations/otel/model/payloads.py +++ b/litellm/integrations/otel/model/payloads.py @@ -12,7 +12,7 @@ from urllib.parse import urlsplit from typing_extensions import ReadOnly, TypedDict -from litellm.integrations.otel.model.metadata import RequestContext, RequestIdentity +from litellm.integrations.otel.model.metadata import RequestContext, RequestIdentity, allowlisted_metadata from litellm.integrations.otel.model.semconv import ( GenAIOperation, GenAIOutputType, @@ -61,6 +61,8 @@ if TYPE_CHECKING: StandardLoggingPayload, ) +_EMPTY_METADATA: Final[Mapping[str, str]] = MappingProxyType({}) + # --- typed sub-structures ---------------------------------------------------- # @@ -437,6 +439,7 @@ class LLMCallSpanData: trace: TraceControls = field(default_factory=TraceControls) session_id: str | None = None embedding_output: EmbeddingOutput | None = None + promoted_metadata: Mapping[str, str] = field(default_factory=lambda: _EMPTY_METADATA) routing_attributes: Mapping[str, RoutingAttributeValue] = field(default_factory=lambda: MappingProxyType({})) @classmethod @@ -449,6 +452,8 @@ class LLMCallSpanData: request_purpose: str | None = None, trace: TraceControls | None = None, session_id: str | None = None, + *, + metadata_keys: tuple[str, ...] = (), ) -> LLMCallSpanData: params: Final = cast(Mapping[str, object], payload.get("model_parameters") or {}) # The single parse of the request's metadata — the request-vs-provider @@ -487,6 +492,7 @@ class LLMCallSpanData: cost=LLMCost.from_breakdown(cast("Mapping[str, object] | None", payload.get("cost_breakdown"))), server=ServerInfo.from_api_base(context.api_base), identity=context.identity, + promoted_metadata=allowlisted_metadata(context.identity.metadata, metadata_keys), is_streaming=as_bool(payload.get("stream")), tools=_extract_tools(params), messages_in=_dicts(payload.get("messages")) if capture_content else (), diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index e11ce2bf472..2d1ef62fde3 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -691,6 +691,24 @@ class PrometheusLogger(CustomLogger): labelnames=self.get_labels_for_metric("litellm_team_rate_limit_used_metric"), ) + self.litellm_project_model_rate_limit_allowed_metric = self._gauge_factory( + "litellm_project_model_rate_limit_allowed_metric", + ( + "Configured rate limit for the Project on the requested model in the current window " + "(model_rpm_limit / model_tpm_limit / model_itpm_limit / model_otpm_limit), by rate_limit_type" + ), + labelnames=self.get_labels_for_metric("litellm_project_model_rate_limit_allowed_metric"), + ) + + self.litellm_project_model_rate_limit_used_metric = self._gauge_factory( + "litellm_project_model_rate_limit_used_metric", + ( + "Requests or tokens the Project has consumed on the requested model in the current rate limit " + "window, by rate_limit_type" + ), + labelnames=self.get_labels_for_metric("litellm_project_model_rate_limit_used_metric"), + ) + ######################################## # LLM API Deployment Metrics / analytics ######################################## @@ -1541,6 +1559,8 @@ class PrometheusLogger(CustomLogger): model_group=standard_logging_payload["model_group"], team=user_api_team, team_alias=user_api_team_alias, + project_id=standard_logging_payload["metadata"].get("user_api_key_project_id"), + project_alias=standard_logging_payload["metadata"].get("user_api_key_project_alias"), org_id=user_api_key_org_id, org_alias=user_api_key_org_alias, user=user_id, @@ -1627,7 +1647,7 @@ class PrometheusLogger(CustomLogger): model_id=enum_values.model_id, ) - self._set_key_and_team_rate_limit_metrics( + self._set_v3_rate_limit_allowed_and_used_metrics( standard_logging_payload=standard_logging_payload, # pyright: ignore[reportArgumentType] # isinstance(dict) above narrows the TypedDict to dict[Unknown, Unknown] enum_values=enum_values, ) @@ -2232,65 +2252,139 @@ class PrometheusLogger(CustomLogger): return None return value - def _set_key_and_team_rate_limit_metrics( + def _set_v3_rate_limit_allowed_and_used_metrics( self, standard_logging_payload: StandardLoggingPayload, enum_values: UserAPIKeyLabelValues, ) -> None: - """ - Export the key-level and team-level RPM / TPM limit and current window - usage from the ``x-ratelimit-{api_key,team}-{limit,remaining}-*`` - headers the v3 rate limiter mirrors into the logging payload. The - limiter already read these counters (from Redis when configured) on - the request path, so no extra store lookup happens here. Descriptors - without a configured limit emit no header, so their series is removed - rather than left at the value from before the limit was dropped. - """ + """Export v3 rate-limit limits and window usage from mirrored logging headers.""" descriptor_gauges: Final[ - tuple[tuple[Literal["api_key", "team"], DEFINED_PROMETHEUS_METRICS, Gauge, Gauge], ...] + tuple[ + tuple[ + Literal[ + "api_key", + "team", + "model_per_project", + "model_per_project_itpm", + "model_per_project_otpm", + ], + Literal["requests", "tokens"], + Literal["requests", "tokens", "input_tokens", "output_tokens"], + DEFINED_PROMETHEUS_METRICS, + Gauge, + Gauge, + ], + ..., + ] ] = ( ( "api_key", + "requests", + "requests", + "litellm_api_key_rate_limit_allowed_metric", + self.litellm_api_key_rate_limit_allowed_metric, + self.litellm_api_key_rate_limit_used_metric, + ), + ( + "api_key", + "tokens", + "tokens", "litellm_api_key_rate_limit_allowed_metric", self.litellm_api_key_rate_limit_allowed_metric, self.litellm_api_key_rate_limit_used_metric, ), ( "team", + "requests", + "requests", "litellm_team_rate_limit_allowed_metric", self.litellm_team_rate_limit_allowed_metric, self.litellm_team_rate_limit_used_metric, ), + ( + "team", + "tokens", + "tokens", + "litellm_team_rate_limit_allowed_metric", + self.litellm_team_rate_limit_allowed_metric, + self.litellm_team_rate_limit_used_metric, + ), + ( + "model_per_project", + "requests", + "requests", + "litellm_project_model_rate_limit_allowed_metric", + self.litellm_project_model_rate_limit_allowed_metric, + self.litellm_project_model_rate_limit_used_metric, + ), + ( + "model_per_project", + "tokens", + "tokens", + "litellm_project_model_rate_limit_allowed_metric", + self.litellm_project_model_rate_limit_allowed_metric, + self.litellm_project_model_rate_limit_used_metric, + ), + ( + "model_per_project_itpm", + "tokens", + "input_tokens", + "litellm_project_model_rate_limit_allowed_metric", + self.litellm_project_model_rate_limit_allowed_metric, + self.litellm_project_model_rate_limit_used_metric, + ), + ( + "model_per_project_otpm", + "tokens", + "output_tokens", + "litellm_project_model_rate_limit_allowed_metric", + self.litellm_project_model_rate_limit_allowed_metric, + self.litellm_project_model_rate_limit_used_metric, + ), ) - for descriptor_key, metric_name, allowed_gauge, used_gauge in descriptor_gauges: - for rate_limit_type in ("requests", "tokens"): - self._set_rate_limit_allowed_and_used_gauges( - standard_logging_payload=standard_logging_payload, - enum_values=enum_values, - descriptor_key=descriptor_key, - metric_name=metric_name, - allowed_gauge=allowed_gauge, - used_gauge=used_gauge, - rate_limit_type=rate_limit_type, - ) + for ( + descriptor_key, + header_rate_limit_type, + rate_limit_type, + metric_name, + allowed_gauge, + used_gauge, + ) in descriptor_gauges: + self._set_rate_limit_allowed_and_used_gauges( + standard_logging_payload=standard_logging_payload, + enum_values=enum_values, + descriptor_key=descriptor_key, + header_rate_limit_type=header_rate_limit_type, + metric_name=metric_name, + allowed_gauge=allowed_gauge, + used_gauge=used_gauge, + rate_limit_type=rate_limit_type, + ) def _set_rate_limit_allowed_and_used_gauges( self, standard_logging_payload: StandardLoggingPayload, enum_values: UserAPIKeyLabelValues, - descriptor_key: Literal["api_key", "team"], + descriptor_key: Literal[ + "api_key", + "team", + "model_per_project", + "model_per_project_itpm", + "model_per_project_otpm", + ], + header_rate_limit_type: Literal["requests", "tokens"], metric_name: DEFINED_PROMETHEUS_METRICS, allowed_gauge: Gauge, used_gauge: Gauge, - rate_limit_type: Literal["requests", "tokens"], + rate_limit_type: Literal["requests", "tokens", "input_tokens", "output_tokens"], ) -> None: limit: Final = self._get_int_from_v3_rate_limit_headers( standard_logging_payload=standard_logging_payload, - header_name=f"x-ratelimit-{descriptor_key}-limit-{rate_limit_type}", + header_name=f"x-ratelimit-{descriptor_key}-limit-{header_rate_limit_type}", ) remaining: Final = self._get_int_from_v3_rate_limit_headers( standard_logging_payload=standard_logging_payload, - header_name=f"x-ratelimit-{descriptor_key}-remaining-{rate_limit_type}", + header_name=f"x-ratelimit-{descriptor_key}-remaining-{header_rate_limit_type}", ) labelled_values: Final = replace(enum_values, rate_limit_type=rate_limit_type) labelnames: Final = self.get_labels_for_metric(metric_name) diff --git a/litellm/litellm_core_utils/duration_parser.py b/litellm/litellm_core_utils/duration_parser.py index 7ac2299b7d8..fe8b71e7623 100644 --- a/litellm/litellm_core_utils/duration_parser.py +++ b/litellm/litellm_core_utils/duration_parser.py @@ -166,6 +166,37 @@ def get_next_standardized_reset_time( return base_midnight + timedelta(days=1) +def _subtract_months(moment: datetime, months: int) -> datetime: + total_months: Final = moment.year * 12 + moment.month - 1 - months + year, month_index = divmod(total_months, 12) + month: Final = month_index + 1 + return moment.replace(year=year, month=month, day=min(moment.day, get_last_day_of_month(year, month))) + + +def get_budget_window_start(duration: str, reset_at: datetime) -> datetime: + """Start of the budget period that ends at `reset_at`, under the same rules + `get_next_standardized_reset_time` used to pick it: `30d` and `Nmo` span calendar + months and a duration it does not recognize resets at the next midnight.""" + value, unit = _parse_duration(_normalize_duration(duration)) + if value is None: + return reset_at - timedelta(days=1) + match unit: + case "mo": + return _subtract_months(reset_at, value) + case "d": + return _subtract_months(reset_at, 1) if value == 30 else reset_at - timedelta(days=value) + case "w": + return reset_at - timedelta(weeks=value) + case "h": + return reset_at - timedelta(hours=value) + case "m": + return reset_at - timedelta(minutes=value) + case "s": + return reset_at - timedelta(seconds=value) + case _: + return reset_at - timedelta(days=1) + + def _setup_timezone(current_time: datetime, timezone_str: str = "UTC") -> tuple[datetime, tzinfo]: """Set up timezone and normalize current time to that timezone.""" try: diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index b5afa962cdc..0ba466ec324 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -2376,6 +2376,10 @@ def exception_type( return original_exception if _is_guardrail_block(original_exception): return original_exception + if isinstance(original_exception, ImportError) and ( + original_exception.name in ("boto3", "botocore") or custom_llm_provider in ("bedrock", "bedrock_mantle") + ): + return original_exception exception_mapping_worked = False exception_provider = custom_llm_provider mappable_exception: Final[_ProviderHTTPException] = cast("_ProviderHTTPException", original_exception) @@ -2398,9 +2402,9 @@ def exception_type( if model or custom_llm_provider: if hasattr(original_exception, "message"): error_str = ( - redact_secret_string(str(original_exception.message)) + redact_secret_string(str(mappable_exception.message)) if _ENABLE_SECRET_REDACTION - else str(original_exception.message) + else str(mappable_exception.message) ) if isinstance(original_exception, BaseException): exception_type = type(original_exception).__name__ diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index 1a48c00bdc2..bab516f54cc 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -67,6 +67,7 @@ OPTIONAL_KWARGS_KEYS: Final = ( "itpm", "otpm", "use_xai_oauth", + "fireworks_forward_user_id", PROVIDER_AFFINITY_HEADER_KWARG_KEY, } ) diff --git a/litellm/litellm_core_utils/optional_imports.py b/litellm/litellm_core_utils/optional_imports.py new file mode 100644 index 00000000000..7b8bbb85d42 --- /dev/null +++ b/litellm/litellm_core_utils/optional_imports.py @@ -0,0 +1,13 @@ +from typing import Final + + +def ensure_optional_import(module: str) -> None: + try: + __import__(module) + except ModuleNotFoundError as error: + if error.name != module: + raise + package: Final = "boto3" if module == "botocore" else module + raise ModuleNotFoundError( + f"Missing optional dependency '{module}'. Run 'pip install {package}'.", name=module + ) from error diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index a781be610a6..473be107cbb 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -782,7 +782,7 @@ class RealTimeStreaming: isinstance(cb, CustomGuardrail) and any( cb.should_run_guardrail( - data=self.request_data, + data={**self.request_data, "stream": True}, event_type=et, ) for et in event_hooks @@ -847,7 +847,7 @@ class RealTimeStreaming: if event_hooks is None: event_hooks = [GuardrailEventHooks.realtime_input_transcription] _realtime_event_types: Final = event_hooks - _check_data: Final = {**self.request_data, "transcript": transcript} + _check_data: Final = {**self.request_data, "transcript": transcript, "stream": True} _already_run: Final[set] = set() for callback in litellm.callbacks: diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 4286d4d6f6a..151830f1589 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -5,7 +5,7 @@ import io import struct from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence from itertools import accumulate -from typing import Final, Literal, cast +from typing import TYPE_CHECKING, Final, Literal, cast import anyio import anyio.lowlevel @@ -30,7 +30,10 @@ from litellm.constants import ( TOKEN_COUNTER_MAX_EXACT_CHARS, ) from litellm.litellm_core_utils.asyncify import asyncify -from litellm.litellm_core_utils.tokenizer import Encoding, HuggingFace, HuggingFaceTokenizer, OpenAIEncoding +from litellm.litellm_core_utils.tokenizer import HuggingFaceTokenizer, OpenAIEncoding + +if TYPE_CHECKING: + from litellm.litellm_core_utils.tokenizer import Encoding, HuggingFace from litellm.litellm_core_utils.url_utils import safe_get from litellm.llms.custom_httpx.http_handler import get_httpx_client from litellm.rust_bridge.tokenizer import get_encoding @@ -680,13 +683,13 @@ def _get_exact_count_function( raise ValueError("Unsupported tokenizer type") -def _encoding_count(encoding: Encoding, text: str) -> int: +def _encoding_count(encoding: "Encoding", text: str) -> int: if isinstance(encoding, OpenAIEncoding): return encoding.count(text) return len(encoding.encode(text, disallowed_special=())) -def openai_tokenizer_encoding(model: str) -> Encoding: +def openai_tokenizer_encoding(model: str) -> "Encoding": """The encoding `token_counter` uses for a model on the `openai_tokenizer` path.""" return get_encoding(openai_tokenizer_encoding_name(model)) diff --git a/litellm/litellm_core_utils/tokenizer.py b/litellm/litellm_core_utils/tokenizer.py index 979b690793b..1c2f3f4e38b 100644 --- a/litellm/litellm_core_utils/tokenizer.py +++ b/litellm/litellm_core_utils/tokenizer.py @@ -14,16 +14,15 @@ from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from functools import partial from pathlib import Path -from types import MappingProxyType +from types import MappingProxyType, UnionType from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias, runtime_checkable import tiktoken -from tokenizers import AddedToken -from tokenizers import Tokenizer as PythonHuggingFaceTokenizer if TYPE_CHECKING: import numpy as np import numpy.typing as npt + from tokenizers import Tokenizer as PythonHuggingFaceTokenizer from litellm.rust_bridge._native import HuggingFaceEncoding from litellm.rust_bridge._native import Tokenizer as NativeTokenizer @@ -217,6 +216,19 @@ class OpenAIEncoding: return allowed +@dataclass(frozen=True, slots=True) +class _NativeAddedToken: + content: str + single_word: bool + lstrip: bool + rstrip: bool + normalized: bool + special: bool + + def __str__(self) -> str: + return self.content + + @dataclass(frozen=True, slots=True) class HuggingFaceTokenizer: """The read-only ``tokenizers.Tokenizer`` surface over the Rust Hugging Face codec.""" @@ -269,7 +281,14 @@ class HuggingFaceTokenizer: def get_vocab_size(self, with_added_tokens: bool = True) -> int: return self._native.get_vocab_size(with_added_tokens) - def get_added_tokens_decoder(self) -> dict[int, AddedToken]: + def get_added_tokens_decoder(self) -> dict[int, _AddedToken]: + try: + from tokenizers import AddedToken + except ModuleNotFoundError as error: + if error.name != "tokenizers": + raise + return {token_id: _NativeAddedToken(*data) for token_id, data in self._native.added_tokens_decoder()} + return { token_id: AddedToken( content, single_word=single_word, lstrip=lstrip, rstrip=rstrip, normalized=normalized, special=special @@ -361,11 +380,40 @@ def _batch_input( Encoding: TypeAlias = tiktoken.Encoding | OpenAIEncoding -HuggingFace: TypeAlias = PythonHuggingFaceTokenizer | HuggingFaceTokenizer -Tokenizer: TypeAlias = Encoding | HuggingFace +if TYPE_CHECKING: + HuggingFace: TypeAlias = PythonHuggingFaceTokenizer | HuggingFaceTokenizer + Tokenizer: TypeAlias = Encoding | HuggingFace + + +def __getattr__(name: str) -> UnionType | type[HuggingFaceTokenizer]: + if name not in {"HuggingFace", "Tokenizer"}: + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + available: Final = HuggingFaceTokenizer if name == "HuggingFace" else Encoding | HuggingFaceTokenizer + try: + from tokenizers import Tokenizer as PythonTokenizer + except ModuleNotFoundError as error: + if error.name == "tokenizers": + return available + raise + return available | PythonTokenizer class _AddedToken(Protocol): + @property + def content(self) -> str: ... + + @property + def single_word(self) -> bool: ... + + @property + def lstrip(self) -> bool: ... + + @property + def rstrip(self) -> bool: ... + + @property + def normalized(self) -> bool: ... + @property def special(self) -> bool: ... diff --git a/litellm/llms/aws_polly/text_to_speech/transformation.py b/litellm/llms/aws_polly/text_to_speech/transformation.py index 133e40dc1ab..01fec02c310 100644 --- a/litellm/llms/aws_polly/text_to_speech/transformation.py +++ b/litellm/llms/aws_polly/text_to_speech/transformation.py @@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Any, Final, Union import httpx from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix +from litellm.litellm_core_utils.optional_imports import ensure_optional_import from litellm.llms.base_llm.text_to_speech.transformation import ( BaseTextToSpeechConfig, TextToSpeechRequestData, @@ -262,11 +263,9 @@ class AWSPollyTextToSpeechConfig(BaseTextToSpeechConfig, BaseAWSLLM): Returns: Tuple of (signed_headers, json_body_string) """ - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - except ImportError: - raise ImportError("Missing boto3 to call AWS Polly. Run 'pip install boto3'.") + ensure_optional_import("botocore") + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest # Get AWS region aws_region_name: Final = litellm_params.get("aws_region_name", self.DEFAULT_REGION) diff --git a/litellm/llms/base_llm/decisions/systemone.py b/litellm/llms/base_llm/decisions/systemone.py deleted file mode 100644 index 43be31b0c30..00000000000 --- a/litellm/llms/base_llm/decisions/systemone.py +++ /dev/null @@ -1,247 +0,0 @@ -"""The Jev / System One wire shape and its translation to and from the OpenAI Decisions shape. - -System One (TypeSafe, Perplexity, OpenRouter, Cloudflare Clef, Strands Decider) takes -{"model", "state", "questions": {name: question}} and answers with {"model", "answers": {name: answer}, "usage"}. -Predicates are `noul` questions, choice options are a `criteria` map, score levels are a `criteria` list. -""" - -import itertools -from collections.abc import Mapping, Sequence -from typing import Final, Literal, TypeAlias - -from pydantic import ConfigDict, TypeAdapter -from typing_extensions import assert_never - -from litellm.llms.base_llm.chat.transformation import BaseLLMException -from litellm.types.llms.base import LiteLLMPydanticObjectBase -from litellm.types.openai_decisions import ( - ChoiceAnswer, - ChoiceProbability, - ChoiceQuestion, - DecisionAnswer, - DecisionChoice, - DecisionInput, - DecisionInputMessage, - DecisionInputPart, - DecisionInputTokensDetails, - DecisionOutputTokensDetails, - DecisionQuestion, - DecisionsRequest, - DecisionsRequestBody, - DecisionsResponse, - DecisionUsage, - PredicateAnswer, - PredicateQuestion, - ScoreAnswer, - ScoreProbability, - ScoreQuestion, -) - - -class SystemOneObjectBase(LiteLLMPydanticObjectBase): - model_config = ConfigDict(extra="allow", frozen=True) - - -class SystemOneNoulAnswer(SystemOneObjectBase): - type: Literal["noul"] - noul: float - - -class SystemOneChoiceAnswer(SystemOneObjectBase): - type: Literal["choice"] - choice: str - confidence: float - probabilities: Mapping[str, float] - - -class SystemOneScoreAnswer(SystemOneObjectBase): - type: Literal["score"] - score: float - confidence: float - probabilities: Mapping[str, float] - - -SystemOneAnswer: TypeAlias = SystemOneNoulAnswer | SystemOneChoiceAnswer | SystemOneScoreAnswer - - -class SystemOneUsage(SystemOneObjectBase): - input_tokens: int = 0 - output_tokens: int = 0 - - -class SystemOneResponse(SystemOneObjectBase): - model: str | None = None - answers: Mapping[str, SystemOneAnswer] - usage: SystemOneUsage | None = None - - -SYSTEM_ONE_RESPONSE_ADAPTER: Final[TypeAdapter[SystemOneResponse]] = TypeAdapter(SystemOneResponse) - - -def _unsupported(what: str, custom_llm_provider: str) -> BaseLLMException: - return BaseLLMException( - status_code=400, - message=f"Decisions provider '{custom_llm_provider}' does not support {what}", - ) - - -def to_system_one_request(model: str, body: DecisionsRequestBody, custom_llm_provider: str) -> dict[str, object]: - keys: Final = question_keys(body.questions, custom_llm_provider) - return { - "model": model, - "state": _state(body.input, custom_llm_provider), - "questions": { - key: _question(question, custom_llm_provider) for key, question in zip(keys, body.questions, strict=True) - }, - } - - -def question_keys(questions: Sequence[DecisionQuestion], custom_llm_provider: str) -> tuple[str, ...]: - """System One keys questions and answers by name, so unnamed questions get a positional key.""" - names: Final = tuple(question.name for question in questions if question.name is not None) - if len(set(names)) != len(names): - raise BaseLLMException( - status_code=400, - message=f"Decisions provider '{custom_llm_provider}' requires a unique name per question", - ) - taken: Final = frozenset(names) - return tuple( - question.name if question.name is not None else _positional_key(index, taken) - for index, question in enumerate(questions) - ) - - -def _positional_key(index: int, taken: frozenset[str]) -> str: - candidates: Final = (f"{'_' * depth}q{index}" for depth in itertools.count()) - return next(key for key in candidates if key not in taken) - - -def _state(input_value: DecisionInput, custom_llm_provider: str) -> str: - if isinstance(input_value, str): - return input_value - return "\n".join(_message_text(message, custom_llm_provider) for message in input_value) - - -def _message_text(message: DecisionInputMessage, custom_llm_provider: str) -> str: - if isinstance(message.content, str): - return message.content - return "\n".join(_part_text(part, custom_llm_provider) for part in message.content) - - -def _part_text(part: DecisionInputPart, custom_llm_provider: str) -> str: - if part.type != "input_text": - raise _unsupported("input_image parts", custom_llm_provider) - return part.text - - -def _question(question: DecisionQuestion, custom_llm_provider: str) -> dict[str, object]: - match question: - case PredicateQuestion(): - return {"type": "noul", "instructions": question.instructions} - case ChoiceQuestion(): - return { - "type": "choice", - "instructions": question.instructions, - "criteria": _choice_criteria(question, custom_llm_provider), - } - case ScoreQuestion(): - return { - "type": "score", - "instructions": question.instructions, - "criteria": [_level_text(level.label, level.description) for level in question.levels], - } - case _: - assert_never(question) - - -def _choice_criteria(question: ChoiceQuestion, custom_llm_provider: str) -> dict[str, str | None]: - values: Final = tuple(_choice_key(choice, custom_llm_provider) for choice in question.choices) - if len(set(values)) != len(values): - raise _unsupported("repeated choice values", custom_llm_provider) - return {value: choice.description for value, choice in zip(values, question.choices, strict=True)} - - -def _choice_key(choice: DecisionChoice, custom_llm_provider: str) -> str: - if not isinstance(choice.value, str): - raise _unsupported("boolean choice values", custom_llm_provider) - return choice.value - - -def _level_text(label: str, description: str | None) -> str: - return description if description is not None else label - - -def to_decisions_response( - system_one: SystemOneResponse, - request: DecisionsRequest, - custom_llm_provider: str, -) -> DecisionsResponse: - keys: Final = question_keys(request.body.questions, custom_llm_provider) - return DecisionsResponse( - model=system_one.model if system_one.model is not None else request.model, - answers=[ - _answer(key, question, system_one.answers, custom_llm_provider) - for key, question in zip(keys, request.body.questions, strict=True) - ], - usage=_usage(system_one.usage), - ) - - -def _answer( - key: str, - question: DecisionQuestion, - answers: Mapping[str, SystemOneAnswer], - custom_llm_provider: str, -) -> DecisionAnswer: - match question, answers.get(key): - case PredicateQuestion(), SystemOneNoulAnswer() as answer: - return PredicateAnswer(type="predicate", name=question.name, probability=answer.noul) - case ChoiceQuestion(), SystemOneChoiceAnswer() as answer: - return ChoiceAnswer( - type="choice", - name=question.name, - choice=answer.choice, - probabilities=_choice_probabilities(question, answer), - confidence=answer.confidence, - ) - case ScoreQuestion(), SystemOneScoreAnswer() as answer: - return ScoreAnswer( - type="score", - name=question.name, - score=answer.score, - probabilities=_score_probabilities(question, answer), - confidence=answer.confidence, - ) - case _: - raise BaseLLMException( - status_code=500, - message=( - f"Decisions provider '{custom_llm_provider}' returned no {question.type} answer " - f"for question '{key}'" - ), - ) - - -def _choice_probabilities(question: ChoiceQuestion, answer: SystemOneChoiceAnswer) -> list[ChoiceProbability]: - return [ - ChoiceProbability(value=choice.value, probability=answer.probabilities.get(str(choice.value), 0.0)) - for choice in question.choices - ] - - -def _score_probabilities(question: ScoreQuestion, answer: SystemOneScoreAnswer) -> list[ScoreProbability]: - return [ - ScoreProbability(value=index, label=level.label, probability=answer.probabilities.get(str(index), 0.0)) - for index, level in enumerate(question.levels) - ] - - -def _usage(usage: SystemOneUsage | None) -> DecisionUsage: - counted: Final = usage if usage is not None else SystemOneUsage() - return DecisionUsage( - input_tokens=counted.input_tokens, - input_tokens_details=DecisionInputTokensDetails(cached_tokens=0, cache_write_tokens=0), - output_tokens=counted.output_tokens, - output_tokens_details=DecisionOutputTokensDetails(reasoning_tokens=0), - total_tokens=counted.input_tokens + counted.output_tokens, - ) diff --git a/litellm/llms/base_llm/files/transformation.py b/litellm/llms/base_llm/files/transformation.py index 30e7d82f840..6c9e534cb7f 100644 --- a/litellm/llms/base_llm/files/transformation.py +++ b/litellm/llms/base_llm/files/transformation.py @@ -135,6 +135,9 @@ class BaseFilesConfig(BaseConfig): ) -> OpenAIFileObject: """Transform file retrieve response into OpenAI format.""" + def is_retrieve_file_response_successful(self, response: httpx.Response) -> bool: + return not httpx.codes.is_error(response.status_code) + @abstractmethod def transform_delete_file_request( self, diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index 0afa5efc29d..bf3c5477886 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -63,7 +63,7 @@ class BedrockAudioTranscriptionRustDispatch: "optional_params": optional_params, "timeout_seconds": timeout_to_seconds(timeout), } - call: Final = NativeCall(args=(), kwargs=fields, bound=fields) + call: Final = NativeCall(args=(), kwargs=fields, base={}) return TranscriptionResponse(**rust(call)) return runtime.run( @@ -96,7 +96,7 @@ class BedrockAudioTranscriptionRustDispatch: "optional_params": optional_params, "timeout_seconds": timeout_to_seconds(timeout), } - call: Final = NativeCall(args=(), kwargs=fields, bound=fields) + call: Final = NativeCall(args=(), kwargs=fields, base={}) return TranscriptionResponse(**await rust(call)) return await runtime.arun( diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 70864070ffd..1e90bf99238 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -32,9 +32,10 @@ from litellm.constants import ( ) from litellm.litellm_core_utils.aws_partition import contains_bedrock_arn, get_aws_dns_suffix from litellm.litellm_core_utils.dd_tracing import tracer +from litellm.litellm_core_utils.optional_imports import ensure_optional_import from litellm.secret_managers.main import get_secret, get_secret_str from litellm.types.llms.base import LiteLLMBaseModel -from litellm.types.llms.bedrock import AWS_AUTH_PARAM_KEYS, AwsAuthParams, AwsSessionTag +from litellm.types.llms.bedrock import AWS_AUTH_PARAM_KEYS, AwsAuthParams, AwsSessionTag, BearerPreparedRequest if TYPE_CHECKING: from botocore.awsrequest import AWSPreparedRequest @@ -446,6 +447,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): # iam_cache: static keys, ambient env (including skip-AssumeRole path), web identity, and # AssumeRole. Do not cache profile / explicit session-token paths here. ######################################################### + ensure_optional_import("botocore") if self._is_auth_with_web_identity_token( aws_web_identity_token, aws_role_name, @@ -775,6 +777,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): aws_region_name = standard_aws_region_name if aws_region_name is None: try: + ensure_optional_import("boto3") import boto3 with tracer.trace("boto3.Session()"): @@ -958,6 +961,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): # For ECS/EC2: call sts:GetCallerIdentity to check if already running as the role try: + ensure_optional_import("boto3") import boto3 with tracer.trace("boto3.client(sts).get_caller_identity"): @@ -1017,6 +1021,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): """ Authenticate with AWS Web Identity Token """ + ensure_optional_import("boto3") import boto3 verbose_logger.debug( @@ -1116,6 +1121,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): aws_session_tags: Sequence[AwsSessionTag] | None = None, ) -> dict: """Handle cross-account role assumption for IRSA.""" + ensure_optional_import("boto3") import boto3 verbose_logger.debug("Cross-account role assumption detected") @@ -1181,6 +1187,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): aws_session_tags: Sequence[AwsSessionTag] | None = None, ) -> dict: """Handle same-account role assumption for IRSA.""" + ensure_optional_import("boto3") import boto3 irsa_sts_kwargs: Final = self._build_sts_client_kwargs( @@ -1287,6 +1294,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): """ Authenticate with AWS Role """ + ensure_optional_import("boto3") import boto3 from botocore.credentials import Credentials @@ -1410,6 +1418,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): """ Authenticate with AWS profile """ + ensure_optional_import("boto3") import boto3 # uses auth values from AWS profile usually stored in ~/.aws/credentials @@ -1448,6 +1457,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): """ Authenticate with AWS Access Key and Secret Key """ + ensure_optional_import("boto3") import boto3 # Check if credentials are already in cache. These credentials have no expiry time. @@ -1468,6 +1478,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): """ Authenticate with AWS Environment Variables """ + ensure_optional_import("boto3") import boto3 with tracer.trace("boto3.Session()"): @@ -1561,10 +1572,6 @@ class BaseAWSLLM(SignsRequestsWithAWS): Returns: Credentials: Boto3 credentials object """ - try: - from botocore.credentials import Credentials - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") aws_region_name: Final = self._get_aws_region_name(optional_params, model) optional_params.pop("aws_region_name", None) auth_params: Final = pop_aws_auth_params(optional_params) @@ -1594,23 +1601,26 @@ class BaseAWSLLM(SignsRequestsWithAWS): headers: dict, api_key: str | None = None, supports_bearer_token: bool = True, - ) -> AWSPreparedRequest: + ) -> AWSPreparedRequest | BearerPreparedRequest: aws_bearer_token: Final = bedrock_bearer_token(api_key) if supports_bearer_token else None if aws_bearer_token is not None: - try: - from botocore.awsrequest import AWSRequest - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") headers["Authorization"] = f"Bearer {aws_bearer_token}" - request = AWSRequest(method="POST", url=endpoint_url, data=data, headers=headers) + bearer_request: Final = httpx.Request("POST", endpoint_url, content=data, headers=headers) + return BearerPreparedRequest( + method="POST", + url=str(bearer_request.url), + headers={ + name.decode("ascii"): value.decode(bearer_request.headers.encoding) + for name, value in bearer_request.headers.raw + }, + body=bearer_request.content, + ) else: - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - from botocore.exceptions import NoCredentialsError - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + ensure_optional_import("botocore") + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest + from botocore.exceptions import NoCredentialsError if credentials is None: raise NoCredentialsError() @@ -1703,12 +1713,10 @@ class BaseAWSLLM(SignsRequestsWithAWS): return headers, json.dumps(request_data).encode() # If no bearer token is set, proceed with the existing SigV4 authentication - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - from botocore.credentials import Credentials - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + ensure_optional_import("botocore") + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest + from botocore.credentials import Credentials auth_params: Final = AwsAuthParams.model_validate(optional_params) aws_region_name: Final = self._get_aws_region_name(optional_params=optional_params, model=model) @@ -1754,11 +1762,9 @@ def sign_aws_json_post( body: str, headers: Mapping[str, str], ) -> AWSPreparedRequest: - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - except ImportError: - raise ImportError(f"Missing boto3 to call {service_name}. Run 'pip install boto3'.") + ensure_optional_import("botocore") + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest aws_request: Final = AWSRequest(method="POST", url=url, data=body, headers=headers) SigV4Auth(get_credentials(), service_name, aws_region_name).add_auth(aws_request) diff --git a/litellm/llms/bedrock/chat/chat_completions/transformation.py b/litellm/llms/bedrock/chat/chat_completions/transformation.py index 728b8166cee..b3697979b3b 100644 --- a/litellm/llms/bedrock/chat/chat_completions/transformation.py +++ b/litellm/llms/bedrock/chat/chat_completions/transformation.py @@ -36,6 +36,7 @@ from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, bedrock_bearer_token from litellm.llms.bedrock.common_utils import ( BedrockError, bedrock_model_is_openai_gpt, + bedrock_runtime_chat_completions_serves_reasoning_inline, split_bedrock_region_path, ) from litellm.llms.openai.chat.gpt_transformation import OpenAIChatCompletionStreamingHandler @@ -213,7 +214,12 @@ def split_reasoning_tag(content: str) -> tuple[str | None, str]: class BedrockRuntimeChatCompletionsStreamingHandler(OpenAIChatCompletionStreamingHandler): - """OpenAI chunk parsing plus the ```` split, tracked per choice index.""" + """OpenAI chunk parsing plus the inline ```` split, tracked per choice index. + + Every chunk echoes the model id litellm sent, so the split engages only when that id's price-map + row carries ``supports_bedrock_runtime_chat_completions_inline_reasoning`` (gpt-oss); a GPT 5.6 or Grok + answer that starts with a literal ```` tag streams as content. + """ def __init__( self, @@ -226,6 +232,8 @@ class BedrockRuntimeChatCompletionsStreamingHandler(OpenAIChatCompletionStreamin def chunk_parser(self, chunk: dict) -> ModelResponseStream: parsed: Final = super().chunk_parser(chunk) + if not bedrock_runtime_chat_completions_serves_reasoning_inline(parsed.model or ""): + return parsed for choice in parsed.choices: next_state, reasoning, content = _split_streamed_content( self._splitters.get(choice.index, ReasoningTagSplitter()), @@ -471,6 +479,8 @@ class AmazonBedrockRuntimeChatCompletionsConfig(OpenAILikeChatConfig): json_mode=json_mode, ) set_provider_response_headers_in_hidden_params(response, raw_response.headers) + if not bedrock_runtime_chat_completions_serves_reasoning_inline(model): + return response for choice in response.choices: if not isinstance(choice, Choices) or not isinstance(choice.message.content, str): continue diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index bc648d1f816..03d54662427 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -12,6 +12,7 @@ from litellm.caching.caching import InMemoryCache from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.litellm_core_utils.core_helpers import map_finish_reason from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.litellm_core_utils.optional_imports import ensure_optional_import from litellm.llms.anthropic.chat.handler import ( ModelResponseIterator as AnthropicModelResponseIterator, ) @@ -264,7 +265,7 @@ async def make_call( ) return completion_stream, response.headers - except BedrockError: + except (BedrockError, ImportError): raise except httpx.HTTPStatusError as err: error_code: Final = err.response.status_code @@ -355,7 +356,7 @@ def make_sync_call( ) return completion_stream, response.headers - except BedrockError: + except (BedrockError, ImportError): raise except httpx.HTTPStatusError as err: error_code: Final = err.response.status_code @@ -440,6 +441,7 @@ class _EventStreamTally: class AWSEventStreamDecoder: def __init__(self, model: str, json_mode: bool | None = False) -> None: + ensure_optional_import("botocore") from botocore.parsers import EventStreamJSONParser self.model = model diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 5b43ff8cae2..44f088510b7 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -27,6 +27,7 @@ from pydantic import ConfigDict, TypeAdapter, ValidationError import litellm from litellm import verbose_logger from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix +from litellm.litellm_core_utils.optional_imports import ensure_optional_import from litellm.llms.base_llm.anthropic_messages.transformation import ( BaseAnthropicMessagesConfig, ) @@ -909,6 +910,17 @@ def bedrock_model_is_openai_gpt(model: str) -> bool: return _openai_gpt_version(model) is not None +def bedrock_runtime_chat_completions_serves_reasoning_inline(model: str) -> bool: + """Whether AWS's native Chat Completions writes this model's reasoning inline in the answer text. + + Data-driven from the price-map ``supports_bedrock_runtime_chat_completions_inline_reasoning`` flag (gpt-oss). + A flagged model opens its answer with a ``...`` block instead of a + ``reasoning_content`` field, so litellm splits that block out for it and keeps every other model's + text as sent. + """ + return _bedrock_price_map_flag(model, "supports_bedrock_runtime_chat_completions_inline_reasoning") + + BEDROCK_CONVERSE_ONLY_REQUEST_KEYS: Final = frozenset( ( "guardrailConfig", @@ -1832,6 +1844,7 @@ class BedrockEventStreamDecoderBase: """ def __init__(self): + ensure_optional_import("botocore") from botocore.parsers import EventStreamJSONParser self.parser = EventStreamJSONParser() @@ -2052,11 +2065,9 @@ class CommonBatchFilesUtils: Returns: Tuple of (signed_headers, signed_data) """ - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + ensure_optional_import("botocore") + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest aws_region_name: Final = self._base_aws.get_aws_region_name(optional_params=optional_params, model="") credentials: Final = self._base_aws.resolve_credentials( diff --git a/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py index 2d02b152c61..ed61707f99e 100644 --- a/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py +++ b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py @@ -3,27 +3,100 @@ Bedrock Token Counter implementation using the CountTokens API. """ from collections.abc import Mapping, Sequence +from dataclasses import dataclass from typing import Any, Final +from pydantic import JsonValue +from typing_extensions import assert_never + from litellm._logging import verbose_logger from litellm.llms.base_llm.base_utils import BaseTokenCounter from litellm.llms.bedrock.common_utils import BedrockError, get_bedrock_base_model from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler -from litellm.types.utils import LlmProviders, TokenCountResponse +from litellm.llms.bedrock.count_tokens.mantle_handler import BedrockMantleCountTokensHandler +from litellm.llms.bedrock_mantle.common_utils import is_mantle_claude_model +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.types.utils import LiteLLMPydanticObjectBase, LlmProviders, TokenCountResponse + +RUNTIME_TOKENIZER_TYPE: Final = "bedrock_api" +MANTLE_TOKENIZER_TYPE: Final = "bedrock_mantle_api" + + +class _CountTokensReply(LiteLLMPydanticObjectBase): + input_tokens: int + + +@dataclass(frozen=True, slots=True) +class CountedTokens: + input_tokens: int + original_response: Mapping[str, JsonValue] + tokenizer_type: str + + +@dataclass(frozen=True, slots=True) +class CountTokensFailure: + status_code: int + message: str + tokenizer_type: str + + +CountTokensOutcome = CountedTokens | CountTokensFailure + + +def _runtime_rejected_claude_model(outcome: CountTokensOutcome, resolved_model: str) -> bool: + return ( + isinstance(outcome, CountTokensFailure) + and outcome.status_code == 400 + and is_mantle_claude_model(resolved_model) + ) class BedrockTokenCounter(BaseTokenCounter): - """Token counter implementation for AWS Bedrock provider using the CountTokens API.""" + """Token counter for AWS Bedrock: bedrock-runtime CountTokens first, and Anthropic's count_tokens + on bedrock-mantle for the Claude models bedrock-runtime answers 400 for""" + + def __init__( + self, + runtime_handler: BedrockCountTokensHandler | None = None, + mantle_handler: BedrockMantleCountTokensHandler | None = None, + client: AsyncHTTPHandler | None = None, + ) -> None: + self._runtime_handler: Final = runtime_handler or BedrockCountTokensHandler() + self._mantle_handler: Final = mantle_handler or BedrockMantleCountTokensHandler() + self._client: Final = client def should_use_token_counting_api( self, custom_llm_provider: str | None = None, ) -> bool: - """ - Returns True if we should use the Bedrock CountTokens API for token counting. - """ return custom_llm_provider == LlmProviders.BEDROCK.value + async def _count_with( + self, + handler: BedrockCountTokensHandler | BedrockMantleCountTokensHandler, + tokenizer_type: str, + request_data: dict[str, object], + litellm_params: dict[str, object], + resolved_model: str, + ) -> CountTokensOutcome: + try: + result: Final = await handler.handle_count_tokens_request( + request_data=request_data, + litellm_params=litellm_params, + resolved_model=resolved_model, + client=self._client, + ) + reply: Final = _CountTokensReply.model_validate(result) + except BedrockError as e: + verbose_logger.debug( + "%s CountTokens API error: status=%s, message=%s", tokenizer_type, e.status_code, e.message + ) + return CountTokensFailure(status_code=e.status_code, message=e.message, tokenizer_type=tokenizer_type) + except Exception as e: + verbose_logger.debug("Error calling %s CountTokens API: %s", tokenizer_type, e) + return CountTokensFailure(status_code=500, message=str(e), tokenizer_type=tokenizer_type) + return CountedTokens(input_tokens=reply.input_tokens, original_response=result, tokenizer_type=tokenizer_type) + async def count_tokens( self, model_to_use: str, @@ -34,81 +107,52 @@ class BedrockTokenCounter(BaseTokenCounter): tools: Sequence[Mapping[str, object]] | None = None, system: object | None = None, ) -> TokenCountResponse | None: - """ - Count tokens using AWS Bedrock's CountTokens API. - - This method calls the existing BedrockCountTokensHandler to make an API call - to Bedrock's token counting endpoint, bypassing the local tiktoken-based counting. - - Args: - model_to_use: The model identifier - messages: The messages to count tokens for - contents: Alternative content format (not used for Bedrock) - deployment: Deployment configuration containing litellm_params - request_model: The original request model name - - Returns: - TokenCountResponse with token count, or None if counting fails - """ if not messages: return None - deployment = deployment or {} - litellm_params: Final = deployment.get("litellm_params", {}) - - # Build request data in the format expected by BedrockCountTokensHandler + litellm_params: Final = (deployment or {}).get("litellm_params", {}) request_data: Final[dict[str, object]] = { "model": model_to_use, "messages": messages, + **({"tools": tools} if tools else {}), + **({"system": system} if system else {}), } - - if tools: - request_data["tools"] = tools - - if system: - request_data["system"] = system - - # Get the resolved model (strip prefixes like bedrock/, converse/, etc.) resolved_model: Final = get_bedrock_base_model(model_to_use) - try: - handler: Final = BedrockCountTokensHandler() - result: Final = await handler.handle_count_tokens_request( - request_data=request_data, - litellm_params=litellm_params, - resolved_model=resolved_model, + runtime: Final = await self._count_with( + self._runtime_handler, RUNTIME_TOKENIZER_TYPE, request_data, litellm_params, resolved_model + ) + outcome: Final = ( + await self._count_with( + self._mantle_handler, MANTLE_TOKENIZER_TYPE, request_data, litellm_params, resolved_model ) - - # Transform response to TokenCountResponse - if result is not None: + if _runtime_rejected_claude_model(runtime, resolved_model) + else runtime + ) + match outcome: + case CountedTokens(): return TokenCountResponse( - total_tokens=result.get("input_tokens", 0), + total_tokens=outcome.input_tokens, request_model=request_model, model_used=model_to_use, - tokenizer_type="bedrock_api", - original_response=result, + tokenizer_type=outcome.tokenizer_type, + original_response=dict(outcome.original_response), ) - except BedrockError as e: - verbose_logger.warning("Bedrock CountTokens API error: status=%s, message=%s", e.status_code, e.message) - return TokenCountResponse( - total_tokens=0, - request_model=request_model, - model_used=model_to_use, - tokenizer_type="bedrock_api", - error=True, - error_message=e.message, - status_code=e.status_code, - ) - except Exception as e: - verbose_logger.warning("Error calling Bedrock CountTokens API: %s", e) - return TokenCountResponse( - total_tokens=0, - request_model=request_model, - model_used=model_to_use, - tokenizer_type="bedrock_api", - error=True, - error_message=str(e), - status_code=500, - ) - - return None + case CountTokensFailure(): + verbose_logger.warning( + "%s CountTokens API error: status=%s, message=%s", + outcome.tokenizer_type, + outcome.status_code, + outcome.message, + ) + return TokenCountResponse( + total_tokens=0, + request_model=request_model, + model_used=model_to_use, + tokenizer_type=outcome.tokenizer_type, + error=True, + error_message=outcome.message, + status_code=outcome.status_code, + ) + case _: + assert_never(outcome) diff --git a/litellm/llms/bedrock/count_tokens/handler.py b/litellm/llms/bedrock/count_tokens/handler.py index d7fb510f057..be8243e48ef 100644 --- a/litellm/llms/bedrock/count_tokens/handler.py +++ b/litellm/llms/bedrock/count_tokens/handler.py @@ -125,7 +125,7 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig): raise except httpx.HTTPStatusError as e: # HTTP errors - preserve the actual status code - verbose_logger.error("HTTP error in CountTokens handler: %s", e) + verbose_logger.debug("HTTP error in CountTokens handler: %s", e) raise BedrockError( status_code=e.response.status_code, message=e.response.text, diff --git a/litellm/llms/bedrock/count_tokens/mantle_handler.py b/litellm/llms/bedrock/count_tokens/mantle_handler.py new file mode 100644 index 00000000000..23c3dc54c95 --- /dev/null +++ b/litellm/llms/bedrock/count_tokens/mantle_handler.py @@ -0,0 +1,99 @@ +from typing import Final + +import httpx +from pydantic import JsonValue, TypeAdapter +from typing_extensions import NotRequired, ReadOnly, TypedDict + +import litellm +from litellm._logging import verbose_logger +from litellm.llms.anthropic.count_tokens.transformation import AnthropicCountTokensConfig +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, run_aws_signing +from litellm.llms.bedrock.common_utils import BedrockError, build_mantle_messages_url +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client + +MANTLE_COUNT_TOKENS_SUFFIX: Final = "/count_tokens" +MANTLE_ANTHROPIC_VERSION: Final = "2023-06-01" + + +class MantleCountTokensRequest(TypedDict): + messages: ReadOnly[list[dict[str, JsonValue]]] + system: ReadOnly[NotRequired[JsonValue]] + tools: ReadOnly[NotRequired[list[dict[str, JsonValue]]]] + + +_COUNT_REQUEST: Final = TypeAdapter(MantleCountTokensRequest) +_COUNT_RESPONSE: Final = TypeAdapter(dict[str, JsonValue]) + + +class BedrockMantleCountTokensHandler(BaseAWSLLM): + """Counts tokens through Anthropic's count_tokens on the bedrock-mantle endpoint. + + Claude models that Bedrock offers only through cross-region inference answer 400 on + bedrock-runtime's CountTokens; AWS documents Mantle's /anthropic/v1/messages/count_tokens + as the way to count them, with the base model id and the deployment's AWS credentials + """ + + def __init__(self, anthropic_config: AnthropicCountTokensConfig | None = None) -> None: + super().__init__() + self._anthropic_config: Final = anthropic_config or AnthropicCountTokensConfig() + + def get_mantle_count_tokens_endpoint(self, aws_region_name: str) -> str: + messages_url: Final = build_mantle_messages_url( + api_base=None, aws_bedrock_runtime_endpoint=None, region=aws_region_name + ) + return f"{messages_url}{MANTLE_COUNT_TOKENS_SUFFIX}" + + async def handle_count_tokens_request( + self, + request_data: dict[str, object], + litellm_params: dict[str, object], + resolved_model: str, + client: AsyncHTTPHandler | None = None, + ) -> dict[str, JsonValue]: + try: + request: Final = _COUNT_REQUEST.validate_python(request_data) + aws_region_name: Final = self._get_aws_region_name( + optional_params=litellm_params, model=resolved_model, model_id=None + ) + endpoint_url: Final = self.get_mantle_count_tokens_endpoint(aws_region_name) + body: Final = self._anthropic_config.transform_request_to_count_tokens( + model=resolved_model, + messages=request["messages"], + tools=request.get("tools"), + system=request.get("system"), + ) + verbose_logger.debug("Making bedrock-mantle count_tokens request to: %s", endpoint_url) + api_key: Final = litellm_params.get("api_key") + signed_headers, signed_body = await run_aws_signing( + self._sign_request, + service_name="bedrock", + headers={"Content-Type": "application/json", "anthropic-version": MANTLE_ANTHROPIC_VERSION}, + optional_params=litellm_params, + request_data=body, + api_base=endpoint_url, + model=resolved_model, + api_key=api_key if isinstance(api_key, str) else None, + ) + async_client: Final = client or get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK) + response: Final = await async_client.post( + endpoint_url, headers=signed_headers, data=signed_body, timeout=30.0 + ) + if response.status_code != 200: + raise BedrockError( + status_code=response.status_code, + message=response.text, + headers=response.headers, + response=response, + ) + return _COUNT_RESPONSE.validate_json(response.content) + except BedrockError: + raise + except httpx.HTTPStatusError as e: + raise BedrockError( + status_code=e.response.status_code, + message=e.response.text, + headers=e.response.headers, + response=e.response, + ) + except Exception as e: + raise BedrockError(status_code=500, message=f"bedrock-mantle count_tokens processing error: {e}") diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py index b72d37e7e32..9a2cf378dcd 100644 --- a/litellm/llms/bedrock/files/transformation.py +++ b/litellm/llms/bedrock/files/transformation.py @@ -8,6 +8,7 @@ from collections.abc import Iterable, Mapping, MutableMapping, Sequence from contextlib import suppress from dataclasses import dataclass from datetime import datetime +from email.utils import parsedate_to_datetime from itertools import chain from types import MappingProxyType from typing import Any, Final, Literal, TypeAlias, TypedDict @@ -16,7 +17,7 @@ from urllib.parse import quote, unquote, urlencode import httpx from httpx import Headers, Response from openai.types.file_deleted import FileDeleted -from pydantic import ConfigDict, Field +from pydantic import ConfigDict, Field, TypeAdapter from typing_extensions import ReadOnly from litellm._logging import verbose_logger @@ -70,12 +71,71 @@ from ..common_utils import ( ) S3_SIGNED_REQUEST_HEADERS_PARAM: Final = "_s3_signed_request_headers" +S3_RETRIEVE_FILE_ID_PARAM: Final = "_s3_retrieve_file_id" +S3_RETRIEVE_FILE_KEY_PARAM: Final = "_s3_retrieve_file_key" +S3_RETRIEVE_FILE_RELATIVE_KEY_PARAM: Final = "_s3_retrieve_file_relative_key" +_S3_SIGNED_REQUEST_HEADERS_ADAPTER: Final = TypeAdapter( + Mapping[str, str], + config=ConfigDict(strict=True), +) LIST_FILES_PURPOSE_PARAM: Final = "_s3_list_files_purpose" LIST_FILES_LOCATION_PARAM: Final = "_s3_list_files_location" +def _is_empty_s3_object_range_error(raw_response: Response) -> bool: + if raw_response.status_code != 416: + return False + if raw_response.headers.get("Content-Range") == "bytes */0": + return True + try: + error_xml: Final = ET.fromstring(raw_response.content) + except ET.ParseError: + return False + return error_xml.findtext("ActualObjectSize") == "0" + + +def _retrieved_s3_file_size(raw_response: Response) -> int: + status_code: Final = raw_response.status_code + if _is_empty_s3_object_range_error(raw_response): + return 0 + if status_code == 206: + content_range: Final = raw_response.headers.get("Content-Range", "") + range_parts: Final = content_range.removeprefix("bytes 0-0/") + if content_range.startswith("bytes 0-0/") and range_parts.isdigit(): + return int(range_parts) + raise BedrockError( + status_code=status_code, + message=f"Invalid S3 Content-Range header: {content_range}", + headers=raw_response.headers, + response=raw_response, + ) + if status_code == 200: + content_length: Final = raw_response.headers.get("Content-Length", "") + if content_length.isdigit(): + return int(content_length) + raise BedrockError( + status_code=status_code, + message=f"Invalid S3 Content-Length header: {content_length}", + headers=raw_response.headers, + response=raw_response, + ) + if status_code >= 400: + raise BedrockError( + status_code=status_code, + message=raw_response.text, + headers=raw_response.headers, + response=raw_response, + ) + raise BedrockError( + status_code=status_code, + message=f"S3 file retrieval returned HTTP {status_code}", + headers=raw_response.headers, + response=raw_response, + ) + + class _S3DeleteContext(LiteLLMBaseModel): file_id: str = Field(min_length=1) @@ -280,6 +340,25 @@ def _resolve_managed_s3_object(file_id: str, litellm_params: Mapping[str, object raise _rejected_file_id(reason) from reason +def _relative_s3_object_key( + bucket_name: str, + object_key: str, + litellm_params: Mapping[str, object], +) -> str: + configured_bucket_prefixes: Final = tuple( + split_configured_cloud_bucket_name(configured_bucket_name) + for configured_bucket_name in get_configured_s3_bucket_names(litellm_params) + ) + matching_prefixes: Final = tuple( + configured_prefix + for configured_bucket, configured_prefix in configured_bucket_prefixes + if configured_bucket == bucket_name + and (not configured_prefix or object_key.startswith(f"{configured_prefix}/")) + ) + configured_prefix: Final = max(matching_prefixes, key=len, default="") + return object_key[len(configured_prefix) + 1 :] if configured_prefix else object_key + + _ANY_MANAGED_LISTING_PREFIX: Final = os.path.commonprefix(BEDROCK_MANAGED_S3_PREFIXES) _MANAGED_LISTING_PREFIX_BY_PURPOSE: Final = MappingProxyType( { @@ -398,6 +477,9 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): def custom_llm_provider(self) -> LlmProviders: return LlmProviders.BEDROCK + def is_retrieve_file_response_successful(self, response: httpx.Response) -> bool: + return not httpx.codes.is_error(response.status_code) or _is_empty_s3_object_range_error(response) + @property def file_upload_http_method(self) -> str: """ @@ -1276,18 +1358,59 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): def transform_retrieve_file_request( self, file_id: str, - optional_params: dict, - litellm_params: dict, - ) -> tuple[str, dict]: - raise NotImplementedError("BedrockFilesConfig does not support file retrieval") + optional_params: Mapping[str, object], + litellm_params: MutableMapping[str, object], + ) -> tuple[str, dict[str, str]]: + """Prepare a ranged S3 GET for file retrieval.""" + bucket_name, object_key = _resolve_managed_s3_object(file_id=file_id, litellm_params=litellm_params) + relative_key: Final = _relative_s3_object_key( + bucket_name=bucket_name, + object_key=object_key, + litellm_params=litellm_params, + ) + url, params = self._transform_s3_file_request( + file_id=file_id, + method="GET", + optional_params=optional_params, + litellm_params=litellm_params, + ) + signed_headers_object: Final = litellm_params.get(S3_SIGNED_REQUEST_HEADERS_PARAM) + if not isinstance(signed_headers_object, Mapping): + raise TypeError("S3 request signing did not produce request headers") + signed_headers: Final = _S3_SIGNED_REQUEST_HEADERS_ADAPTER.validate_python(signed_headers_object) + range_headers: Final = MappingProxyType({**signed_headers, "Range": "bytes=0-0"}) + litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM] = range_headers # rebind-ok: handed to validate_environment + litellm_params[S3_RETRIEVE_FILE_ID_PARAM] = file_id # rebind-ok: required by response transform + litellm_params[S3_RETRIEVE_FILE_KEY_PARAM] = object_key # rebind-ok: required by response transform + litellm_params[S3_RETRIEVE_FILE_RELATIVE_KEY_PARAM] = relative_key # rebind-ok: required by response transform + return url, params def transform_retrieve_file_response( self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj, - litellm_params: dict, + litellm_params: Mapping[str, object], ) -> OpenAIFileObject: - raise NotImplementedError("BedrockFilesConfig does not support file retrieval") + """Build file metadata, accepting 416 only when S3 proves the object is empty.""" + file_id: Final = litellm_params.get(S3_RETRIEVE_FILE_ID_PARAM) + object_key: Final = litellm_params.get(S3_RETRIEVE_FILE_KEY_PARAM) + relative_key: Final = litellm_params.get(S3_RETRIEVE_FILE_RELATIVE_KEY_PARAM) + if not isinstance(file_id, str) or not isinstance(object_key, str) or not isinstance(relative_key, str): + raise TypeError("S3 retrieve response is missing request context") + + file_size: Final = _retrieved_s3_file_size(raw_response) + + last_modified: Final = raw_response.headers.get("Last-Modified", "") + created_at: Final = int(parsedate_to_datetime(last_modified).timestamp()) if last_modified else 0 + return OpenAIFileObject( + id=file_id, + bytes=file_size, + created_at=created_at, + filename=posixpath.basename(object_key), + object="file", + purpose="batch_output" if relative_key.startswith(BEDROCK_MANAGED_S3_OUTPUT_PREFIX) else "batch", + status="processed", + ) def transform_delete_file_request( self, diff --git a/litellm/llms/bedrock/image_edit/handler.py b/litellm/llms/bedrock/image_edit/handler.py index 76ec537b68d..53bb332f511 100644 --- a/litellm/llms/bedrock/image_edit/handler.py +++ b/litellm/llms/bedrock/image_edit/handler.py @@ -27,6 +27,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_httpx_client, ) from litellm.types.llms.base import LiteLLMBaseModel +from litellm.types.llms.bedrock import BearerPreparedRequest from litellm.types.utils import ImageResponse from ..base_aws_llm import BaseAWSLLM, bedrock_bearer_token @@ -44,7 +45,7 @@ class BedrockImageEditPreparedRequest(LiteLLMBaseModel): """ endpoint_url: str - prepped: AWSPreparedRequest + prepped: AWSPreparedRequest | BearerPreparedRequest body: bytes data: dict diff --git a/litellm/llms/bedrock/image_generation/image_handler.py b/litellm/llms/bedrock/image_generation/image_handler.py index 018e92fd79f..e81dfe76942 100644 --- a/litellm/llms/bedrock/image_generation/image_handler.py +++ b/litellm/llms/bedrock/image_generation/image_handler.py @@ -27,6 +27,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_httpx_client, ) from litellm.types.llms.base import LiteLLMBaseModel +from litellm.types.llms.bedrock import BearerPreparedRequest from litellm.types.utils import ImageResponse from ..base_aws_llm import BaseAWSLLM, bedrock_bearer_token @@ -44,7 +45,7 @@ class BedrockImagePreparedRequest(LiteLLMBaseModel): """ endpoint_url: str - prepped: AWSPreparedRequest + prepped: AWSPreparedRequest | BearerPreparedRequest body: bytes data: dict diff --git a/litellm/llms/bedrock/passthrough/transformation.py b/litellm/llms/bedrock/passthrough/transformation.py index efad2e73e50..df4845828f5 100644 --- a/litellm/llms/bedrock/passthrough/transformation.py +++ b/litellm/llms/bedrock/passthrough/transformation.py @@ -22,6 +22,13 @@ if TYPE_CHECKING: from litellm.types.utils import CostResponseTypes +BEDROCK_STREAMING_ACTIONS: Final = frozenset({"invoke-with-response-stream", "converse-stream"}) + + +def is_bedrock_streaming_endpoint(endpoint: str) -> bool: + return endpoint.partition("?")[0].rstrip("/").rsplit("/", 1)[-1] in BEDROCK_STREAMING_ACTIONS + + _TEXT_ONLY_DELTA_FIELDS: Final = frozenset({"content", "role"}) diff --git a/litellm/llms/bedrock_mantle/common_utils.py b/litellm/llms/bedrock_mantle/common_utils.py index 8f9cf5c4f6b..75e34cf3ae5 100644 --- a/litellm/llms/bedrock_mantle/common_utils.py +++ b/litellm/llms/bedrock_mantle/common_utils.py @@ -14,15 +14,10 @@ global state. import re from collections.abc import Mapping +from functools import partial from typing import Final, Literal -from botocore.exceptions import ( - CredentialRetrievalError, - NoCredentialsError, - PartialCredentialsError, - ProfileNotFound, -) - +from litellm.litellm_core_utils.optional_imports import ensure_optional_import from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, SignsRequestsWithAWS from litellm.llms.bedrock.common_utils import AmazonBedrockGlobalConfig from litellm.secret_managers.main import get_secret_str @@ -72,7 +67,7 @@ class BedrockMantleAuthMixin(SignsRequestsWithAWS): return resolve_mantle_bearer_token(api_key) @staticmethod - def _resolve_region(params: dict) -> str: + def _resolve_region(params: Mapping[str, object]) -> str: return resolve_mantle_region(params) def sign_request( @@ -87,32 +82,39 @@ class BedrockMantleAuthMixin(SignsRequestsWithAWS): fake_stream: bool | None = None, ) -> tuple[dict, bytes | None]: bearer: Final = self._resolve_bearer_token(api_key) - if not bearer: - # Pin the credential-scope region to the region of the actual signing URL - # so the SigV4 scope and URL host can never disagree, even when a stale - # api_base and aws_region_name point at different regions. - host_match: Final = MANTLE_HOST_RE.match(api_base.rstrip("/")) - optional_params = { - **optional_params, - "aws_region_name": ( - host_match.group(1) - if host_match - else self._resolve_region({**optional_params, "api_base": api_base}) - ), - } - headers = {k: v for k, v in headers.items() if k.lower() != "authorization"} + sign: Final[partial[tuple[dict[str, str | bytes], bytes | None]]] = partial( + self._aws_signer._sign_request, + service_name="bedrock", + request_data=request_data, + api_base=api_base, + api_key=bearer, + model=model, + stream=stream, + fake_stream=fake_stream, + ) + if bearer: + return sign(headers=headers, optional_params=optional_params) + ensure_optional_import("botocore") + from botocore.exceptions import ( + CredentialRetrievalError, + NoCredentialsError, + PartialCredentialsError, + ProfileNotFound, + ) + + # Pin the credential-scope region to the region of the actual signing URL + # so the SigV4 scope and URL host can never disagree, even when a stale + # api_base and aws_region_name point at different regions. + host_match: Final = MANTLE_HOST_RE.match(api_base.rstrip("/")) + optional_params = { + **optional_params, + "aws_region_name": ( + host_match.group(1) if host_match else self._resolve_region({**optional_params, "api_base": api_base}) + ), + } + headers = {k: v for k, v in headers.items() if k.lower() != "authorization"} try: - return self._aws_signer._sign_request( - service_name="bedrock", - headers=headers, - optional_params=optional_params, - request_data=request_data, - api_base=api_base, - api_key=bearer, - model=model, - stream=stream, - fake_stream=fake_stream, - ) + return sign(headers=headers, optional_params=optional_params) except ( NoCredentialsError, PartialCredentialsError, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 6298ee6ffa7..895305d6b6f 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -4756,7 +4756,8 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) - self._raise_for_provider_error_status(response=response, provider_config=provider_config) + if not provider_config.is_retrieve_file_response_successful(response): + self._raise_for_provider_error_status(response=response, provider_config=provider_config) return provider_config.transform_retrieve_file_response( raw_response=response, logging_obj=logging_obj, @@ -4814,7 +4815,8 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) - self._raise_for_provider_error_status(response=response, provider_config=provider_config) + if not provider_config.is_retrieve_file_response_successful(response): + self._raise_for_provider_error_status(response=response, provider_config=provider_config) return provider_config.transform_retrieve_file_response( raw_response=response, logging_obj=logging_obj, diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index 485341c59af..d2b24489f20 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -43,7 +43,9 @@ from ..common_utils import ( FIREROUTER, FireworksAIException, FireworksAIMixin, + get_fireworks_forwarded_user_id, resolve_fireworks_resource_name, + without_caller_user, ) if TYPE_CHECKING: @@ -672,13 +674,24 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): **stream_options, "include_usage": True, } - return super().transform_request( + request: Final = super().transform_request( model=resolved_model, messages=messages, optional_params=optional_params, litellm_params=litellm_params, headers=headers, ) + forwarded_user_id: Final = get_fireworks_forwarded_user_id(litellm_params) + return request if forwarded_user_id is None else {**request, "user": forwarded_user_id} + + def transform_extra_body( + self, + extra_body: Mapping[str, object], + request: Mapping[str, object], + model: str, + litellm_params: Mapping[str, object], + ) -> Mapping[str, object]: + return without_caller_user(extra_body, get_fireworks_forwarded_user_id(litellm_params)) def _handle_message_content_with_tool_calls( self, diff --git a/litellm/llms/fireworks_ai/common_utils.py b/litellm/llms/fireworks_ai/common_utils.py index 17fadf7ae0f..10314d3fac4 100644 --- a/litellm/llms/fireworks_ai/common_utils.py +++ b/litellm/llms/fireworks_ai/common_utils.py @@ -49,6 +49,33 @@ def with_fireworks_session_affinity( return MappingProxyType({**headers, "x-session-affinity": session_id}) +FIREWORKS_FORWARD_USER_ID_PARAM: Final = "fireworks_forward_user_id" + + +def _authenticated_user_id(metadata: object) -> str | None: + user_id: Final = metadata.get("user_api_key_user_id") if isinstance(metadata, Mapping) else None + return user_id if isinstance(user_id, str) and user_id else None + + +def get_fireworks_forwarded_user_id(litellm_params: Mapping[str, object]) -> str | None: + if litellm_params.get(FIREWORKS_FORWARD_USER_ID_PARAM) is not True: + return None + return next( + ( + user_id + for key in ("metadata", "litellm_metadata") + if (user_id := _authenticated_user_id(litellm_params.get(key))) is not None + ), + None, + ) + + +def without_caller_user(extra_body: Mapping[str, object], forwarded_user_id: str | None) -> Mapping[str, object]: + if forwarded_user_id is None: + return extra_body + return MappingProxyType({key: value for key, value in extra_body.items() if key != "user"}) + + def resolve_fireworks_api_key(api_key: str | None) -> str | None: return api_key or ( get_secret_str("FIREWORKS_API_KEY") diff --git a/litellm/llms/fireworks_ai/responses/transformation.py b/litellm/llms/fireworks_ai/responses/transformation.py index 30e6f053cc1..3adc5a48490 100644 --- a/litellm/llms/fireworks_ai/responses/transformation.py +++ b/litellm/llms/fireworks_ai/responses/transformation.py @@ -7,9 +7,12 @@ import httpx from openai.types.responses import EasyInputMessageParam, ResponseInputContentParam, ResponseInputItemParam from litellm.llms.fireworks_ai.common_utils import ( + FIREWORKS_FORWARD_USER_ID_PARAM, + get_fireworks_forwarded_user_id, resolve_fireworks_api_key, resolve_fireworks_resource_name, with_fireworks_session_affinity, + without_caller_user, ) from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.secret_managers.main import get_secret_str @@ -31,6 +34,16 @@ def _session_params(litellm_params: GenericLiteLLMParams) -> Mapping[str, object ) +def _forwarded_user_params(litellm_params: GenericLiteLLMParams) -> Mapping[str, object]: + extras: Final[Mapping[str, object]] = litellm_params.model_extra or MappingProxyType({}) + return MappingProxyType( + { + FIREWORKS_FORWARD_USER_ID_PARAM: litellm_params.fireworks_forward_user_id, + "litellm_metadata": extras.get("litellm_metadata"), + } + ) + + _INSTRUCTION_ROLES: Final = frozenset({"system", "developer"}) @@ -163,13 +176,24 @@ class FireworksAIResponsesAPIConfig(OpenAIResponsesAPIConfig): *instruction_entries, ) } - return super().transform_responses_api_request( + request: Final = super().transform_responses_api_request( model=resolve_fireworks_resource_name(model), input=folded_input, response_api_optional_request_params=folded_params, litellm_params=litellm_params, headers=headers, ) + forwarded_user_id: Final = get_fireworks_forwarded_user_id(_forwarded_user_params(litellm_params)) + return request if forwarded_user_id is None else {**request, "user": forwarded_user_id} + + def transform_extra_body( + self, + extra_body: Mapping[str, object], + request: Mapping[str, object], + model: str, + litellm_params: GenericLiteLLMParams, + ) -> Mapping[str, object]: + return without_caller_user(extra_body, get_fireworks_forwarded_user_id(_forwarded_user_params(litellm_params))) def transform_delete_response_api_response( self, diff --git a/litellm/llms/pass_through/guardrail_translation/handler.py b/litellm/llms/pass_through/guardrail_translation/handler.py index 1f295a6e656..fd58ca2ddbb 100644 --- a/litellm/llms/pass_through/guardrail_translation/handler.py +++ b/litellm/llms/pass_through/guardrail_translation/handler.py @@ -10,6 +10,7 @@ from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, Optional from litellm._logging import verbose_proxy_logger +from litellm.integrations.custom_guardrail import without_server_streaming_classification from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.proxy._types import PassThroughGuardrailSettings from litellm.types.utils import GenericGuardrailAPIInputs @@ -80,7 +81,9 @@ class PassThroughEndpointHandler(BaseTranslation): from litellm.litellm_core_utils.safe_json_dumps import safe_dumps payload_to_check: Final = { - k: v for k, v in data.items() if not k.startswith("_") and k not in ("metadata", "litellm_logging_obj") + k: v + for k, v in without_server_streaming_classification(data).items() + if not k.startswith("_") and k not in ("metadata", "litellm_logging_obj") } verbose_proxy_logger.debug("PassThroughEndpointHandler: Using full payload for guardrail") return safe_dumps(payload_to_check) diff --git a/litellm/llms/sagemaker/chat/handler.py b/litellm/llms/sagemaker/chat/handler.py index 10be9ef384c..3e671ce5d07 100644 --- a/litellm/llms/sagemaker/chat/handler.py +++ b/litellm/llms/sagemaker/chat/handler.py @@ -6,6 +6,7 @@ from typing import Final import httpx from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix +from litellm.litellm_core_utils.optional_imports import ensure_optional_import from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, pop_aws_auth_params from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.utils import ModelResponse, get_secret @@ -19,10 +20,9 @@ class SagemakerChatHandler(BaseAWSLLM): self, optional_params: dict, ): - try: - from botocore.credentials import Credentials - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + ensure_optional_import("botocore") + from botocore.credentials import Credentials + auth_params: Final = pop_aws_auth_params(optional_params) aws_region_name = optional_params.pop("aws_region_name", None) optional_params.pop("aws_bedrock_runtime_endpoint", None) @@ -54,11 +54,9 @@ class SagemakerChatHandler(BaseAWSLLM): aws_region_name: str, extra_headers: dict | None = None, ): - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + ensure_optional_import("botocore") + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest sigv4: Final = SigV4Auth(credentials, "sagemaker", aws_region_name) dns_suffix: Final = get_aws_dns_suffix(aws_region_name) diff --git a/litellm/llms/sagemaker/completion/handler.py b/litellm/llms/sagemaker/completion/handler.py index 46d204cf385..069518021fa 100644 --- a/litellm/llms/sagemaker/completion/handler.py +++ b/litellm/llms/sagemaker/completion/handler.py @@ -10,6 +10,7 @@ from litellm._logging import verbose_logger from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.optional_imports import ensure_optional_import from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, pop_aws_auth_params from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, @@ -42,10 +43,9 @@ class SagemakerLLM(BaseAWSLLM): self, optional_params: dict, ): - try: - from botocore.credentials import Credentials - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + ensure_optional_import("botocore") + from botocore.credentials import Credentials + auth_params: Final = pop_aws_auth_params(optional_params) aws_region_name = optional_params.pop("aws_region_name", None) optional_params.pop("aws_bedrock_runtime_endpoint", None) @@ -79,11 +79,9 @@ class SagemakerLLM(BaseAWSLLM): aws_region_name: str, extra_headers: dict | None = None, ): - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + ensure_optional_import("botocore") + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest sigv4: Final = SigV4Auth(credentials, "sagemaker", aws_region_name) dns_suffix: Final = get_aws_dns_suffix(aws_region_name) diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 8c77413bace..6b759065706 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -25,6 +25,22 @@ from litellm.types.llms.vertex_ai import ( from litellm.types.utils import TokenCountResponse from litellm.utils import supports_response_schema, supports_system_messages +GEMINI_VIDEO_METADATA_KEYS: Final = MappingProxyType( + {"fps": "fps", "start_offset": "startOffset", "end_offset": "endOffset"} +) + + +GEMINI_FILES_API_URI_PREFIX: Final = "https://generativelanguage.googleapis.com/v1beta/files/" + + +def gemini_video_metadata_from_openai(video_metadata: Mapping[str, object]) -> dict[str, object]: + return { + gemini_key: video_metadata[openai_key] + for openai_key, gemini_key in GEMINI_VIDEO_METADATA_KEYS.items() + if openai_key in video_metadata + } + + VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS: Final = frozenset( { "audio", diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index da5b974f186..09ec2ba6ee4 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -67,7 +67,7 @@ from litellm.types.llms.openai import ( OpenAIFilesPurpose, PathLike, ) -from litellm.types.llms.vertex_ai import GcsBucketResponse, GeminiEmbeddingInput +from litellm.types.llms.vertex_ai import GcsBucketResponse, GeminiEmbeddingElement, GeminiEmbeddingInput from litellm.types.utils import ( Embedding, EmbeddingResponse, @@ -555,18 +555,27 @@ def _is_responses_batch_entry(openai_entry: Mapping[str, object]) -> bool: return path == "responses" or path.endswith("/responses") +def _own_embedding_input( + element: GeminiEmbeddingElement | list[str] | list[GeminiEmbeddingElement], +) -> GeminiEmbeddingInput: + if isinstance(element, (str, list)): + return element + file_block_alone: Final[list[GeminiEmbeddingElement]] = [element] + return file_block_alone + + def _openai_embedding_input_elements( embedding_input: GeminiEmbeddingInput, -) -> tuple[str | list[str], ...]: +) -> tuple[GeminiEmbeddingInput, ...]: """ Split an OpenAI `input` into the elements that each get their own embedding. - A string is one embedding, a flat array is one embedding per element, and a nested - array is one combined embedding per inner array, matching the online - `batchEmbedContents` path. + A string or a file content block is one embedding, a flat array is one embedding + per element, and a nested array is one combined embedding per inner array, + matching the online `batchEmbedContents` path. """ if isinstance(embedding_input, list): - return tuple(embedding_input) + return tuple(_own_embedding_input(element) for element in embedding_input) return (embedding_input,) diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 8749b81d7f1..15e5e440b91 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -57,7 +57,9 @@ from litellm.types.llms.vertex_ai import ( from litellm.types.utils import GenericImageParsingChunk, LlmProviders from ..common_utils import ( + GEMINI_FILES_API_URI_PREFIX, check_text_in_content, + gemini_video_metadata_from_openai, get_supports_response_schema, get_supports_system_message, ) @@ -69,7 +71,6 @@ _GCS_METADATA_VERTEX_BASE: object | None = None # Shared sync client for GCS JSON API metadata reads so proxy/SSL settings # from litellm's HTTP stack apply (see Greptile review on PR #27278). _GCS_METADATA_HTTP_HANDLER: HTTPHandler | None = None -GEMINI_FILES_API_URI_PREFIX: Final = "https://generativelanguage.googleapis.com/v1beta/files/" _GEMINI_MIME_TYPE_ALIASES: Final[dict[str, str]] = { "image/jpg": "image/jpeg", } @@ -197,13 +198,7 @@ def _apply_gemini_metadata( part_dict["media_resolution"] = media_resolution_enum if video_metadata is not None: - gemini_video_metadata: Final = {} - if "fps" in video_metadata: - gemini_video_metadata["fps"] = video_metadata["fps"] - if "start_offset" in video_metadata: - gemini_video_metadata["startOffset"] = video_metadata["start_offset"] - if "end_offset" in video_metadata: - gemini_video_metadata["endOffset"] = video_metadata["end_offset"] + gemini_video_metadata: Final = gemini_video_metadata_from_openai(video_metadata) if gemini_video_metadata: part_dict["video_metadata"] = gemini_video_metadata diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py index d15cf89cf47..0a800497de1 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py @@ -22,6 +22,8 @@ from litellm.types.utils import EmbeddingResponse from ..gemini.vertex_and_google_ai_studio_gemini import VertexLLM from .batch_embed_content_transformation import ( + file_reference_name, + flatten_media_sources, is_file_reference, process_embed_content_response, process_response, @@ -38,11 +40,8 @@ class GoogleBatchEmbeddings(VertexLLM): def _flatten_and_detect_file_refs( input: GeminiEmbeddingInput, ) -> tuple[list[str], bool]: - """Flatten nested input lists and detect file references.""" - input_list: Final = [input] if isinstance(input, str) else input - flat_elements: Final = [ - e for item in input_list for e in (item if isinstance(item, list) else [item]) if isinstance(e, str) - ] + """Flatten nested input lists and file content blocks into their sources and detect file references.""" + flat_elements: Final = list(flatten_media_sources(input)) has_file_refs: Final = any(is_file_reference(e) for e in flat_elements) return flat_elements, has_file_refs @@ -68,7 +67,7 @@ class GoogleBatchEmbeddings(VertexLLM): for element in input_list: if isinstance(element, str) and is_file_reference(element): - url = f"https://generativelanguage.googleapis.com/v1beta/{element}" + url = f"https://generativelanguage.googleapis.com/v1beta/{file_reference_name(element)}" headers = {"x-goog-api-key": api_key} response = sync_handler.get(url=url, headers=headers) @@ -105,7 +104,7 @@ class GoogleBatchEmbeddings(VertexLLM): for element in input_list: if isinstance(element, str) and is_file_reference(element): - url = f"https://generativelanguage.googleapis.com/v1beta/{element}" + url = f"https://generativelanguage.googleapis.com/v1beta/{file_reference_name(element)}" headers = {"x-goog-api-key": api_key} response = await async_handler.get(url=url, headers=headers) @@ -139,6 +138,7 @@ class GoogleBatchEmbeddings(VertexLLM): timeout=300, client=None, extra_headers: dict | None = None, + drop_params: bool = False, ) -> EmbeddingResponse: _auth_header, vertex_project = self._ensure_access_token( credentials=vertex_credentials, @@ -207,22 +207,26 @@ class GoogleBatchEmbeddings(VertexLLM): api_key=api_key, optional_params=optional_params, logging_obj=logging_obj, + drop_params=drop_params, ) ### TRANSFORMATION (sync path) ### request_data: VertexAIBatchEmbeddingsRequestBody | dict[str, object] + flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input) if use_embed_content: resolved_files = {} if api_key: - resolved_files = self._resolve_file_references(input=input, api_key=api_key, sync_handler=sync_handler) + resolved_files = self._resolve_file_references( + input=flat_elements, api_key=api_key, sync_handler=sync_handler + ) request_data = transform_openai_input_gemini_embed_content( input=input, model=model, optional_params=optional_params, resolved_files=resolved_files, + drop_params=drop_params, ) else: - flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input) if has_file_refs and not api_key: raise ValueError( "An API key is required to resolve Gemini file references (files/...). " @@ -238,6 +242,7 @@ class GoogleBatchEmbeddings(VertexLLM): model=model, optional_params=optional_params, resolved_files=resolved_files, + drop_params=drop_params, ) ## LOGGING @@ -294,6 +299,7 @@ class GoogleBatchEmbeddings(VertexLLM): api_key: str | None = None, optional_params: dict | None = None, logging_obj: "LiteLLMLoggingObj | None" = None, + drop_params: bool = False, ) -> EmbeddingResponse: if client is None: _params: Final = {} @@ -312,20 +318,21 @@ class GoogleBatchEmbeddings(VertexLLM): async_handler = client ### TRANSFORMATION (async path) ### + flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input) if use_embed_content: resolved_files = {} if api_key: resolved_files = await self._async_resolve_file_references( - input=input, api_key=api_key, async_handler=async_handler + input=flat_elements, api_key=api_key, async_handler=async_handler ) data = transform_openai_input_gemini_embed_content( input=input, model=model, optional_params=optional_params or {}, resolved_files=resolved_files, + drop_params=drop_params, ) else: - flat_elements, has_file_refs = self._flatten_and_detect_file_refs(input) if has_file_refs and not api_key: raise ValueError( "An API key is required to resolve Gemini file references (files/...). " @@ -341,6 +348,7 @@ class GoogleBatchEmbeddings(VertexLLM): model=model, optional_params=optional_params or {}, resolved_files=resolved_files, + drop_params=drop_params, ) ## LOGGING diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py index 9ee2b71d30c..96a785c06c9 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py @@ -4,22 +4,27 @@ Transformation logic from OpenAI /v1/embeddings format to Google AI Studio /batc Why separate file? Make it easy to see how transformation works """ -from collections.abc import Mapping, Sequence -from typing import Final +from collections.abc import Iterator, Mapping, Sequence +from typing import Annotated, Final, Literal, cast -from pydantic import TypeAdapter, ValidationError +from pydantic import ConfigDict, Field, TypeAdapter, ValidationError +from litellm.exceptions import BadRequestError +from litellm.llms.vertex_ai.common_utils import GEMINI_FILES_API_URI_PREFIX, gemini_video_metadata_from_openai +from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.llms.vertex_ai import ( BlobType, ContentType, EmbedContentRequest, FileDataType, + GeminiEmbeddingElement, GeminiEmbeddingInput, PartType, PromptTokensDetails, UsageMetadata, VertexAIBatchEmbeddingsRequestBody, VertexAIBatchEmbeddingsResponseObject, + VideoMetadataType, ) from litellm.types.utils import ( Embedding, @@ -40,9 +45,16 @@ SUPPORTED_EMBEDDING_MIME_TYPES: Final = { } +_GEMINI_API_V1BETA: Final = GEMINI_FILES_API_URI_PREFIX.removesuffix("files/") + + def is_file_reference(s: str) -> bool: - """Check if string is a Gemini file reference (files/...).""" - return isinstance(s, str) and s.startswith("files/") + """A `files/...` name or the Files API URI that `/v1/files` returns as the file id.""" + return isinstance(s, str) and (s.startswith("files/") or s.startswith(GEMINI_FILES_API_URI_PREFIX)) + + +def file_reference_name(reference: str) -> str: + return reference.removeprefix(_GEMINI_API_V1BETA) _is_file_reference = is_file_reference @@ -88,12 +100,13 @@ def _infer_mime_type_from_gcs_url(gcs_url: str) -> str: ) -def _parse_data_url(data_url: str) -> tuple[str, str]: +def _parse_data_url(data_url: str, mime_type_override: str | None = None) -> tuple[str, str]: """ Parse a data URL to extract the media type and base64 data. Args: data_url: Data URL in format: data:image/jpeg;base64,/9j/4AAQ... + mime_type_override: An explicit media type that replaces the declared one, skipping the allowlist Returns: tuple: (media_type, base64_data) @@ -110,50 +123,225 @@ def _parse_data_url(data_url: str) -> tuple[str, str]: raise ValueError(f"Invalid data URL format (missing comma): {data_url[:50]}...") metadata, base64_data = data_url.split(",", 1) + declared_media_type: Final = metadata[5:].split(";")[0] - metadata = metadata[5:] + if mime_type_override is not None: + return mime_type_override, base64_data - if ";" in metadata: - media_type = metadata.split(";")[0] - else: - media_type = metadata - - if media_type not in SUPPORTED_EMBEDDING_MIME_TYPES: + if declared_media_type not in SUPPORTED_EMBEDDING_MIME_TYPES: raise ValueError( - f"Unsupported MIME type for embedding: {media_type}. " + f"Unsupported MIME type for embedding: {declared_media_type}. " f"Supported types: {', '.join(sorted(SUPPORTED_EMBEDDING_MIME_TYPES))}" ) - return media_type, base64_data + return declared_media_type, base64_data + + +def _is_data_url(s: str) -> bool: + return s.startswith("data:") and ";base64," in s + + +class _EmbeddingVideoMetadata(LiteLLMBaseModel): + model_config = ConfigDict(extra="forbid", strict=True) + + fps: float | None = None + start_offset: str | None = None + end_offset: str | None = None + + +class _EmbeddingFile(LiteLLMBaseModel): + model_config = ConfigDict(extra="forbid") + + file_id: str | None = None + file_data: str | None = None + filename: str | None = None + format: Annotated[str, Field(min_length=1)] | None = None + video_metadata: _EmbeddingVideoMetadata | None = None + + +class _EmbeddingFileBlock(LiteLLMBaseModel): + model_config = ConfigDict(extra="forbid") + + type: Literal["file"] + file: _EmbeddingFile + + +_file_block_adapter: Final = TypeAdapter(_EmbeddingFileBlock) +_video_metadata_adapter: Final = TypeAdapter(VideoMetadataType) +_input_shape_adapter: Final[TypeAdapter[str | list[object]]] = TypeAdapter(str | list[object]) +_mapping_adapter: Final = TypeAdapter(dict[str, object]) +_BLOCK_FIELDS: Final = frozenset(_EmbeddingFileBlock.model_fields) +_FILE_FIELDS: Final = frozenset(_EmbeddingFile.model_fields) +_VIDEO_METADATA_FIELDS: Final = frozenset(_EmbeddingVideoMetadata.model_fields) +_FILE_SOURCE_FORMS: Final = "a data: URI, a gs:// URL, a files/ reference, or a Gemini Files API URI" + + +def _invalid_input(message: str) -> BadRequestError: + return BadRequestError(message=message, model=None, llm_provider="gemini") + + +def _validation_error_summary(error: ValidationError) -> str: + return "; ".join(f"{'.'.join(str(loc) for loc in detail['loc'])}: {detail['msg']}" for detail in error.errors()) + + +def _as_mapping(value: object) -> dict[str, object] | None: + try: + return _mapping_adapter.validate_python(value) + except ValidationError: + return None + + +def _only_fields(mapping: Mapping[str, object], fields: frozenset[str]) -> dict[str, object]: + return {key: value for key, value in mapping.items() if key in fields} + + +def _dropping_unsupported_keys(block: Mapping[str, object]) -> dict[str, object]: + """What `drop_params` keeps of a file content block: the keys this surface understands, at every level.""" + kept_block: Final = _only_fields(block, _BLOCK_FIELDS) + file: Final = _as_mapping(block.get("file")) + if file is None: + return kept_block + kept_file: Final = _only_fields(file, _FILE_FIELDS) + video_metadata: Final = _as_mapping(file.get("video_metadata")) + if video_metadata is None: + return {**kept_block, "file": kept_file} + return { + **kept_block, + "file": {**kept_file, "video_metadata": _only_fields(video_metadata, _VIDEO_METADATA_FIELDS)}, + } + + +def _parse_file_block(element: object, drop_params: bool) -> _EmbeddingFileBlock: + if not isinstance(element, Mapping): + raise _invalid_input( + f"Embedding input elements must be strings or file content blocks, got {type(element).__name__}" + ) + block: Final = cast(Mapping[str, object], element) # cast-ok: isinstance leaves the key and value types unknown + try: + return _file_block_adapter.validate_python(_dropping_unsupported_keys(block) if drop_params else block) + except ValidationError as error: + raise _invalid_input( + f"Invalid file content block in embedding input: {_validation_error_summary(error)}" + ) from error + + +def _file_block_source(block: _EmbeddingFileBlock) -> str: + match (block.file.file_id, block.file.file_data): + case (str() as file_id, None): + return file_id + case (None, str() as file_data): + return file_data + case (None, None): + raise _invalid_input("A file content block in embedding input needs file.file_id or file.file_data") + case _: + raise _invalid_input( + "A file content block in embedding input takes file.file_id or file.file_data, not both" + ) + + +def _gemini_video_metadata(video_metadata: _EmbeddingVideoMetadata) -> VideoMetadataType: + return _video_metadata_adapter.validate_python( + gemini_video_metadata_from_openai(video_metadata.model_dump(exclude_none=True)) + ) + + +def _source_mime_type( + source: str, + mime_type_override: str | None, + resolved_files: Mapping[str, Mapping[str, str]], +) -> str | None: + if mime_type_override is not None: + return mime_type_override + if _is_data_url(source): + try: + return _parse_data_url(source)[0] + except ValueError: + return None + if _is_gcs_url(source): + try: + return _infer_mime_type_from_gcs_url(source) + except ValueError: + return None + if is_file_reference(source): + file_info: Final = resolved_files.get(source) + return None if file_info is None else file_info.get("mime_type") + return None + + +def _media_part( + source: str, + mime_type_override: str | None, + resolved_files: Mapping[str, Mapping[str, str]], +) -> PartType: + if _is_data_url(source): + mime_type, base64_data = _parse_data_url(source, mime_type_override) + return PartType(inline_data=BlobType(mime_type=mime_type, data=base64_data)) + if _is_gcs_url(source): + gcs_mime_type: Final = mime_type_override or _infer_mime_type_from_gcs_url(source) + return PartType(file_data=FileDataType(mime_type=gcs_mime_type, file_uri=source)) + if is_file_reference(source): + file_info: Final = resolved_files.get(source) + if file_info is None: + raise _invalid_input( + f"File reference {source!r} could not be resolved: " + "Gemini Files API references are only supported through the gemini/ provider" + ) + return PartType( + file_data=FileDataType(mime_type=mime_type_override or file_info["mime_type"], file_uri=file_info["uri"]) + ) + raise _invalid_input(f"A file content block source must be {_FILE_SOURCE_FORMS}, got {source[:50]!r}") + + +def _top_level_elements( + input: GeminiEmbeddingInput, +) -> Sequence[GeminiEmbeddingElement | list[str] | list[GeminiEmbeddingElement]]: + try: + _input_shape_adapter.validate_python(input) + except ValidationError as error: + raise _invalid_input( + f"Embedding input must be a string or a list of strings and file content blocks, got {type(input).__name__}" + ) from error + return [input] if isinstance(input, str) else input + + +def _elements(input: GeminiEmbeddingInput) -> Iterator[GeminiEmbeddingElement]: + for element in _top_level_elements(input): + if isinstance(element, list): + yield from element + else: + yield element + + +def _element_source(element: GeminiEmbeddingElement) -> str: + if isinstance(element, str): + return element + return _file_block_source(_parse_file_block(element, drop_params=True)) + + +def flatten_media_sources(input: GeminiEmbeddingInput) -> tuple[str, ...]: + """Every string element plus every file content block's source, in input order.""" + return tuple(_element_source(element) for element in _elements(input)) def _is_multimodal_input(input: GeminiEmbeddingInput) -> bool: """ Check if the input contains multimodal data (data URIs, file references, - GCS URLs, or nested lists for combined embeddings). + GCS URLs, file content blocks, or nested lists for combined embeddings). Args: - input: GeminiEmbeddingInput — str, List[str], or List[List[str]] for combined embeddings + input: GeminiEmbeddingInput — str, List[element], or List[List[element]] for combined embeddings Returns: - bool: True if any element is multimodal or a nested list + bool: True if any element is multimodal """ - if isinstance(input, str): - return _is_multimodal_element(input) - - for element in input: - if isinstance(element, list): - if any(_is_multimodal_element(sub) for sub in element if isinstance(sub, str)): - return True - elif isinstance(element, str) and _is_multimodal_element(element): - return True - - return False + return any(_is_multimodal_element(element) for element in _elements(input)) -def _is_multimodal_element(element: str) -> bool: - """Check if a single string element is multimodal.""" - if element.startswith("data:") and ";base64," in element: +def _is_multimodal_element(element: GeminiEmbeddingElement) -> bool: + """Check if a single element is multimodal.""" + if not isinstance(element, str): + return True + if _is_data_url(element): return True if is_file_reference(element): return True @@ -163,37 +351,29 @@ def _is_multimodal_element(element: str) -> bool: def _build_part_for_input( - element: str, - resolved_files: dict[str, dict[str, str]] | None = None, + element: GeminiEmbeddingElement, + resolved_files: Mapping[str, Mapping[str, str]] | None = None, + drop_params: bool = False, ) -> PartType: """ Build a single PartType for an input element, handling text, data URIs, - file references, and GCS URLs. + file references, GCS URLs, and file content blocks carrying a mime type + and video_metadata. """ - resolved_files = resolved_files or {} + files: Final = resolved_files or {} - if element.startswith("data:") and ";base64," in element: - mime_type, base64_data = _parse_data_url(element) - blob: Final[BlobType] = {"mime_type": mime_type, "data": base64_data} - return PartType(inline_data=blob) - elif _is_gcs_url(element): - mime_type = _infer_mime_type_from_gcs_url(element) - file_data: Final[FileDataType] = { - "mime_type": mime_type, - "file_uri": element, - } - return PartType(file_data=file_data) - elif is_file_reference(element): - if element not in resolved_files: - raise ValueError(f"File reference {element} not resolved") - file_info: Final = resolved_files[element] - file_data_ref: Final[FileDataType] = { - "mime_type": file_info["mime_type"], - "file_uri": file_info["uri"], - } - return PartType(file_data=file_data_ref) - else: - return PartType(text=element) + if isinstance(element, str): + return _media_part(element, None, files) if _is_multimodal_element(element) else PartType(text=element) + + block: Final = _parse_file_block(element, drop_params) + part: Final = _media_part(_file_block_source(block), block.file.format, files) + if block.file.video_metadata is None: + return part + video_metadata: Final = _gemini_video_metadata(block.file.video_metadata) + if not video_metadata: + return part + part_with_metadata: Final[PartType] = {**part, "video_metadata": video_metadata} + return part_with_metadata _SUPPORTED_EMBED_PARAMS: Final = {"outputDimensionality", "taskType", "title"} @@ -214,6 +394,7 @@ def transform_openai_input_gemini_content( model: str, optional_params: dict, resolved_files: dict[str, dict[str, str]] | None = None, + drop_params: bool = False, ) -> VertexAIBatchEmbeddingsRequestBody: """ Transform OpenAI embedding input to Gemini batchEmbedContents format. @@ -234,19 +415,18 @@ def transform_openai_input_gemini_content( gemini_params: Final = _filter_embed_params(optional_params) - input_list: Final = [input] if isinstance(input, str) else input + input_list: Final = _top_level_elements(input) requests: Final[list[EmbedContentRequest]] = [] for element in input_list: if isinstance(element, list): if not element: raise ValueError("Nested input list must not be empty") - for sub in element: - if not isinstance(sub, str): - raise ValueError(f"Elements inside a nested input list must be strings, got {type(sub)}") - parts = [_build_part_for_input(sub, resolved_files=resolved_files) for sub in element] + parts = [ + _build_part_for_input(sub, resolved_files=resolved_files, drop_params=drop_params) for sub in element + ] else: - parts = [_build_part_for_input(element, resolved_files=resolved_files)] + parts = [_build_part_for_input(element, resolved_files=resolved_files, drop_params=drop_params)] request = EmbedContentRequest( model=gemini_model_name, content=ContentType(parts=parts), @@ -262,6 +442,7 @@ def transform_openai_input_gemini_embed_content( model: str, optional_params: dict, resolved_files: dict[str, dict[str, str]] | None = None, + drop_params: bool = False, ) -> dict: """ Transform OpenAI embedding input to Gemini embedContent format (multimodal). @@ -279,7 +460,7 @@ def transform_openai_input_gemini_embed_content( gemini_params: Final = _filter_embed_params(optional_params) - input_list: Final = [input] if isinstance(input, str) else input + input_list: Final = _top_level_elements(input) parts: Final[list[PartType]] = [] for element in input_list: @@ -288,9 +469,7 @@ def transform_openai_input_gemini_embed_content( "Nested (combined) embeddings are not supported on the embedContent path. " "Use the batchEmbedContents path or pass a flat list instead." ) - if not isinstance(element, str): - raise ValueError(f"Unsupported input type: {type(element)}") - parts.append(_build_part_for_input(element, resolved_files=resolved_files)) + parts.append(_build_part_for_input(element, resolved_files=resolved_files, drop_params=drop_params)) request_body: Final[dict] = { "content": ContentType(parts=parts), @@ -313,31 +492,18 @@ def _parse_usage_metadata(raw_usage_metadata: object) -> UsageMetadata | None: return None -def _flatten_input(input: GeminiEmbeddingInput) -> tuple[str, ...]: - if isinstance(input, str): - return (input,) - return tuple(sub for element in input for sub in (element if isinstance(element, list) else [element])) +def _flatten_input(input: GeminiEmbeddingInput) -> tuple[GeminiEmbeddingElement, ...]: + return tuple(_elements(input)) def _is_image_element( - element: str, + element: GeminiEmbeddingElement, resolved_files: Mapping[str, Mapping[str, str]], ) -> bool: - if element.startswith("data:") and ";base64," in element: - try: - mime_type, _ = _parse_data_url(element) - except ValueError: - return False - return mime_type in _IMAGE_MIME_TYPES - if _is_gcs_url(element): - try: - return _infer_mime_type_from_gcs_url(element) in _IMAGE_MIME_TYPES - except ValueError: - return False - if is_file_reference(element): - file_info: Final = resolved_files.get(element) - return file_info is not None and file_info.get("mime_type") in _IMAGE_MIME_TYPES - return False + if isinstance(element, str): + return _source_mime_type(element, None, resolved_files) in _IMAGE_MIME_TYPES + block: Final = _parse_file_block(element, drop_params=True) + return _source_mime_type(_file_block_source(block), block.file.format, resolved_files) in _IMAGE_MIME_TYPES def _is_image_only_input( diff --git a/litellm/llms/voyage/common_utils.py b/litellm/llms/voyage/common_utils.py new file mode 100644 index 00000000000..2f4f63a8a43 --- /dev/null +++ b/litellm/llms/voyage/common_utils.py @@ -0,0 +1,35 @@ +""" +Shared helpers for the Voyage (VoyageAI by MongoDB) provider. +""" + +from typing import Final + +from litellm.secret_managers.main import get_secret_str + +VOYAGE_API_BASE: Final = "https://api.voyageai.com/v1" +MONGODB_API_BASE: Final = "https://ai.mongodb.com/v1" +MONGODB_API_KEY_PREFIX: Final = "al-" + + +def get_voyage_api_key(api_key: str | None = None) -> str | None: + """Resolve the key a Voyage request will authenticate with, explicit value first.""" + return ( + api_key + or get_secret_str("VOYAGE_API_KEY") + or get_secret_str("VOYAGE_AI_API_KEY") + or get_secret_str("VOYAGE_AI_TOKEN") + ) + + +def get_default_base_url(api_key: str | None = None) -> str: + """ + Pick the host that issued the key: MongoDB-issued keys (``al-`` prefix) are only + valid on ai.mongodb.com, every other key on api.voyageai.com. + + Mirrors ``voyageai.util.get_default_base_url`` in the official SDK: + https://github.com/voyage-ai/voyageai-python/blob/main/voyageai/util.py + """ + resolved: Final = get_voyage_api_key(api_key) + if resolved is not None and resolved.startswith(MONGODB_API_KEY_PREFIX): + return MONGODB_API_BASE + return VOYAGE_API_BASE diff --git a/litellm/llms/voyage/embedding/transformation.py b/litellm/llms/voyage/embedding/transformation.py index 7d74b1e00c4..a9232ad27e2 100644 --- a/litellm/llms/voyage/embedding/transformation.py +++ b/litellm/llms/voyage/embedding/transformation.py @@ -5,7 +5,7 @@ import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig -from litellm.secret_managers.main import get_secret_str +from litellm.llms.voyage.common_utils import get_default_base_url, get_voyage_api_key from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues from litellm.types.utils import EmbeddingResponse, Usage @@ -49,7 +49,7 @@ class VoyageEmbeddingConfig(BaseEmbeddingConfig): if not api_base.endswith("/embeddings"): api_base = f"{api_base}/embeddings" return api_base - return "https://api.voyageai.com/v1/embeddings" + return f"{get_default_base_url(api_key)}/embeddings" def get_supported_openai_params(self, model: str) -> list: return [ @@ -85,14 +85,8 @@ class VoyageEmbeddingConfig(BaseEmbeddingConfig): api_key: str | None = None, api_base: str | None = None, ) -> dict: - if api_key is None: - api_key = ( - get_secret_str("VOYAGE_API_KEY") - or get_secret_str("VOYAGE_AI_API_KEY") - or get_secret_str("VOYAGE_AI_TOKEN") - ) return { - "Authorization": f"Bearer {api_key}", + "Authorization": f"Bearer {get_voyage_api_key(api_key)}", } def transform_embedding_request( diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index 1ce2e6e3f29..9f81c8e40ac 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -11,7 +11,7 @@ import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig -from litellm.secret_managers.main import get_secret_str +from litellm.llms.voyage.common_utils import get_default_base_url, get_voyage_api_key from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues from litellm.types.utils import EmbeddingResponse, Usage @@ -58,7 +58,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): if not api_base.endswith("/contextualizedembeddings"): api_base = f"{api_base}/contextualizedembeddings" return api_base - return "https://api.voyageai.com/v1/contextualizedembeddings" + return f"{get_default_base_url(api_key)}/contextualizedembeddings" def get_supported_openai_params(self, model: str) -> list: # mutable-ok: base class signature return ["encoding_format", "dimensions"] @@ -91,14 +91,8 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): api_key: str | None = None, api_base: str | None = None, ) -> dict: - if api_key is None: - api_key = ( - get_secret_str("VOYAGE_API_KEY") - or get_secret_str("VOYAGE_AI_API_KEY") - or get_secret_str("VOYAGE_AI_TOKEN") - ) return { - "Authorization": f"Bearer {api_key}", + "Authorization": f"Bearer {get_voyage_api_key(api_key)}", } AUTO_CHUNK_SIZE: Final = 32000 diff --git a/litellm/llms/voyage/embedding/transformation_multimodal.py b/litellm/llms/voyage/embedding/transformation_multimodal.py index 814d5ab7eb0..035765b6691 100644 --- a/litellm/llms/voyage/embedding/transformation_multimodal.py +++ b/litellm/llms/voyage/embedding/transformation_multimodal.py @@ -13,7 +13,7 @@ import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig -from litellm.secret_managers.main import get_secret_str +from litellm.llms.voyage.common_utils import get_default_base_url, get_voyage_api_key from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues from litellm.types.utils import EmbeddingResponse, Usage @@ -58,7 +58,7 @@ class VoyageMultimodalEmbeddingConfig(BaseEmbeddingConfig): if not api_base.endswith("/multimodalembeddings"): api_base = f"{api_base}/multimodalembeddings" return api_base - return "https://api.voyageai.com/v1/multimodalembeddings" + return f"{get_default_base_url(api_key)}/multimodalembeddings" def get_supported_openai_params(self, model: str) -> list: return ["dimensions"] @@ -84,19 +84,14 @@ class VoyageMultimodalEmbeddingConfig(BaseEmbeddingConfig): api_key: str | None = None, api_base: str | None = None, ) -> dict: - if api_key is None: - api_key = ( - get_secret_str("VOYAGE_API_KEY") - or get_secret_str("VOYAGE_AI_API_KEY") - or get_secret_str("VOYAGE_AI_TOKEN") - ) - if not api_key: + resolved_api_key: Final = get_voyage_api_key(api_key) + if not resolved_api_key: raise ValueError( "Voyage API key is required for multimodal embeddings. " "Set VOYAGE_API_KEY / VOYAGE_AI_API_KEY / VOYAGE_AI_TOKEN " "or pass `api_key` explicitly." ) - return {"Authorization": f"Bearer {api_key}"} + return {"Authorization": f"Bearer {resolved_api_key}"} def _normalize_content_item(self, item: dict[str, object]) -> dict[str, object]: item_type: Final = item.get("type") diff --git a/litellm/llms/voyage/rerank/transformation.py b/litellm/llms/voyage/rerank/transformation.py index 3acb2f2ed58..121c18e82ae 100644 --- a/litellm/llms/voyage/rerank/transformation.py +++ b/litellm/llms/voyage/rerank/transformation.py @@ -13,7 +13,7 @@ from pydantic import ConfigDict, TypeAdapter from litellm._uuid import uuid from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig -from litellm.secret_managers.main import get_secret_str +from litellm.llms.voyage.common_utils import get_default_base_url, get_voyage_api_key from litellm.types.rerank import ( RerankBilledUnits, RerankResponse, @@ -31,6 +31,16 @@ _STR: Final = TypeAdapter(str) class VoyageRerankConfig(BaseRerankConfig): + """ + ``validate_environment`` stores the credential it authenticates with so ``get_complete_url`` + can select the host that issued it. ``ProviderConfigManager.get_provider_rerank_config`` + builds this config per request, so that key never reaches another one. + """ + + def __init__(self) -> None: + super().__init__() + self._api_key: str | None = None + def get_supported_cohere_rerank_params(self, model: str) -> list: return ["query", "documents", "top_n", "return_documents"] @@ -66,7 +76,7 @@ class VoyageRerankConfig(BaseRerankConfig): optional_params: dict | None = None, ) -> str: if api_base is None: - return "https://api.voyageai.com/v1/rerank" + return f"{get_default_base_url(self._api_key)}/rerank" api_base = api_base.rstrip("/") if not api_base.endswith("/v1/rerank"): if api_base.endswith("/v1"): @@ -148,12 +158,12 @@ class VoyageRerankConfig(BaseRerankConfig): optional_params: dict | None = None, litellm_params: Mapping[str, object] | None = None, ) -> dict: - if api_key is None: - api_key = get_secret_str("VOYAGE_API_KEY") or get_secret_str("VOYAGE_AI_API_KEY") - if api_key is None: + resolved_api_key: Final = get_voyage_api_key(api_key) + if resolved_api_key is None: raise ValueError("Voyage AI API key is required. Set via `api_key` parameter or `VOYAGE_API_KEY` env var.") + self._api_key = resolved_api_key return { - "Authorization": f"Bearer {api_key}", + "Authorization": f"Bearer {resolved_api_key}", "content-type": "application/json", } diff --git a/litellm/main.py b/litellm/main.py index 37b290f51ba..b02ccca9ecb 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -5737,6 +5737,7 @@ def completion( *ANTHROPIC_WIF_KWARGS_KEYS, *OPENAI_WIF_KWARGS_KEYS, PROVIDER_AFFINITY_HEADER_KWARG_KEY, + "fireworks_forward_user_id", ) if key in kwargs }, @@ -6378,6 +6379,10 @@ def embedding( """ azure: Final = kwargs.get("azure", None) client: Final = kwargs.pop("client", None) + drop_params_kwarg: Final = ( + cast(object, kwargs["drop_params"]) if "drop_params" in kwargs else None # cast-ok: untyped request kwargs + ) + drop_unsupported_params: Final = litellm.drop_params is True or normalize_drop_params(drop_params_kwarg) is True shared_session: Final = kwargs.get("shared_session", None) max_retries: Final = kwargs.get("max_retries", None) litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") @@ -6858,6 +6863,7 @@ def embedding( api_base=api_base, client=client, extra_headers=headers, + drop_params=drop_unsupported_params, ) elif custom_llm_provider == "vertex_ai": @@ -6910,6 +6916,7 @@ def embedding( api_base=api_base, client=client, extra_headers=headers, + drop_params=drop_unsupported_params, ) elif ( "image" in optional_params diff --git a/litellm/messages/dispatch.py b/litellm/messages/dispatch.py index 74736011f20..c7865132840 100644 --- a/litellm/messages/dispatch.py +++ b/litellm/messages/dispatch.py @@ -60,21 +60,23 @@ def _public_request( max_tokens: Final = fields.get("max_tokens") if not isinstance(model, str) or messages is None or not isinstance(max_tokens, int): return None - return native_call(args, kwargs, fields) + return native_call(legacy, args, kwargs) def _resolved_provider(request: NativeCall) -> str | None: try: - return get_llm_provider(str(request.bound["model"]), optional_str(request.bound.get("custom_llm_provider")))[1] + return get_llm_provider( + str(request.resolved["model"]), optional_str(request.resolved.get("custom_llm_provider")) + )[1] except BadRequestError: - return optional_str(request.bound.get("custom_llm_provider")) + return optional_str(request.resolved.get("custom_llm_provider")) def _context(request: NativeCall) -> RouteContext: return RouteContext( Route.MESSAGES, provider=_resolved_provider(request), - model=str(request.bound["model"]), + model=str(request.resolved["model"]), ) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index f34606ce127..7930a823e07 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -342,12 +342,14 @@ }, "writer.palmyra-vision-7b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 4096, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_vision": true }, @@ -486,13 +488,16 @@ }, "us.amazon.nova-2-lite-v1:0": { "cache_read_input_token_cost": 8.25e-08, + "cache_read_input_token_cost_flex": 4.125e-08, "input_cost_per_token": 3.3e-07, + "input_cost_per_token_flex": 1.65e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 1000000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 2.75e-06, + "output_cost_per_token_flex": 1.375e-06, "supports_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -556,13 +561,16 @@ }, "amazon.nova-pro-v1:0": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_flex": 1e-07, "input_cost_per_token": 8e-07, + "input_cost_per_token_flex": 4e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 300000, "max_output_tokens": 10000, "max_tokens": 10000, "mode": "chat", "output_cost_per_token": 3.2e-06, + "output_cost_per_token_flex": 1.6e-06, "supports_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -3188,13 +3196,16 @@ }, "apac.amazon.nova-pro-v1:0": { "cache_read_input_token_cost": 2.1e-07, + "cache_read_input_token_cost_flex": 1.05e-07, "input_cost_per_token": 8.4e-07, + "input_cost_per_token_flex": 4.2e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 300000, "max_output_tokens": 10000, "max_tokens": 10000, "mode": "chat", "output_cost_per_token": 3.36e-06, + "output_cost_per_token_flex": 1.68e-06, "supports_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -13076,12 +13087,14 @@ }, "bedrock/ap-northeast-1/deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -13089,32 +13102,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/ap-northeast-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/ap-northeast-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13122,43 +13139,52 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/ap-northeast-1/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.3e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.03e-06, + "output_cost_per_token_flex": 1.52e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/ap-northeast-1/moonshotai.kimi-k2.5": { "input_cost_per_token": 7.2e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.6e-06, + "output_cost_per_token_flex": 1.8e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/ap-northeast-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13168,13 +13194,17 @@ "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-coder-next.html" }, "bedrock/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.3e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.03e-06, + "output_cost_per_token_flex": 1.52e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, @@ -13182,11 +13212,11 @@ "input_cost_per_token": 6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.03e-06, - "source": "https://platform.moonshot.ai/docs/guide/kimi-k2-5-quickstart", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13216,12 +13246,14 @@ }, "bedrock/ap-south-1/deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -13229,32 +13261,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/ap-south-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/ap-south-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13262,43 +13298,52 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/ap-south-1/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.1e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.94e-06, + "output_cost_per_token_flex": 1.47e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/ap-south-1/moonshotai.kimi-k2.5": { "input_cost_per_token": 7.2e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.6e-06, + "output_cost_per_token_flex": 1.8e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/ap-south-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13309,12 +13354,13 @@ }, "bedrock/ap-southeast-2/minimax.minimax-m2.5": { "input_cost_per_token": 3.1e-07, + "input_cost_per_token_flex": 1.55e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13322,16 +13368,19 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.24e-06 + "output_cost_per_token": 1.24e-06, + "output_cost_per_token_flex": 6.2e-07 }, "bedrock/ap-southeast-3/deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -13339,32 +13388,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/ap-southeast-3/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/ap-southeast-3/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13372,32 +13425,37 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/ap-southeast-3/moonshotai.kimi-k2.5": { "input_cost_per_token": 7.2e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.6e-06, + "output_cost_per_token_flex": 1.8e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/ap-southeast-3/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13426,12 +13484,14 @@ }, "bedrock/eu-north-1/deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -13439,32 +13499,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/eu-north-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/eu-north-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13472,23 +13536,26 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/eu-north-1/moonshotai.kimi-k2.5": { "input_cost_per_token": 7.2e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.6e-06, + "output_cost_per_token_flex": 1.8e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/eu-central-1/1-month-commitment/anthropic.claude-instant-v1": { "cost_per_second": 0.01635, @@ -13578,29 +13645,33 @@ "supports_tool_choice": true }, "bedrock/eu-central-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/eu-central-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13608,16 +13679,19 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/eu-central-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13645,29 +13719,33 @@ "output_cost_per_token": 6.5e-07 }, "bedrock/eu-west-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/eu-west-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13675,16 +13753,19 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/eu-west-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13712,29 +13793,33 @@ "output_cost_per_token": 7.8e-07 }, "bedrock/eu-west-2/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 4.7e-07, + "input_cost_per_token_flex": 2.35e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.86e-06, + "output_cost_per_token_flex": 9.3e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/eu-west-2/minimax.minimax-m2.5": { "input_cost_per_token": 4.7e-07, + "input_cost_per_token_flex": 2.35e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13742,16 +13827,19 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.86e-06 + "output_cost_per_token": 1.86e-06, + "output_cost_per_token_flex": 9.3e-07 }, "bedrock/eu-west-2/nvidia.nemotron-super-3-120b": { "input_cost_per_token": 2.3e-07, + "input_cost_per_token_flex": 1.15e-07, "litellm_provider": "bedrock", "max_input_tokens": 256000, "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 1.01e-06, + "output_cost_per_token_flex": 5.05e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -13763,12 +13851,14 @@ }, "bedrock/eu-west-2/qwen.qwen3-coder-next": { "input_cost_per_token": 7.8e-07, + "input_cost_per_token_flex": 3.9e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.86e-06, + "output_cost_per_token_flex": 9.3e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13809,29 +13899,33 @@ "supports_tool_choice": true }, "bedrock/eu-south-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/eu-south-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13839,16 +13933,19 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/eu-south-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13895,12 +13992,14 @@ }, "bedrock/sa-east-1/deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -13908,32 +14007,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/sa-east-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/sa-east-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13941,43 +14044,52 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/sa-east-1/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.3e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.03e-06, + "output_cost_per_token_flex": 1.52e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/sa-east-1/moonshotai.kimi-k2.5": { "input_cost_per_token": 7.2e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.6e-06, + "output_cost_per_token_flex": 1.8e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/sa-east-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -14125,12 +14237,14 @@ }, "bedrock/us-east-1/deepseek.v3.2": { "input_cost_per_token": 6.2e-07, + "input_cost_per_token_flex": 3.1e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.85e-06, + "output_cost_per_token_flex": 9.25e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -14138,29 +14252,34 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/us-east-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/us-east-1/minimax.minimax-m2.5": { "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", @@ -14170,43 +14289,51 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html" }, "bedrock/us-east-1/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_flex": 1.25e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/us-east-1/moonshotai.kimi-k2.5": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3e-06, + "output_cost_per_token_flex": 1.5e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/us-east-1/qwen.qwen3-coder-next": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -14217,12 +14344,14 @@ }, "bedrock/us-east-2/deepseek.v3.2": { "input_cost_per_token": 6.2e-07, + "input_cost_per_token_flex": 3.1e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.85e-06, + "output_cost_per_token_flex": 9.25e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -14230,32 +14359,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/us-east-2/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/us-east-2/minimax.minimax-m2.5": { "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -14263,43 +14396,52 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.2e-06 + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07 }, "bedrock/us-east-2/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_flex": 1.25e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/us-east-2/moonshotai.kimi-k2.5": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3e-06, + "output_cost_per_token_flex": 1.5e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/us-east-2/qwen.qwen3-coder-next": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -14755,12 +14897,14 @@ }, "bedrock/us-west-2/deepseek.v3.2": { "input_cost_per_token": 6.2e-07, + "input_cost_per_token_flex": 3.1e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.85e-06, + "output_cost_per_token_flex": 9.25e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -14768,29 +14912,34 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/us-west-2/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/us-west-2/minimax.minimax-m2.5": { "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", @@ -14800,43 +14949,51 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html" }, "bedrock/us-west-2/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_flex": 1.25e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/us-west-2/moonshotai.kimi-k2.5": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3e-06, + "output_cost_per_token_flex": 1.5e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/us-west-2/qwen.qwen3-coder-next": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -15039,6 +15196,7 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "claude-haiku-4-5-20251001": { + "supports_web_search": true, "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, "cache_creation_input_token_cost_batches": 6.25e-07, @@ -15181,6 +15339,7 @@ "source": "https://docs.anthropic.com/en/docs/about-claude/pricing" }, "claude-sonnet-5": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, @@ -15271,6 +15430,7 @@ "supports_web_search": true }, "claude-sonnet-4-6": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 3.75e-06, "cache_creation_input_token_cost_above_1hr": 6e-06, @@ -15344,6 +15504,7 @@ "output_cost_per_token_batches": 7.5e-06 }, "claude-opus-4-5-20251101": { + "supports_web_search": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_creation_input_token_cost_batches": 3.125e-06, @@ -15413,6 +15574,7 @@ "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-6": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, @@ -15498,6 +15660,7 @@ "prompt_cache_min_tokens": 4096 }, "claude-opus-4-7": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, @@ -15585,6 +15748,7 @@ "prompt_cache_min_tokens": 2048 }, "claude-fable-5": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, @@ -15629,6 +15793,7 @@ "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-fable-5-1": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, @@ -15674,6 +15839,7 @@ "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-5-5": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 5e-06, "cache_creation_input_token_cost_above_1hr": 8e-06, @@ -15719,6 +15885,7 @@ "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-5": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, @@ -15766,6 +15933,7 @@ "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-8": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, @@ -22487,14 +22655,15 @@ "supports_tool_choice": true }, "deepseek.v3-v1:0": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 5.8e-07, "litellm_provider": "bedrock_converse", - "max_input_tokens": 163840, - "max_output_tokens": 81920, - "max_tokens": 81920, + "max_input_tokens": 128000, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.68e-06, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-1.html", "supports_function_calling": true, "supports_reasoning": true, "supports_tool_choice": true, @@ -22502,12 +22671,14 @@ }, "deepseek.v3.2": { "input_cost_per_token": 6.2e-07, + "input_cost_per_token_flex": 3.1e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 164000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.85e-06, + "output_cost_per_token_flex": 9.25e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -23183,13 +23354,16 @@ }, "eu.amazon.nova-pro-v1:0": { "cache_read_input_token_cost": 2.625e-07, + "cache_read_input_token_cost_flex": 1.3125e-07, "input_cost_per_token": 1.05e-06, + "input_cost_per_token_flex": 5.25e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 300000, "max_output_tokens": 10000, "max_tokens": 10000, "mode": "chat", "output_cost_per_token": 4.2e-06, + "output_cost_per_token_flex": 2.1e-06, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_function_calling": true, "supports_pdf_input": true, @@ -29253,7 +29427,7 @@ }, "gemini/gemini-2.5-flash-image": { "cache_read_input_token_cost": 3e-08, - "deprecation_date": "2026-10-02", + "deprecation_date": "2027-03-15", "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 3e-07, "input_cost_per_token_batches": 1.5e-07, @@ -29703,6 +29877,7 @@ "web_search_billing_unit": "per_query" }, "gemini/gemini-2.5-flash-preview-tts": { + "deprecation_date": "2026-11-17", "input_cost_per_token": 5e-07, "input_cost_per_token_batches": 2.5e-07, "litellm_provider": "gemini", @@ -30276,7 +30451,7 @@ "supports_video_input": true, "supports_vision": true, "tpm": 800000, - "deprecation_date": "2026-09-30" + "deprecation_date": "2026-10-22" }, "gemini/gemini-3.1-pro-preview": { "prompt_cache_min_tokens": 4096, @@ -30783,6 +30958,7 @@ }, "gemini/gemini-2.5-pro-preview-tts": { "cache_read_input_token_cost": 1.25e-07, + "deprecation_date": "2026-11-17", "input_cost_per_audio_token": 7e-07, "input_cost_per_token": 1e-06, "input_cost_per_token_batches": 5e-07, @@ -32129,14 +32305,17 @@ "output_cost_per_token": 2e-06 }, "google.gemma-3-12b-it": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 9e-08, + "input_cost_per_token_flex": 5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.9e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 1.5e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-google-gemma-3-12b-it.html", "supports_audio_input": false, "supports_function_calling": true, "supports_response_schema": true, @@ -32144,14 +32323,17 @@ "supports_vision": true }, "google.gemma-3-27b-it": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 2.3e-07, + "input_cost_per_token_flex": 1.2e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 3.8e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 1.9e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-google-gemma-3-27b-pt.html", "supports_audio_input": false, "supports_function_calling": true, "supports_response_schema": true, @@ -32159,14 +32341,17 @@ "supports_vision": true }, "google.gemma-3-4b-it": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 4e-08, + "input_cost_per_token_flex": 2e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 8e-08, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 4e-08, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-google-gemma-3-4b-it.html", "supports_audio_input": false, "supports_function_calling": true, "supports_system_messages": true, @@ -38434,14 +38619,17 @@ "supports_tool_choice": true }, "minimax.minimax-m2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 1000000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 6e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2.html", "supports_audio_input": false, "supports_function_calling": true, "supports_system_messages": true, @@ -38450,24 +38638,29 @@ "supports_vision": false }, "minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 196000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html", "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false }, "minimax.minimax-m2.5": { "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 196000, "max_output_tokens": 8000, @@ -38607,29 +38800,34 @@ "max_output_tokens": 128000 }, "mistral.devstral-2-123b": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 4e-07, + "input_cost_per_token_flex": 2e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 2e-06, + "output_cost_per_token_flex": 1e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-mistral-ai-devstral-2-123b.html", "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false }, "mistral.magistral-small-2509": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 40000, "max_tokens": 40000, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_flex": 7.5e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -38640,12 +38838,14 @@ }, "mistral.ministral-3-14b-instruct": { "input_cost_per_token": 2e-07, + "input_cost_per_token_flex": 1e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2e-07, + "output_cost_per_token_flex": 1e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -38656,12 +38856,14 @@ }, "mistral.ministral-3-3b-instruct": { "input_cost_per_token": 1e-07, + "input_cost_per_token_flex": 5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1e-07, + "output_cost_per_token_flex": 5e-08, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -38672,12 +38874,14 @@ }, "mistral.ministral-3-8b-instruct": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.5e-07, + "output_cost_per_token_flex": 7e-08, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -38722,12 +38926,14 @@ }, "mistral.mistral-large-3-675b-instruct": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_flex": 7.5e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -38761,27 +38967,33 @@ "supports_tool_choice": true }, "mistral.voxtral-mini-3b-2507": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 4e-08, + "input_cost_per_token_flex": 2e-08, "litellm_provider": "bedrock_converse", - "max_input_tokens": 128000, + "max_input_tokens": 32768, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 4e-08, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 2e-08, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-mistral-ai-voxtral-mini-3b-2507.html", "supports_audio_input": true, "supports_system_messages": true, "supports_native_structured_output": true }, "mistral.voxtral-small-24b-2507": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 1e-07, + "input_cost_per_token_flex": 5e-08, "litellm_provider": "bedrock_converse", - "max_input_tokens": 128000, + "max_input_tokens": 32768, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 3e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 1.5e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-mistral-ai-voxtral-small-24b-2507.html", "supports_audio_input": true, "supports_system_messages": true, "supports_native_structured_output": true @@ -39262,8 +39474,8 @@ "cache_read_input_token_cost": 6.8e-08, "input_cost_per_token": 6.8e-07, "litellm_provider": "mistral", - "max_input_tokens": 524288, - "max_tokens": 524288, + "max_input_tokens": 1048576, + "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 2.09e-06, "reasoning_effort_levels": [ @@ -39283,8 +39495,8 @@ "cache_read_input_token_cost": 6.8e-08, "input_cost_per_token": 6.8e-07, "litellm_provider": "mistral", - "max_input_tokens": 524288, - "max_tokens": 524288, + "max_input_tokens": 1048576, + "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 2.09e-06, "reasoning_effort_levels": [ @@ -39590,6 +39802,7 @@ "supports_vision": true }, "moonshot.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, @@ -39597,7 +39810,7 @@ "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.5e-06, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_audio_input": false, "supports_function_calling": true, "supports_reasoning": true, @@ -39608,12 +39821,14 @@ }, "moonshotai.kimi-k2.5": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3e-06, + "output_cost_per_token_flex": 1.5e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -40708,14 +40923,17 @@ "source": "https://tokenfactory.nebius.com/models/catalog/embedding/Qwen%2FQwen3-Embedding-8B" }, "nvidia.nemotron-nano-12b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 2e-07, + "input_cost_per_token_flex": 1e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 3e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-12b-v2-vl-bf16.html", "supports_audio_input": false, "supports_function_calling": true, "supports_response_schema": true, @@ -40723,14 +40941,17 @@ "supports_vision": true }, "nvidia.nemotron-nano-9b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-08, + "input_cost_per_token_flex": 3e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.3e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 1.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-9b-v2.html", "supports_audio_input": false, "supports_function_calling": true, "supports_response_schema": true, @@ -40739,12 +40960,14 @@ }, "nvidia.nemotron-nano-3-30b": { "input_cost_per_token": 6e-08, + "input_cost_per_token_flex": 3e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.4e-07, + "output_cost_per_token_flex": 1.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -40756,12 +40979,14 @@ }, "nvidia.nemotron-super-3-120b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 6.5e-07, + "output_cost_per_token_flex": 3.25e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -42095,14 +42320,15 @@ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "input_cost_per_token": 1.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -42113,14 +42339,15 @@ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "input_cost_per_token": 7e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 3e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -42128,14 +42355,17 @@ }, "openai.gpt-oss-safeguard-120b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_response_schema": true, "supports_system_messages": true, @@ -42143,14 +42373,17 @@ }, "openai.gpt-oss-safeguard-20b": { "input_cost_per_token": 7e-08, + "input_cost_per_token_flex": 3e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2e-07, + "output_cost_per_token_flex": 1e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_response_schema": true, "supports_system_messages": true, @@ -42415,6 +42648,33 @@ "prompt_cache_min_tokens": 512, "supports_sampling_params": false }, + "openrouter/anthropic/claude-haiku-5.5:batch": { + "cache_creation_input_token_cost": 6.25e-08, + "cache_creation_input_token_cost_above_100k_tokens": 3.125e-07, + "cache_creation_input_token_cost_above_1hr": 1e-07, + "cache_creation_input_token_cost_above_1hr_above_100k_tokens": 5e-07, + "cache_read_input_token_cost": 5e-09, + "cache_read_input_token_cost_above_100k_tokens": 2.5e-08, + "input_cost_per_token": 5e-08, + "input_cost_per_token_above_100k_tokens": 2.5e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 2.5e-07, + "output_cost_per_token_above_100k_tokens": 1.25e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "openrouter/anthropic/claude-haiku-4.5": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, @@ -42600,14 +42860,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v3.2": { - "cache_read_input_token_cost": 2.8e-08, + "cache_read_input_token_cost": 1.35e-07, "deprecation_date": "2026-09-28", - "input_cost_per_token": 2.8e-07, - "input_cost_per_token_cache_hit": 2.8e-08, + "input_cost_per_token": 2.59e-07, + "input_cost_per_token_cache_hit": 1.35e-07, "litellm_provider": "openrouter", "max_input_tokens": 163840, - "max_output_tokens": 65536, - "max_tokens": 65536, + "max_output_tokens": 147456, + "max_tokens": 147456, "mode": "chat", "output_cost_per_token": 4.2e-07, "supports_assistant_prefill": true, @@ -42690,14 +42950,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "cache_read_input_token_cost": 1.74e-08, - "input_cost_per_token": 2.088e-07, + "cache_read_input_token_cost": 2.39395e-08, + "input_cost_per_token": 2.87274e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 4.176e-07, + "output_cost_per_token": 5.74548e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42710,14 +42970,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "cache_read_input_token_cost": 1.5e-08, - "input_cost_per_token": 4.4e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 3e-07, + "output_cost_per_token": 1.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -43304,14 +43564,14 @@ "supports_web_search": false }, "openrouter/nvidia/nemotron-3.5-lightning": { - "cache_read_input_token_cost": 3e-08, - "input_cost_per_token": 6e-08, + "cache_read_input_token_cost": 2.45e-08, + "input_cost_per_token": 4.9e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.6e-07, + "output_cost_per_token": 1.4e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -44055,13 +44315,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.5-35b-a3b": { - "input_cost_per_token": 1.5e-07, + "input_cost_per_token": 1.625e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 235929, - "max_tokens": 235929, + "max_output_tokens": 65536, + "max_tokens": 65536, "mode": "chat", - "output_cost_per_token": 1e-06, + "output_cost_per_token": 1.3e-06, "source": "https://openrouter.ai/api/v1/models", "supports_function_calling": true, "supports_reasoning": true, @@ -44075,13 +44335,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.5-27b": { - "input_cost_per_token": 1.95e-07, + "input_cost_per_token": 2.6e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 65536, - "max_tokens": 65536, + "max_output_tokens": 81920, + "max_tokens": 81920, "mode": "chat", - "output_cost_per_token": 1.56e-06, + "output_cost_per_token": 2.6e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -44153,14 +44413,14 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.5-397b-a17b": { - "cache_read_input_token_cost": 2.2e-07, - "input_cost_per_token": 4.5e-07, + "cache_read_input_token_cost": 2.25e-07, + "input_cost_per_token": 5.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 81920, - "max_tokens": 81920, + "max_output_tokens": 235929, + "max_tokens": 235929, "mode": "chat", - "output_cost_per_token": 3e-06, + "output_cost_per_token": 3.5e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -45287,12 +45547,14 @@ }, "qwen.qwen3-235b-a22b-2507-v1:0": { "input_cost_per_token": 2.2e-07, + "input_cost_per_token_flex": 1.1e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 262144, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 8.8e-07, + "output_cost_per_token_flex": 4.4e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_tool_choice": true, @@ -45318,12 +45580,14 @@ }, "qwen.qwen3-32b-v1:0": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 32768, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-32b.html", "supports_function_calling": true, "supports_reasoning": true, @@ -45460,12 +45724,14 @@ }, "qwen.qwen3-coder-next": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -47124,7 +47390,6 @@ }, "together_ai/deepseek-ai/DeepSeek-V4-Flash-0731": { "cache_read_input_token_cost": 3e-08, - "deprecation_date": "2026-09-29", "input_cost_per_token": 1.4e-07, "litellm_provider": "together_ai", "max_input_tokens": 1048576, @@ -47154,7 +47419,6 @@ }, "together_ai/deepseek-ai/DeepSeek-V4-Pro-0813": { "cache_read_input_token_cost": 1.3e-07, - "deprecation_date": "2026-09-29", "input_cost_per_token": 1.32e-06, "litellm_provider": "together_ai", "max_input_tokens": 1048576, @@ -47366,13 +47630,16 @@ }, "us.amazon.nova-pro-v1:0": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_flex": 1e-07, "input_cost_per_token": 8e-07, + "input_cost_per_token_flex": 4e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 300000, "max_output_tokens": 10000, "max_tokens": 10000, "mode": "chat", "output_cost_per_token": 3.2e-06, + "output_cost_per_token_flex": 1.6e-06, "supports_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -47781,6 +48048,7 @@ "supports_vision": false }, "us-gov.nvidia.nemotron-nano-12b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 2.4e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, @@ -47788,6 +48056,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-12b-v2-vl-bf16.html", "supports_system_messages": true, "supports_vision": true, "supports_audio_input": false, @@ -47795,6 +48064,7 @@ "supports_response_schema": true }, "us-gov.nvidia.nemotron-nano-9b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.2e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, @@ -47802,6 +48072,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 2.76e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-9b-v2.html", "supports_system_messages": true, "supports_audio_input": false, "supports_function_calling": true, @@ -47829,14 +48100,16 @@ "input_cost_per_token": 8.4e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 3.6e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -47846,14 +48119,16 @@ "input_cost_per_token": 1.8e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -48082,12 +48357,14 @@ }, "us.deepseek.v3.2": { "input_cost_per_token": 6.2e-07, + "input_cost_per_token_flex": 3.1e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 163840, "max_output_tokens": 163840, "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 1.85e-06, + "output_cost_per_token_flex": 9.25e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -48098,12 +48375,14 @@ }, "eu.deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 163840, "max_output_tokens": 163840, "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -52898,18 +53177,21 @@ "supports_web_search": true }, "zai.glm-4.7": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 203000, "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 2.2e-06, + "output_cost_per_token_flex": 1.1e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-zai-glm-4-7.html", "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false @@ -52933,18 +53215,21 @@ "supports_vision": false }, "zai.glm-4.7-flash": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7e-08, + "input_cost_per_token_flex": 3.5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 203000, "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 4e-07, + "output_cost_per_token_flex": 2e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-zai-glm-4-7-flash.html", "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false @@ -53201,7 +53486,7 @@ ] }, "azure/sora-2": { - "deprecation_date": "2026-10-15", + "deprecation_date": "2026-11-02", "litellm_provider": "azure", "mode": "video_generation", "output_cost_per_video_per_second": 0.1, @@ -57519,6 +57804,7 @@ "gemini_native_audio": true }, "gemini-2.5-flash-native-audio-preview-12-2025": { + "deprecation_date": "2026-11-17", "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", @@ -57549,6 +57835,7 @@ "supports_web_search": true }, "gemini-3.1-flash-live-preview": { + "deprecation_date": "2026-11-17", "input_cost_per_audio_token": 3e-06, "input_cost_per_image_token": 1e-06, "input_cost_per_token": 7.5e-07, @@ -57710,6 +57997,7 @@ "gemini_native_audio": true }, "gemini/gemini-2.5-flash-native-audio-preview-12-2025": { + "deprecation_date": "2026-11-17", "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", @@ -57743,6 +58031,7 @@ "supports_reasoning": true }, "gemini/gemini-3.1-flash-live-preview": { + "deprecation_date": "2026-11-17", "input_cost_per_audio_token": 3e-06, "input_cost_per_image_token": 1e-06, "input_cost_per_token": 7.5e-07, @@ -57782,6 +58071,7 @@ "supports_reasoning": true }, "gemini/gemini-3.1-flash-tts-preview": { + "deprecation_date": "2026-11-17", "input_cost_per_token": 1e-06, "input_cost_per_token_batches": 5e-07, "litellm_provider": "gemini", @@ -57853,6 +58143,7 @@ "supports_prompt_caching": true }, "gemini-2.5-flash-preview-tts": { + "deprecation_date": "2026-11-17", "input_cost_per_token": 5e-07, "input_cost_per_token_batches": 2.5e-07, "litellm_provider": "gemini", @@ -58243,12 +58534,15 @@ }, "bedrock_mantle/openai.gpt-oss-120b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -58261,12 +58555,15 @@ }, "bedrock_mantle/openai.gpt-oss-20b": { "input_cost_per_token": 7e-08, + "input_cost_per_token_flex": 3.5e-08, "output_cost_per_token": 3e-07, + "output_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -58279,12 +58576,15 @@ }, "bedrock_mantle/openai.gpt-oss-safeguard-120b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7e-08, "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 65536, - "max_tokens": 65536, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-safeguard-120b.html", "supported_endpoints": [ "/v1/chat/completions" ], @@ -58295,12 +58595,15 @@ }, "bedrock_mantle/openai.gpt-oss-safeguard-20b": { "input_cost_per_token": 7e-08, + "input_cost_per_token_flex": 3e-08, "output_cost_per_token": 2e-07, + "output_cost_per_token_flex": 1e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 65536, - "max_tokens": 65536, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-safeguard-20b.html", "supported_endpoints": [ "/v1/chat/completions" ], @@ -58344,6 +58647,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58387,6 +58691,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58419,6 +58724,7 @@ "supports_function_calling": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58495,6 +58801,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58882,6 +59189,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58920,6 +59228,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58958,6 +59267,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -59357,7 +59667,9 @@ }, "bedrock_mantle/google.gemma-4-31b": { "input_cost_per_token": 1.4e-07, + "input_cost_per_token_flex": 7e-08, "output_cost_per_token": 4e-07, + "output_cost_per_token_flex": 2e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 256000, @@ -59376,7 +59688,9 @@ }, "bedrock_mantle/google.gemma-4-26b-a4b": { "input_cost_per_token": 1.3e-07, + "input_cost_per_token_flex": 6.5e-08, "output_cost_per_token": 4e-07, + "output_cost_per_token_flex": 2e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 256000, @@ -59395,7 +59709,9 @@ }, "bedrock_mantle/google.gemma-4-e2b": { "input_cost_per_token": 4e-08, + "input_cost_per_token_flex": 2e-08, "output_cost_per_token": 8e-08, + "output_cost_per_token_flex": 4e-08, "litellm_provider": "bedrock_mantle", "max_input_tokens": 128000, "max_output_tokens": 128000, @@ -65580,6 +65896,7 @@ "supports_vision": false }, "bedrock/us-gov-west-1/nvidia.nemotron-nano-12b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 2.4e-07, "litellm_provider": "bedrock", "max_input_tokens": 128000, @@ -65587,6 +65904,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-12b-v2-vl-bf16.html", "supports_system_messages": true, "supports_vision": true, "supports_audio_input": false, @@ -65594,6 +65912,7 @@ "supports_response_schema": true }, "bedrock/us-gov-west-1/nvidia.nemotron-nano-9b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.2e-08, "litellm_provider": "bedrock", "max_input_tokens": 128000, @@ -65601,6 +65920,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 2.76e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-9b-v2.html", "supports_system_messages": true, "supports_audio_input": false, "supports_function_calling": true, @@ -65628,14 +65948,16 @@ "input_cost_per_token": 8.4e-08, "litellm_provider": "bedrock", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 3.6e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -65645,14 +65967,16 @@ "input_cost_per_token": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -65842,6 +66166,7 @@ "supports_vision": false }, "bedrock/us-gov-east-1/nvidia.nemotron-nano-12b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 2.4e-07, "litellm_provider": "bedrock", "max_input_tokens": 128000, @@ -65849,6 +66174,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-12b-v2-vl-bf16.html", "supports_system_messages": true, "supports_vision": true, "supports_audio_input": false, @@ -65856,6 +66182,7 @@ "supports_response_schema": true }, "bedrock/us-gov-east-1/nvidia.nemotron-nano-9b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.2e-08, "litellm_provider": "bedrock", "max_input_tokens": 128000, @@ -65863,6 +66190,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 2.76e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-9b-v2.html", "supports_system_messages": true, "supports_audio_input": false, "supports_function_calling": true, @@ -65890,14 +66218,16 @@ "input_cost_per_token": 8.4e-08, "litellm_provider": "bedrock", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 3.6e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -65907,14 +66237,16 @@ "input_cost_per_token": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -66292,9 +66624,10 @@ "output_cost_per_token": 3.6e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -66310,9 +66643,10 @@ "output_cost_per_token": 7.2e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -66449,9 +66783,10 @@ "output_cost_per_token": 3.6e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -66467,9 +66802,10 @@ "output_cost_per_token": 7.2e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -66481,8 +66817,11 @@ "supports_tool_choice": true }, "bedrock_mantle/deepseek.v3.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 5.8e-07, + "input_cost_per_token_flex": 2.9e-07, "output_cost_per_token": 1.68e-06, + "output_cost_per_token_flex": 8.4e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 128000, "max_output_tokens": 8000, @@ -66496,8 +66835,11 @@ "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-1.html" }, "bedrock_mantle/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "output_cost_per_token": 2.5e-06, + "output_cost_per_token_flex": 1.25e-06, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 16000, @@ -66513,7 +66855,9 @@ }, "bedrock_mantle/qwen.qwen3-235b-a22b-2507": { "input_cost_per_token": 2.2e-07, + "input_cost_per_token_flex": 1.1e-07, "output_cost_per_token": 8.8e-07, + "output_cost_per_token_flex": 4.4e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 8000, @@ -66529,7 +66873,9 @@ }, "bedrock_mantle/qwen.qwen3-32b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 32000, "max_output_tokens": 8000, @@ -66546,7 +66892,9 @@ "bedrock_mantle/qwen.qwen3-coder-30b-a3b-instruct": { "deprecation_date": "2027-03-30", "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 16000, @@ -66562,7 +66910,9 @@ "bedrock_mantle/qwen.qwen3-coder-480b-a35b-instruct": { "deprecation_date": "2027-03-30", "input_cost_per_token": 4.5e-07, + "input_cost_per_token_flex": 2.25e-07, "output_cost_per_token": 1.8e-06, + "output_cost_per_token_flex": 9e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 128000, "max_output_tokens": 16000, @@ -66577,7 +66927,9 @@ }, "bedrock_mantle/qwen.qwen3-next-80b-a3b-instruct": { "input_cost_per_token": 1.4e-07, + "input_cost_per_token_flex": 7e-08, "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 8000, @@ -66593,7 +66945,9 @@ }, "bedrock_mantle/qwen.qwen3-vl-235b-a22b-instruct": { "input_cost_per_token": 5.3e-07, + "input_cost_per_token_flex": 2.6e-07, "output_cost_per_token": 2.66e-06, + "output_cost_per_token_flex": 1.33e-06, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 8000, @@ -67997,14 +68351,14 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "cache_read_input_token_cost": 6.5e-08, - "input_cost_per_token": 7e-08, + "cache_read_input_token_cost": 3.8e-08, + "input_cost_per_token": 3.9e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 7e-06, + "output_cost_per_token": 3.39e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68135,8 +68489,8 @@ "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-flash-0731": { - "cache_read_input_token_cost": 1.63e-08, - "input_cost_per_token": 1.63e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 6e-09, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -68226,14 +68580,14 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k3": { - "cache_read_input_token_cost": 2.6e-07, - "input_cost_per_token": 7.9e-07, + "cache_read_input_token_cost": 4e-07, + "input_cost_per_token": 7.6e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 1.5e-05, + "output_cost_per_token": 1.3e-05, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68351,14 +68705,14 @@ "supports_web_search": true }, "openrouter/z-ai/glm-5.2": { - "cache_read_input_token_cost": 1.5e-07, - "input_cost_per_token": 1.52e-07, + "cache_read_input_token_cost": 6.8e-08, + "input_cost_per_token": 6.9e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 1.2e-05, + "output_cost_per_token": 4.3e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68654,13 +69008,13 @@ }, "openrouter/qwen/qwen3.6-27b": { "cache_read_input_token_cost": 3e-08, - "input_cost_per_token": 3e-07, + "input_cost_per_token": 4.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", - "output_cost_per_token": 2e-06, + "output_cost_per_token": 2.7e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68714,8 +69068,8 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "cache_read_input_token_cost": 3e-08, - "input_cost_per_token": 3e-08, + "cache_read_input_token_cost": 7.5e-09, + "input_cost_per_token": 7.5e-09, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -68734,14 +69088,14 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k2.6": { - "cache_read_input_token_cost": 1.6e-07, - "input_cost_per_token": 9.5e-07, + "cache_read_input_token_cost": 9.75e-08, + "input_cost_per_token": 4.65e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, "max_tokens": 235929, "mode": "chat", - "output_cost_per_token": 4e-06, + "output_cost_per_token": 2.45e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68968,6 +69322,7 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.5-9b": { + "deprecation_date": "2026-10-21", "input_cost_per_token": 1e-07, "output_cost_per_token": 1.5e-07, "litellm_provider": "openrouter", @@ -69161,8 +69516,8 @@ "supports_web_search": false }, "openrouter/nvidia/nemotron-3-nano-30b-a3b": { - "input_cost_per_token": 5e-08, - "output_cost_per_token": 2e-07, + "input_cost_per_token": 6e-08, + "output_cost_per_token": 2.4e-07, "cache_read_input_token_cost": 3e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, @@ -69183,7 +69538,7 @@ "openrouter/z-ai/glm-4.6v": { "input_cost_per_token": 3e-07, "output_cost_per_token": 9e-07, - "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost": 5.5e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, @@ -69609,13 +69964,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-next-80b-a3b-instruct": { - "input_cost_per_token": 1e-07, + "input_cost_per_token": 9e-08, "output_cost_per_token": 1.1e-06, "cache_read_input_token_cost": 7e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 235929, - "max_tokens": 235929, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -69771,13 +70126,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-30b-a3b-instruct-2507": { - "input_cost_per_token": 4.815e-08, + "input_cost_per_token": 1e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 32000, - "max_tokens": 32000, + "max_output_tokens": 235929, + "max_tokens": 235929, "mode": "chat", - "output_cost_per_token": 1.9305e-07, + "output_cost_per_token": 3e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -70530,7 +70885,7 @@ "supports_web_search": false }, "openrouter/mistralai/mistral-nemo": { - "input_cost_per_token": 1.9e-08, + "input_cost_per_token": 2.9e-08, "output_cost_per_token": 3e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, @@ -73416,16 +73771,21 @@ "supports_web_search": true }, "openrouter/~anthropic/claude-haiku-latest": { - "cache_creation_input_token_cost": 1.25e-06, - "cache_creation_input_token_cost_above_1hr": 2e-06, - "cache_read_input_token_cost": 1e-07, - "input_cost_per_token": 1e-06, + "input_cost_per_token_above_100k_tokens": 5e-07, + "output_cost_per_token_above_100k_tokens": 2.5e-06, + "cache_read_input_token_cost_above_100k_tokens": 5e-08, + "cache_creation_input_token_cost_above_100k_tokens": 6.25e-07, + "cache_creation_input_token_cost_above_1hr_above_100k_tokens": 1e-06, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_read_input_token_cost": 1e-08, + "input_cost_per_token": 1e-07, "litellm_provider": "openrouter", - "max_input_tokens": 200000, - "max_output_tokens": 64000, - "max_tokens": 64000, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 5e-06, + "output_cost_per_token": 5e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -73462,7 +73822,7 @@ "openrouter/~anthropic/claude-sonnet-latest": { "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, - "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 2e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, @@ -73684,8 +74044,8 @@ "openrouter/~openai/gpt-sol-latest": { "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_272k_tokens": 5e-06, - "cache_read_input_token_cost": 2e-07, - "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, "input_cost_per_token": 2e-06, "input_cost_per_token_above_272k_tokens": 4e-06, "litellm_provider": "openrouter", @@ -79771,7 +80131,7 @@ "supports_vision": true, "supports_pdf_input": true, "supports_audio_input": false, - "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost": 1e-07, "supports_prompt_caching": true, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, @@ -80109,33 +80469,10 @@ "supports_tool_choice": true, "supports_vision": true }, - "xai.grok-4.7": { - "supports_regex_lookaround": false, - "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_flex": 2.5e-07, - "cache_read_input_token_cost_priority": 8.75e-07, - "input_cost_per_token": 2e-06, - "input_cost_per_token_flex": 1e-06, - "input_cost_per_token_priority": 3.5e-06, - "litellm_provider": "bedrock_converse", - "max_input_tokens": 500000, - "max_output_tokens": 500000, - "max_tokens": 500000, - "mode": "chat", - "output_cost_per_token": 6e-06, - "output_cost_per_token_flex": 3e-06, - "output_cost_per_token_priority": 1.05e-05, - "source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json", - "supports_function_calling": true, - "supports_prompt_caching": false, - "supports_reasoning": true, - "supports_tool_choice": true, - "supports_vision": true - }, "openrouter/anthropic/claude-sonnet-5.5:batch": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, - "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost": 5e-08, "input_cost_per_token": 1e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, @@ -80155,12 +80492,12 @@ "supports_web_search": true }, "baseten/deepseek-ai/DeepSeek-V4.1-Flash-Fast": { - "cache_read_input_token_cost": 1.4e-07, + "cache_read_input_token_cost": 1.4e-08, "input_cost_per_token": 6e-07, "litellm_provider": "baseten", "max_input_tokens": 1048576, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 2.4e-06, "source": "https://inference.baseten.co/v1/models", @@ -80177,25 +80514,31 @@ "cache_creation_input_token_cost_above_272k_tokens_batches": 2.5e-06, "cache_creation_input_token_cost_above_272k_tokens_flex": 2.5e-06, "cache_creation_input_token_cost_above_272k_tokens_priority": 1e-05, + "cache_creation_input_token_cost_above_272k_tokens_ultrafast": 3e-05, "cache_creation_input_token_cost_batches": 1.25e-06, "cache_creation_input_token_cost_flex": 1.25e-06, "cache_creation_input_token_cost_priority": 5e-06, + "cache_creation_input_token_cost_ultrafast": 1.5e-05, "cache_read_input_token_cost": 1e-07, "cache_read_input_token_cost_above_272k_tokens": 2e-07, "cache_read_input_token_cost_above_272k_tokens_batches": 1e-07, "cache_read_input_token_cost_above_272k_tokens_flex": 1e-07, "cache_read_input_token_cost_above_272k_tokens_priority": 4e-07, + "cache_read_input_token_cost_above_272k_tokens_ultrafast": 1.2e-06, "cache_read_input_token_cost_batches": 5e-08, "cache_read_input_token_cost_flex": 5e-08, "cache_read_input_token_cost_priority": 2e-07, + "cache_read_input_token_cost_ultrafast": 6e-07, "input_cost_per_token": 2e-06, "input_cost_per_token_above_272k_tokens": 4e-06, "input_cost_per_token_above_272k_tokens_batches": 2e-06, "input_cost_per_token_above_272k_tokens_flex": 2e-06, "input_cost_per_token_above_272k_tokens_priority": 8e-06, + "input_cost_per_token_above_272k_tokens_ultrafast": 2.4e-05, "input_cost_per_token_batches": 1e-06, "input_cost_per_token_flex": 1e-06, "input_cost_per_token_priority": 4e-06, + "input_cost_per_token_ultrafast": 1.2e-05, "litellm_provider": "openai", "max_input_tokens": 922000, "max_output_tokens": 128000, @@ -80206,9 +80549,11 @@ "output_cost_per_token_above_272k_tokens_batches": 7.5e-06, "output_cost_per_token_above_272k_tokens_flex": 7.5e-06, "output_cost_per_token_above_272k_tokens_priority": 3e-05, + "output_cost_per_token_above_272k_tokens_ultrafast": 9e-05, "output_cost_per_token_batches": 5e-06, "output_cost_per_token_flex": 5e-06, "output_cost_per_token_priority": 2e-05, + "output_cost_per_token_ultrafast": 6e-05, "regional_processing_uplift_multiplier_eu": 1.1, "regional_processing_uplift_multiplier_us": 1.1, "search_context_cost_per_query": { @@ -80285,39 +80630,6 @@ "supports_sampling_params": false, "supports_xhigh_reasoning_effort": true }, - "openai.gpt-6.1-sol": { - "cache_creation_input_token_cost": 2.5e-06, - "cache_creation_input_token_cost_above_272k_tokens": 5e-06, - "cache_read_input_token_cost": 1e-07, - "cache_read_input_token_cost_above_272k_tokens": 2e-07, - "input_cost_per_token": 2e-06, - "input_cost_per_token_above_272k_tokens": 4e-06, - "litellm_provider": "bedrock_converse", - "max_input_tokens": 1050000, - "max_output_tokens": 131072, - "max_tokens": 131072, - "mode": "chat", - "output_cost_per_token": 1e-05, - "output_cost_per_token_above_272k_tokens": 1.5e-05, - "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html", - "supported_modalities": [ - "text", - "image" - ], - "supported_output_modalities": [ - "text" - ], - "supports_function_calling": true, - "supports_max_reasoning_effort": true, - "supports_minimal_reasoning_effort": false, - "supports_none_reasoning_effort": false, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_sampling_params": false, - "supports_xhigh_reasoning_effort": true - }, "bedrock_mantle/openai.gpt-6.1-sol": { "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, @@ -80349,6 +80661,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -80740,7 +81053,7 @@ "cache_read_input_token_cost": 7e-08, "input_cost_per_token": 6.8e-07, "litellm_provider": "openrouter", - "max_input_tokens": 524288, + "max_input_tokens": 1048576, "max_output_tokens": 262144, "max_tokens": 262144, "mode": "chat", @@ -81529,5 +81842,66 @@ "thinking_always_on": false, "prompt_cache_min_tokens": 512, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing" + }, + "openrouter/inclusionai/ling-3.0-flash-sante": { + "cache_read_input_token_cost": 8.4e-09, + "input_cost_per_token": 4.2e-08, + "litellm_provider": "openrouter", + "max_input_tokens": 262144, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 1.232e-07, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": false, + "supports_tool_choice": true, + "supports_vision": false, + "supports_web_search": false + }, + "openrouter/stepfun/step-5-preview": { + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 1e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 2.7e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": false, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": false + }, + "us.twelvelabs.pegasus-1-5-v1:0": { + "input_cost_per_video_per_second": 0.00049, + "litellm_provider": "bedrock", + "max_output_tokens": 98304, + "max_tokens": 98304, + "mode": "chat", + "output_cost_per_token": 7.5e-06, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", + "supports_video_input": true + }, + "global.twelvelabs.pegasus-1-5-v1:0": { + "input_cost_per_video_per_second": 0.00049, + "litellm_provider": "bedrock", + "max_output_tokens": 98304, + "max_tokens": 98304, + "mode": "chat", + "output_cost_per_token": 7.5e-06, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", + "supports_video_input": true } } diff --git a/litellm/ocr/dispatch.py b/litellm/ocr/dispatch.py index a6a6c5d0c50..4fb4a6c709a 100644 --- a/litellm/ocr/dispatch.py +++ b/litellm/ocr/dispatch.py @@ -9,7 +9,7 @@ from litellm.rust_bridge import runtime from litellm.rust_bridge.catalog import Route, RouteContext from litellm.rust_bridge.dispatch import PublicDispatch from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR -from litellm.rust_bridge.public_call import NativeCall, native_call, native_call_hook, optional_str +from litellm.rust_bridge.public_call import NativeCall, native_call, native_call_hook, optional_str, signature __all__ = ("aocr", "ocr") @@ -38,18 +38,21 @@ def _bind_request( ) +_OCR: Final = signature(_bind_request) + + def _public_request(name: str, args: tuple[object, ...], kwargs: Mapping[str, object]) -> NativeCall: try: - fields: Final = _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds the public arguments before native validation - return native_call(args, kwargs, fields) + _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds the public arguments before native validation except TypeError as error: raise TypeError(str(error).replace("_bind_request()", f"{name}()")) from None + return native_call(_OCR, args, kwargs) def _context(request: NativeCall) -> RouteContext: - prefix, separator, _ = str(request.bound["model"]).partition("/") - provider: Final = optional_str(request.bound.get("custom_llm_provider")) or (prefix if separator else None) - return RouteContext(Route.OCR, provider=provider, model=str(request.bound["model"])) + prefix, separator, _ = str(request.resolved["model"]).partition("/") + provider: Final = optional_str(request.resolved.get("custom_llm_provider")) or (prefix if separator else None) + return RouteContext(Route.OCR, provider=provider, model=str(request.resolved["model"])) _DISPATCH: Final = PublicDispatch( diff --git a/litellm/provider_endpoints_support_backup.json b/litellm/provider_endpoints_support_backup.json index 0e325bb61fe..d9a6feb625d 100644 --- a/litellm/provider_endpoints_support_backup.json +++ b/litellm/provider_endpoints_support_backup.json @@ -2354,7 +2354,7 @@ } }, "voyage": { - "display_name": "Voyage AI (`voyage`)", + "display_name": "VoyageAI by MongoDB (`voyage`)", "url": "https://docs.litellm.ai/docs/providers/voyage", "endpoints": { "chat_completions": false, diff --git a/litellm/proxy/_experimental/mcp_server/README.md b/litellm/proxy/_experimental/mcp_server/README.md index 1123a9779ac..7d1f58965b0 100644 --- a/litellm/proxy/_experimental/mcp_server/README.md +++ b/litellm/proxy/_experimental/mcp_server/README.md @@ -13,3 +13,20 @@ Listing failures retain per-server outcome metadata. An incomplete upstream cata For a scoped rollout, keep the previous source build serving the control pool and send only selected clients to a separate candidate pool. All candidate replicas must share configuration and salt. Verify page one on one candidate replica and continuation on another, plus a fresh listing and tool call on the control pool. Do not mirror tool calls between pools For rollback, stop sending new requests to the candidate pool, drain its in-flight operations, and return selected clients to the control pool. Clients must discard candidate cursors and start a fresh listing when crossing versions; older gateways do not validate these cursors. Keep the registry-revision migration installed when rolling back pagination. Verify a fresh listing and a tool call after switching pools + +## Elicitation + +Elicitation is disabled by default. Enable it per upstream in `config.yaml`; `allow_elicitation` is currently YAML-only and is not editable through the Admin UI or database-backed server API + +```yaml +mcp_servers: + interactive: + url: https://mcp.example.com/mcp + transport: http + allow_elicitation: true + timeout: 60 +``` + +The downstream MCP client must advertise the requested form or URL capability during initialization. Use the gateway's legacy SSE endpoint (`/mcp/sse`) for the verified interactive form and URL relay path. The current Streamable HTTP endpoint (`/mcp`) can lose initialization state before a tool call, so even a client that advertised support receives an explicit elicitation error instead of an input request. Successful interactive relay over Streamable HTTP is not currently verified. Stateless calls and LLM tool bridges with no downstream MCP client also receive explicit errors + +The relay wait uses the upstream server's existing `timeout` setting, or `LITELLM_MCP_CLIENT_TIMEOUT` (60 seconds by default). The enclosing tool call also retains its existing timeout. Unsupported modes, disconnects and relay failures return errors, never a fabricated user decline. Actual user accept, decline and cancel responses are preserved; cancellation stops the relay diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index c9da8c1fcaa..b4845948ede 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -72,6 +72,7 @@ def _configuration_identity(server: MCPServer) -> str: "token_url", "registration_url", "authorization_response_iss_parameter_supported", + "client_id_metadata_document_supported", ) ) | (frozenset() if server.issuer_is_anchored else frozenset(("issuer",))), diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index e2d26054c80..823f0ed675c 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -223,6 +223,7 @@ class _OAuthCredentialAccessToken(TypedDict): class OAuthCredentialPayload(_OAuthCredentialAccessToken, total=False): identity_binding_proof: ReadOnly[str] + cimd_client_id: ReadOnly[str] type: str refresh_token: str expires_at: str @@ -1765,6 +1766,7 @@ async def store_user_oauth_credential( scopes: list[str] | None = None, skip_byok_guard: bool = False, identity_binding_proof: str | None = None, + cimd_client_id: str | None = None, ) -> None: """Persist an OAuth2 access token for a user+server pair. @@ -1782,6 +1784,7 @@ async def store_user_oauth_credential( "access_token": access_token, "connected_at": datetime.now(timezone.utc).isoformat(), **({"identity_binding_proof": identity_binding_proof} if identity_binding_proof else {}), + **({"cimd_client_id": cimd_client_id} if cimd_client_id else {}), } if refresh_token: payload["refresh_token"] = refresh_token @@ -2063,6 +2066,7 @@ async def refresh_user_oauth_token( auth_method=getattr(server, "token_endpoint_auth_method", None), client_id=client_id, client_secret=client_secret, + cimd_client_id=cred.get("cimd_client_id"), ) token_data: Final[dict[str, str]] = { "grant_type": "refresh_token", @@ -2132,6 +2136,7 @@ async def refresh_user_oauth_token( expires_in=expires_in, scopes=scopes, identity_binding_proof=binding_proof, + cimd_client_id=cred.get("cimd_client_id"), skip_byok_guard=True, # Row is already OAuth2; skip the extra find_unique check ) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index da6ce954b1d..0253719f410 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -79,9 +79,13 @@ from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( enforce_oauth_identity_binding, ) from litellm.proxy._experimental.mcp_server.oauth_utils import ( + CIMD_METADATA_PATH, TOKEN_NO_CACHE_HEADERS, build_upstream_oauth2_token_request, + get_cimd_client_id, + get_cimd_document_url, get_request_base_url, + needs_cimd_discovery, oauth_client_registration_matches, resolve_upstream_resource, validate_trusted_redirect_uri, @@ -638,6 +642,7 @@ async def _store_per_user_token_server_side( user_id: str, token_response: dict[str, Any], identity_binding_proof: str | None = None, + cimd_client_id: str | None = None, ) -> None: """Persist the OAuth token server-side and warm the Redis cache. @@ -680,6 +685,7 @@ async def _store_per_user_token_server_side( expires_in=expires_in, scopes=scopes, identity_binding_proof=identity_binding_proof, + **({"cimd_client_id": cimd_client_id} if cimd_client_id is not None else {}), ) verbose_logger.info( "_store_per_user_token_server_side: stored token for user=%s server=%s", @@ -778,20 +784,20 @@ async def _server_with_oauth_endpoints( mcp_server: MCPServer, needed_endpoint: Callable[[MCPServer], str | None], ) -> MCPServer: - """Join deferred OAuth discovery only when the endpoint this caller needs is still missing. + """Join deferred discovery for missing endpoints or unknown public-client metadata. - Admin-entered endpoints live on ``configured_*`` after an anchored issuer empties the - resolved fields. A caller whose needed endpoint already resolves never awaits discovery - and cannot 503 over a leftover pin. A server still missing it joins the deferred task; - no slot is a no-op and the caller 400s. + Capability discovery is optional when manual endpoints already resolve; the manager + preserves those endpoints if discovery fails. No discovery slot remains a no-op. """ - if needed_endpoint(mcp_server) is not None: + if needed_endpoint(mcp_server) is not None and not needs_cimd_discovery(mcp_server): return mcp_server from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # circular import with mcp_server_manager at module load global_mcp_server_manager, ) - return await global_mcp_server_manager.ensure_oauth_metadata_discovered(mcp_server) + if needed_endpoint(mcp_server) is None: + return await global_mcp_server_manager.ensure_oauth_metadata_discovered(mcp_server) + return await global_mcp_server_manager.ensure_oauth_metadata_discovered(mcp_server, needed_endpoint=needed_endpoint) def _raise_unless_oauth2_discovery_server( @@ -979,7 +985,8 @@ async def authorize_with_server( binding: Final = resolved_server.oauth_identity_binding enforce_binding: Final = binding is not None and binding.mode == "enforce" - if enforce_binding: + cimd_client_id: Final = get_cimd_client_id(resolved_server) + if enforce_binding or cimd_client_id: _require_s256_pkce(code_challenge, code_challenge_method) if resolved_server.is_dcr_bridge: @@ -1043,7 +1050,7 @@ async def authorize_with_server( relay_state: Final = secrets.token_urlsafe(_OAUTH_STATE_HANDLE_BYTES) params: Final = { - "client_id": resolved_server.client_id if resolved_server.client_id else client_id, + "client_id": resolved_server.client_id or cimd_client_id or client_id, "redirect_uri": f"{request_base_url}/callback", "state": relay_state, "response_type": response_type or "code", @@ -1080,7 +1087,37 @@ def _token_credential_source(mcp_server: MCPServer) -> CredentialSource: """Mirrors the resolved-client rule in :func:`exchange_token_with_server`: when the server has a stored client_id the gateway presents its own credentials upstream, so a credential rejection is the operator's fault, not the caller's.""" - return "gateway_stored" if mcp_server.client_id else "caller_supplied" + return "gateway_stored" if mcp_server.client_id or get_cimd_client_id(mcp_server) else "caller_supplied" + + +async def _saved_cimd_refresh_client_id( + server: MCPServer, user_id: str | None, refresh_token: str | None +) -> str | None: + """Reuse only the client identity bound to this caller's presented refresh grant.""" + if ( + not user_id + or not refresh_token + or not server.needs_user_oauth_token + or server.auth_type != MCPAuth.oauth2 + or server.client_id + or server.client_secret + ): + return None + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server.db import get_user_oauth_credential + + if proxy_server.prisma_client is None: + return None + try: + credential: Final = await get_user_oauth_credential(proxy_server.prisma_client, user_id, server.server_id) + except Exception: # noqa: BLE001 # optional storage must not prevent a caller-owned OAuth exchange + return None + if credential is None: + return None + stored_refresh: Final = credential.get("refresh_token") + if not stored_refresh or not secrets.compare_digest(stored_refresh.encode(), refresh_token.encode()): + return None + return credential.get("cimd_client_id") async def exchange_token_with_server( @@ -1117,6 +1154,17 @@ async def exchange_token_with_server( ), ) + request_user_id: Final = ( + await extract_user_id_from_request(request) + if resolved_server.needs_user_oauth_token or resolved_server.oauth_identity_binding is not None + else None + ) + cimd_client_id: Final = ( + await _saved_cimd_refresh_client_id(resolved_server, request_user_id, refresh_token) + if grant_type == "refresh_token" + else None + ) or get_cimd_client_id(resolved_server) + # The id, secret, and token-endpoint auth method must come from the same source. When the # server-side client_id wins, falling back to the caller's secret pairs the persisted client # with a foreign secret; the register short-circuit hands clients a placeholder secret @@ -1138,16 +1186,11 @@ async def exchange_token_with_server( auth_method=resolved_auth_method, client_id=resolved_client_id, client_secret=resolved_client_secret, + cimd_client_id=cimd_client_id, ) except TokenEndpointAuthConfigError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc - request_user_id: Final = ( - await extract_user_id_from_request(request) - if resolved_server.needs_user_oauth_token or resolved_server.oauth_identity_binding is not None - else None - ) - bridge_identity: _BridgeAuthorizationCode | None = None bridge_mint_ready: _BridgeMintReady | None = None bridge_upstream_refresh: SecretStr | None = None @@ -1265,7 +1308,7 @@ async def exchange_token_with_server( except httpx.HTTPStatusError as exc: fault: Final = classify_upstream_token_rejection( exc.response, - credential_source=_token_credential_source(resolved_server), + credential_source="gateway_stored" if cimd_client_id else _token_credential_source(resolved_server), log_context=resolved_server.server_id, ) upstream_rejected_bridge_refresh: Final = ( @@ -1336,6 +1379,7 @@ async def exchange_token_with_server( user_id=user_id, token_response=token_response, identity_binding_proof=binding_proof, + **({"cimd_client_id": cimd_client_id} if cimd_client_id is not None else {}), ) else: verbose_logger.warning( @@ -1991,8 +2035,21 @@ async def register_client_with_server( ), ) + cimd_client_id: Final = get_cimd_client_id(resolved_server) + if cimd_client_id: + return { + "client_id": cimd_client_id, + "token_endpoint_auth_method": "none", + "redirect_uris": client_facing_redirect_uris, + } registration_url: Final = resolved_server.effective_registration_url if registration_url is None: + if resolved_server.client_id_metadata_document_supported and resolved_server.is_gateway_managed_oauth2: + raise HTTPException( + status_code=400, + detail="CIMD requires a stable HTTPS PROXY_BASE_URL and public-client authentication; " + "configure these or provide a pre-registered OAuth client", + ) return dummy_return bridge_relay: Final = _dcr_bridge_relays_client_registration(resolved_server) @@ -3162,3 +3219,25 @@ async def register_client(request: Request, mcp_server_name: str | None = None): client_redirect_uris=client_redirect_uris, client_application_type=client_application_type, ) + + +@router.get(CIMD_METADATA_PATH, include_in_schema=False) +async def oauth_client_metadata() -> JSONResponse: + from mcp.shared.auth import OAuthClientInformationFull + from pydantic import AnyUrl + + document_url: Final = get_cimd_document_url() + if document_url is None: + raise HTTPException(status_code=404, detail="CIMD requires a configured HTTPS PROXY_BASE_URL") + base_url: Final = document_url.removesuffix(CIMD_METADATA_PATH) + metadata: Final = OAuthClientInformationFull( + client_id=document_url, + client_name="LiteLLM MCP Gateway", + redirect_uris=[AnyUrl(f"{base_url}/callback")], + token_endpoint_auth_method="none", + grant_types=["authorization_code", "refresh_token"], + response_types=["code"], + ) + return JSONResponse( + metadata.model_dump(mode="json", exclude_none=True), headers={"Cache-Control": "public, max-age=300"} + ) diff --git a/litellm/proxy/_experimental/mcp_server/elicitation_handler.py b/litellm/proxy/_experimental/mcp_server/elicitation_handler.py index 57d2d86d506..245a0612435 100644 --- a/litellm/proxy/_experimental/mcp_server/elicitation_handler.py +++ b/litellm/proxy/_experimental/mcp_server/elicitation_handler.py @@ -1,30 +1,35 @@ -""" -MCP Elicitation Handler -Handles `elicitation/create` requests from upstream MCP servers by either: -1. Relaying them to the connected downstream MCP client (if it supports elicitation) -2. Returning a decline/error response (if no downstream client or unsupported) -Supports both Form mode (structured data collection) and URL mode (external URL -navigation for sensitive interactions like OAuth). -MCP Spec Reference: - https://modelcontextprotocol.io/specification/2025-11-25/client/elicitation -""" +"""Relay upstream elicitation through the initiating downstream MCP request.""" -from typing import TYPE_CHECKING, Final, Protocol, Union +from __future__ import annotations + +import asyncio +from typing import TYPE_CHECKING, Final, Protocol from litellm._logging import verbose_logger +from litellm.constants import MCP_CLIENT_TIMEOUT if TYPE_CHECKING: from mcp.types import ( + INTERNAL_ERROR, + INVALID_REQUEST, + REQUEST_TIMEOUT, + ClientCapabilities, + ElicitRequestedSchema, ElicitRequestFormParams, ElicitRequestParams, ElicitRequestURLParams, ElicitResult, ErrorData, + RequestId, ) -# Guard imports that require the mcp package try: from mcp.types import ( + INTERNAL_ERROR, + INVALID_REQUEST, + REQUEST_TIMEOUT, + ClientCapabilities, + ElicitRequestedSchema, ElicitRequestFormParams, ElicitRequestParams, ElicitRequestURLParams, @@ -38,136 +43,65 @@ except ImportError: class _DownstreamElicitSession(Protocol): - """The downstream MCP client session methods this module relays elicitation requests through.""" + async def elicit_url( + self, message: str, url: str, elicitation_id: str, related_request_id: RequestId | None = None + ) -> ElicitResult: ... - async def elicit_url(self, message: str, url: str, elicitation_id: str) -> "ElicitResult": ... - - async def elicit_form(self, message: str, requested_schema: dict[str, object]) -> "ElicitResult": ... - - async def elicit(self, message: str, requested_schema: dict[str, object]) -> "ElicitResult": ... + async def elicit_form( + self, message: str, requested_schema: ElicitRequestedSchema, related_request_id: RequestId | None = None + ) -> ElicitResult: ... async def handle_elicitation_request( context: object, - params: "ElicitRequestParams", + params: ElicitRequestParams, downstream_session: _DownstreamElicitSession | None = None, - downstream_capabilities: object = None, -) -> Union["ElicitResult", "ErrorData"]: - """ - Handle an MCP elicitation/create request from an upstream MCP server. - In Gateway mode (Mode A), we relay the elicitation request to the - connected downstream client if they declared elicitation capabilities. - In Tool Bridge mode (Mode B), there's no persistent downstream MCP - client, so we return a decline response. - Args: - context: MCP RequestContext from the upstream server connection. - params: The ElicitRequestParams (either form or URL mode). - downstream_session: The ServerSession to the downstream client, - if available (for relaying). - downstream_capabilities: The downstream client's declared - capabilities, used to check elicitation support. - Returns: - ElicitResult with the user's response, or ErrorData on failure. - """ + downstream_capabilities: ClientCapabilities | None = None, + related_request_id: RequestId | None = None, + timeout: float = MCP_CLIENT_TIMEOUT, +) -> ElicitResult | ErrorData: if not MCP_ELICITATION_AVAILABLE: - return ErrorData( - code=-1, - message="MCP elicitation is not available (mcp package not installed)", - ) + return ErrorData(code=INTERNAL_ERROR, message="MCP elicitation is not available") + if downstream_session is None: + return ErrorData(code=INVALID_REQUEST, message="MCP elicitation requires a connected downstream MCP client") try: - mode: Final = getattr(params, "mode", "form") - verbose_logger.info( - "MCP elicitation: received request mode=%s, message=%s", - mode, - getattr(params, "message", ""), + return await asyncio.wait_for( + _relay_elicitation_to_downstream(params, downstream_session, downstream_capabilities, related_request_id), + timeout=timeout, ) - # Check if we have a downstream session to relay to - if downstream_session is not None: - return await _relay_elicitation_to_downstream( - params=params, - downstream_session=downstream_session, - downstream_capabilities=downstream_capabilities, - ) - # No downstream session — we're in Tool Bridge mode - # or the client doesn't support elicitation - verbose_logger.info("MCP elicitation: no downstream session available, declining") - return ElicitResult( - action="decline", - ) - except Exception as e: - verbose_logger.exception("MCP elicitation handler failed: %s", e) + except asyncio.TimeoutError: + return ErrorData(code=REQUEST_TIMEOUT, message="MCP elicitation timed out waiting for the downstream client") + except Exception: + verbose_logger.warning("MCP elicitation: downstream relay failed") return ErrorData( - code=-1, - message=f"Elicitation failed: {e}", + code=INTERNAL_ERROR, message="MCP elicitation failed while communicating with the downstream client" ) async def _relay_elicitation_to_downstream( - params: "ElicitRequestParams", + params: ElicitRequestParams, downstream_session: _DownstreamElicitSession, - downstream_capabilities: object = None, -) -> Union["ElicitResult", "ErrorData"]: - """ - Relay an elicitation request to the downstream MCP client. - Uses the ServerSession's elicit_form() or elicit_url() methods to - send the elicitation request back to the connected client. - Args: - params: The elicitation request parameters. - downstream_session: The ServerSession connected to the downstream client. - downstream_capabilities: Client capabilities to check support. - Returns: - ElicitResult from the downstream client. - """ - mode: Final = getattr(params, "mode", "form") - # Check if the downstream client supports the requested mode - if downstream_capabilities is not None: - elicit_caps: Final[object] = getattr(downstream_capabilities, "elicitation", None) - if elicit_caps is None: - verbose_logger.info("MCP elicitation: downstream client does not support elicitation") - return ElicitResult(action="decline") - if mode == "url": - url_cap: Final[object] = getattr(elicit_caps, "url", None) - if url_cap is None: - verbose_logger.info("MCP elicitation: downstream client does not support URL mode") - return ElicitResult(action="decline") - if mode == "form": - form_cap: Final[object] = getattr(elicit_caps, "form", None) - if form_cap is None: - verbose_logger.info("MCP elicitation: downstream client does not support form mode") - return ElicitResult(action="decline") - try: - if mode == "url" and isinstance(params, ElicitRequestURLParams): - # URL mode: relay URL to client for external navigation - verbose_logger.info( - "MCP elicitation: relaying URL mode to downstream, url=%s", - getattr(params, "url", ""), - ) - result = await downstream_session.elicit_url( - message=params.message, - url=params.url, - elicitation_id=params.elicitation_id, - ) - elif isinstance(params, ElicitRequestFormParams): - # Form mode: relay structured form to client - verbose_logger.info("MCP elicitation: relaying form mode to downstream") - result = await downstream_session.elicit_form( - message=params.message, - requested_schema=params.requested_schema, - ) - else: - # Fallback for generic ElicitRequestParams — pass an empty schema - # since elicit() requires requested_schema as a positional arg. - verbose_logger.info("MCP elicitation: relaying generic elicitation to downstream") - result = await downstream_session.elicit( - message=getattr(params, "message", ""), - requested_schema=getattr(params, "requested_schema", {}), - ) - verbose_logger.info( - "MCP elicitation: downstream responded with action=%s", - getattr(result, "action", "unknown"), + downstream_capabilities: ClientCapabilities | None = None, + related_request_id: RequestId | None = None, +) -> ElicitResult | ErrorData: + capabilities: Final = downstream_capabilities.elicitation if downstream_capabilities is not None else None + if capabilities is None: + return ErrorData(code=INVALID_REQUEST, message="Downstream client has not advertised elicitation support") + if isinstance(params, ElicitRequestURLParams): + if capabilities.url is None: + return ErrorData(code=INVALID_REQUEST, message="Downstream client does not support URL elicitation") + return await downstream_session.elicit_url( + message=params.message, + url=params.url, + elicitation_id=params.elicitation_id, + related_request_id=related_request_id, ) - return result - except Exception as e: - verbose_logger.warning("MCP elicitation: failed to relay to downstream: %s", e) - # If relay fails, decline gracefully - return ElicitResult(action="decline") + if capabilities.form is None and capabilities.url is not None: + return ErrorData(code=INVALID_REQUEST, message="Downstream client does not support form elicitation") + if not isinstance(params, ElicitRequestFormParams): + return ErrorData(code=INVALID_REQUEST, message="Unsupported MCP elicitation parameters") + return await downstream_session.elicit_form( + message=params.message, + requested_schema=params.requested_schema, + related_request_id=related_request_id, + ) diff --git a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py index f52d2a006d2..13c3d265429 100644 --- a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py +++ b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py @@ -11,7 +11,9 @@ from mcp.types import ( ErrorData, ) +from litellm.constants import MCP_CLIENT_TIMEOUT from litellm.proxy._experimental.mcp_server.contracts import OperationContext +from litellm.proxy._experimental.mcp_server.mcp_context import get_active_mcp_request_ctx from litellm.proxy._types import UserAPIKeyAuth @@ -62,11 +64,19 @@ def create_sampling_callback( return callback -def create_elicitation_callback() -> ElicitationCallback: +def create_elicitation_callback(timeout: float | None = None) -> ElicitationCallback: from litellm.proxy._experimental.mcp_server.server import get_active_mcp_session downstream_session: Final = get_active_mcp_session() - downstream_capabilities: Final = getattr(downstream_session, "capabilities", None) + request: Final = get_active_mcp_request_ctx() + client_params: Final = downstream_session.client_params if downstream_session is not None else None + downstream_capabilities: Final = ( + client_params.capabilities.model_copy(deep=True) if client_params is not None else None + ) + related_request_id: Final = ( + request.request_id if request is not None and request.session is downstream_session else None + ) + relay_timeout: Final = timeout if timeout is not None else MCP_CLIENT_TIMEOUT async def callback(context: object, params: ElicitRequestParams) -> ElicitResult | ErrorData: from litellm.proxy._experimental.mcp_server.elicitation_handler import handle_elicitation_request @@ -76,6 +86,8 @@ def create_elicitation_callback() -> ElicitationCallback: params=params, downstream_session=downstream_session, downstream_capabilities=downstream_capabilities, + related_request_id=related_request_id, + timeout=relay_timeout, ) return callback diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 15145a20e33..8185ca508c0 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -105,6 +105,7 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import ( # noqa: F401 _redact_mcp_resource_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export canonicalize_url_identity, get_byok_www_authenticate, + needs_cimd_discovery, redact_mcp_resource_url, ) from litellm.proxy._experimental.mcp_server.outbound_credentials import ( @@ -452,6 +453,7 @@ class _AuthorizationServerMetadataPayload(TypedDict, total=False): authorization_endpoint: str token_endpoint: str registration_endpoint: str + client_id_metadata_document_supported: ReadOnly[object] scopes_supported: Sequence[str] grant_types_supported: Sequence[str] token_endpoint_auth_methods_supported: Sequence[str] @@ -767,14 +769,16 @@ def _flow_endpoints_missing( return authorization_url is None or token_url is None -def oauth_endpoints_unresolved(server: MCPServer) -> bool: - """``_flow_endpoints_missing`` over a built registry entry, for the reload fast-path check. +def oauth_endpoints_unresolved(server: MCPServer, *, include_client_metadata: bool = True) -> bool: + """Whether endpoint or eligible client metadata discovery is still pending. The flow comes from ``effective_oauth2_flow``, the one column-first, shape-fallback judge every flow decision uses, not from the raw column: a legacy row the startup backfill deliberately left unstamped (the ambiguous M2M shape) serves M2M at request time, and reading the bare column here would classify it as interactive-missing-endpoints and re-run discovery on every reload. """ + if include_client_metadata and needs_cimd_discovery(server): + return True if ( server.auth_type == MCPAuth.oauth2_token_exchange and server.token_exchange_profile == "entra_obo" @@ -868,6 +872,8 @@ def carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serve may_carry: Final = _endpoints_corroborate_authorization_url( previous_server.authorization_url, new_server.authorization_url ) + if may_carry and new_server.client_id_metadata_document_supported is None: + new_server.client_id_metadata_document_supported = previous_server.client_id_metadata_document_supported if may_carry and new_server.issuer is None: new_server.issuer = previous_server.issuer new_server.authorization_response_iss_parameter_supported = ( # rebind-ok: publish on the existing rebuild object @@ -908,7 +914,7 @@ def _restrict_discovery_to_corroborated_authorization_server( return metadata if _endpoints_corroborate_authorization_url(metadata.authorization_url, manual_authorization_url): return metadata - if not metadata.token_url and not metadata.registration_url: + if not metadata.token_url and not metadata.registration_url and not metadata.client_id_metadata_document_supported: return metadata bridge_note: Final = ( " The discovered registration_url is rejected with it, so this dcr_bridge server stays on the" @@ -926,7 +932,9 @@ def _restrict_discovery_to_corroborated_authorization_server( _normalized_authorize_endpoint(manual_authorization_url), bridge_note, ) - return metadata.model_copy(update={"token_url": None, "registration_url": None}) + return metadata.model_copy( + update={"token_url": None, "registration_url": None, "client_id_metadata_document_supported": False} + ) def _redacted_origin_list(urls: Sequence[str]) -> str: @@ -1727,12 +1735,12 @@ def _create_sampling_callback( return create_sampling_callback(user_api_key_auth, raw_headers, client_ip, operation_context) -def _create_elicitation_callback(): +def _create_elicitation_callback(timeout: float | None = None): if not MCP_ELICITATION_AVAILABLE: return None from litellm.proxy._experimental.mcp_server.legacy_callbacks import create_elicitation_callback - return create_elicitation_callback() + return create_elicitation_callback(timeout=timeout) def _record_mcp_guardrail_evaluations( @@ -2060,6 +2068,7 @@ class MCPServerManager: update={ "scopes": server.scopes or metadata.scopes, "issuer": server.issuer or discovered_issuer, + "client_id_metadata_document_supported": metadata.client_id_metadata_document_supported, "authorization_response_iss_parameter_supported": ( metadata.authorization_response_iss_parameter_supported if discovered_issuer is not None @@ -2224,12 +2233,27 @@ class MCPServerManager: if should_defer != has_slot: self._set_oauth_discovery_deferred(server.server_id, should_defer) - async def ensure_oauth_metadata_discovered(self, server: MCPServer, *, _retry_stale: bool = True) -> MCPServer: + async def ensure_oauth_metadata_discovered( + self, + server: MCPServer, + *, + needed_endpoint: Callable[[MCPServer], str | None] | None = None, + _retry_stale: bool = True, + ) -> MCPServer: return await self.catalog.resolve_oauth_metadata( - server, lambda selected: self._ensure_oauth_metadata_discovered(selected, _retry_stale=_retry_stale) + server, + lambda selected: self._ensure_oauth_metadata_discovered( + selected, needed_endpoint=needed_endpoint, _retry_stale=_retry_stale + ), ) - async def _ensure_oauth_metadata_discovered(self, server: MCPServer, *, _retry_stale: bool = True) -> MCPServer: + async def _ensure_oauth_metadata_discovered( + self, + server: MCPServer, + *, + needed_endpoint: Callable[[MCPServer], str | None] | None = None, + _retry_stale: bool = True, + ) -> MCPServer: """Join the bounded discovery task and return the resolved server. Concurrent callers share one task per server. A failed attempt remains @@ -2237,9 +2261,12 @@ class MCPServerManager: Args: server: The MCP server whose OAuth metadata must be resolved. + needed_endpoint: A caller-specific endpoint that may remain usable when + optional capability discovery fails. Returns: - The resolved server; the registered server when no discovery is + The resolved server; a configured caller endpoint remains usable on + optional capability-discovery failure. The registered server when no discovery is pending, or when discovery failed for a client-forwarded-token server, whose session consumes no discovered endpoint. @@ -2257,18 +2284,26 @@ class MCPServerManager: outcome: Final = await asyncio.shield(task) except asyncio.CancelledError: if task.cancelled() and not self._oauth_discovery_slot_is_current(server.server_id, generation): - return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale) + return await self._rejoin_oauth_metadata_discovery( + server, needed_endpoint=needed_endpoint, retry_stale=_retry_stale + ) raise match outcome: case _OAuthDiscoveryResolved(resolved_server): self.catalog.assert_current(resolved_server) return resolved_server case _OAuthDiscoveryStale(): - return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale) + return await self._rejoin_oauth_metadata_discovery( + server, needed_endpoint=needed_endpoint, retry_stale=_retry_stale + ) case _OAuthDiscoveryFailed(timed_out=timed_out): current: Final = self._registered_server(server) self.catalog.assert_current(current) - if current.is_client_forwarded_token: + if ( + current.is_client_forwarded_token + or not oauth_endpoints_unresolved(current, include_client_metadata=False) + or (needed_endpoint is not None and needed_endpoint(current) is not None) + ): return current server_ref: Final = current.alias or current.server_name or current.name or current.server_id reason: Final = "timed out" if timed_out else "returned incomplete metadata" @@ -2279,11 +2314,19 @@ class MCPServerManager: return assert_never(outcome) - async def _rejoin_oauth_metadata_discovery(self, server: MCPServer, *, retry_stale: bool) -> MCPServer: + async def _rejoin_oauth_metadata_discovery( + self, server: MCPServer, *, needed_endpoint: Callable[[MCPServer], str | None] | None = None, retry_stale: bool + ) -> MCPServer: if retry_stale: - return await self.ensure_oauth_metadata_discovered(server, _retry_stale=False) + return await self.ensure_oauth_metadata_discovered( + server, needed_endpoint=needed_endpoint, _retry_stale=False + ) current: Final = self._registered_server(server) - if not oauth_endpoints_unresolved(current) or current.is_client_forwarded_token: + if ( + not oauth_endpoints_unresolved(current, include_client_metadata=False) + or current.is_client_forwarded_token + or (needed_endpoint is not None and needed_endpoint(current)) + ): return current raise HTTPException(status_code=503, detail="OAuth metadata discovery changed repeatedly; retry shortly") @@ -2616,6 +2659,9 @@ class MCPServerManager: scopes=resolved_scopes, configured_scopes=tuple(configured_scopes) if configured_scopes else None, issuer=effective_issuer, + client_id_metadata_document_supported=( + gated_oauth_metadata.client_id_metadata_document_supported if gated_oauth_metadata else None + ), authorization_response_iss_parameter_supported=( gated_oauth_metadata.authorization_response_iss_parameter_supported if gated_oauth_metadata @@ -3194,6 +3240,9 @@ class MCPServerManager: scopes=resolved_scopes, configured_scopes=configured_scopes, issuer=effective_issuer, + client_id_metadata_document_supported=( + gated_oauth_metadata.client_id_metadata_document_supported if gated_oauth_metadata else None + ), authorization_response_iss_parameter_supported=( gated_oauth_metadata.authorization_response_iss_parameter_supported if gated_oauth_metadata else False ), @@ -3267,7 +3316,7 @@ class MCPServerManager: or "rfc8693", timeout=getattr(mcp_server, "timeout", None), max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None), - rpm=getattr(mcp_server, "rpm", None), + rpm=mcp_server.rpm, ) _warn_legacy_delegate_auth_if_applicable(new_server, source="database") if register_oauth_discovery: @@ -4201,7 +4250,11 @@ class MCPServerManager: if resolved_server.allow_sampling else None ), - elicitation_callback=(_create_elicitation_callback() if resolved_server.allow_elicitation else None), + elicitation_callback=( + _create_elicitation_callback(timeout=resolved_server.timeout) + if resolved_server.allow_elicitation + else None + ), ) _create_mcp_client = create_mcp_client @@ -5304,6 +5357,7 @@ class MCPServerManager: authorization_url=data.get("authorization_endpoint"), token_url=data.get("token_endpoint"), registration_url=data.get("registration_endpoint"), + client_id_metadata_document_supported=data.get("client_id_metadata_document_supported") is True, discovered_issuer=claimed_issuer if isinstance(claimed_issuer, str) and claimed_issuer else None, authorization_response_iss_parameter_supported=data.get( "authorization_response_iss_parameter_supported" @@ -5336,6 +5390,7 @@ class MCPServerManager: return MCPOAuthMetadata( authorization_url=f"{base}/oauth2/v2.0/authorize", token_url=f"{base}/oauth2/v2.0/token", + client_id_metadata_document_supported=False, ) @staticmethod diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index e60c4ef1c30..276cf250b30 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -130,6 +130,57 @@ def _resolve_proxy_base_url_env() -> str | None: return None +CIMD_METADATA_PATH: Final = "/oauth/client-metadata.json" + + +def get_cimd_document_url() -> str | None: + try: + configured: Final = _resolve_proxy_base_url_env() + parsed: Final = urlparse(configured or "") + except ValueError: + return None + if parsed.scheme != "https" or not parsed.hostname or parsed.username is not None or parsed.password is not None: + return None + return f"{configured}{CIMD_METADATA_PATH}" + + +def _can_use_cimd(server: "MCPServer") -> bool: + return ( + server.is_gateway_managed_oauth2 + and server.needs_user_oauth_token + and not server.client_id + and not server.client_secret + and server.token_endpoint_auth_method != "client_secret_basic" + ) + + +def needs_cimd_discovery(server: "MCPServer") -> bool: + """Resolve unknown client metadata support even when OAuth endpoints are configured.""" + return ( + getattr(server, "client_id_metadata_document_supported", False) is None + and _can_use_cimd(server) + and get_cimd_document_url() is not None + ) + + +def _deployment_prefers_cimd() -> bool: + from litellm.proxy.proxy_server import general_settings + + return general_settings.get("mcp_prefer_client_id_metadata_document") is True + + +def _dynamic_registration_takes_precedence(server: "MCPServer") -> bool: + return server.effective_registration_url is not None and not _deployment_prefers_cimd() + + +def get_cimd_client_id(server: "MCPServer") -> str | None: + if getattr(server, "client_id_metadata_document_supported", False) is not True or not _can_use_cimd(server): + return None + if _dynamic_registration_takes_precedence(server): + return None + return get_cimd_document_url() + + BYOK_RESOURCE_METADATA_PATH: Final = "/v1/mcp/oauth/protected-resource" @@ -748,6 +799,7 @@ def build_upstream_oauth2_token_request( auth_method: object, client_id: str | None, client_secret: str | None, + cimd_client_id: str | None = None, ) -> TokenEndpointClientAuth: """Client auth plus the RFC 8707 ``resource`` for one upstream plain-OAuth2 token request. @@ -757,10 +809,11 @@ def build_upstream_oauth2_token_request( authenticate as the caller's own client rather than the server's; ``resource`` always comes from the server, so no leg can choose or forget it. """ + selected_cimd_id: Final = cimd_client_id or get_cimd_client_id(mcp_server) client_auth: Final = build_token_endpoint_client_auth( - auth_method=normalize_token_endpoint_auth_method(auth_method), - client_id=client_id, - client_secret=client_secret, + auth_method=None if selected_cimd_id else normalize_token_endpoint_auth_method(auth_method), + client_id=selected_cimd_id or client_id, + client_secret=None if selected_cimd_id else client_secret, ) resource: Final = resolve_upstream_resource(mcp_server) if not resource: diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py index af1e82eab82..b926ca8a6a4 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py @@ -25,7 +25,7 @@ from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( RefreshTokenPresented, enforce_oauth_identity_binding, ) -from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request +from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request, get_cimd_client_id from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( OAuthToken, ) @@ -47,6 +47,7 @@ class CredentialPersist(Protocol): expires_in: int | None, scopes: tuple[str, ...] | None, identity_binding_proof: str | None = None, + cimd_client_id: str | None = None, ) -> None: ... @@ -112,12 +113,14 @@ class AuthorizationCodeRefresher: if not token_url: return None + cimd_client_id: Final = token.cimd_client_id or get_cimd_client_id(server) try: token_request: Final = build_upstream_oauth2_token_request( server, auth_method=server.token_endpoint_auth_method, client_id=server.client_id, client_secret=server.client_secret, + cimd_client_id=cimd_client_id, ) except TokenEndpointAuthConfigError as exc: verbose_logger.warning("MCP OAuth refresh misconfigured for server %s: %s", server_id, exc) @@ -155,22 +158,21 @@ class AuthorizationCodeRefresher: expires_in: Final = _parse_expires_in(body.get("expires_in")) scopes: Final = _parse_scopes(body.get("scope")) or token.scopes - if binding_proof is not None: - await self._persist( - user_id, - server_id, - access_token, - new_refresh, - expires_in, - scopes or None, - identity_binding_proof=binding_proof, - ) - else: - await self._persist(user_id, server_id, access_token, new_refresh, expires_in, scopes or None) + await self._persist( + user_id, + server_id, + access_token, + new_refresh, + expires_in, + scopes or None, + **({"identity_binding_proof": binding_proof} if binding_proof is not None else {}), + **({"cimd_client_id": cimd_client_id} if cimd_client_id is not None else {}), + ) return OAuthToken( access_token=access_token, expires_at=self._clock() + expires_in if expires_in is not None else None, refresh_token=new_refresh, scopes=scopes, identity_binding_proof=binding_proof, + cimd_client_id=cimd_client_id, ) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py index 0ac4296f498..b9f09c266d5 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py @@ -42,6 +42,7 @@ class OAuthToken: refresh_token: str | None = None scopes: tuple[str, ...] = () identity_binding_proof: str | None = None + cimd_client_id: str | None = None def __repr__(self) -> str: has_refresh: Final = self.refresh_token is not None diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py index 2897c3e8e4a..d8094d5facb 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py @@ -72,6 +72,7 @@ async def _persist_credential( expires_in: int | None, scopes: tuple[str, ...] | None, identity_binding_proof: str | None = None, + cimd_client_id: str | None = None, ) -> None: from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 store_user_oauth_credential, @@ -90,6 +91,7 @@ async def _persist_credential( scopes=list(scopes) if scopes else None, skip_byok_guard=True, identity_binding_proof=identity_binding_proof, + **({"cimd_client_id": cimd_client_id} if cimd_client_id is not None else {}), ) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py index 0f18931d118..56ab8e6bbdc 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py @@ -47,12 +47,14 @@ def _to_oauth_token(payload: Mapping[str, object]) -> OAuthToken | None: refresh_token: Final = payload.get("refresh_token") expires_at: Final = payload.get("expires_at") binding_proof: Final = payload.get("identity_binding_proof") + cimd_client_id: Final = payload.get("cimd_client_id") return OAuthToken( access_token=access_token, expires_at=_iso_to_epoch(expires_at) if isinstance(expires_at, str) else None, refresh_token=refresh_token if isinstance(refresh_token, str) else None, scopes=_to_scopes(payload.get("scopes")), identity_binding_proof=binding_proof if isinstance(binding_proof, str) else None, + cimd_client_id=cimd_client_id if isinstance(cimd_client_id, str) else None, ) diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 517c0916e13..68188800438 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -166,6 +166,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( name="mcp_discoverable", module_path="litellm.proxy._experimental.mcp_server.discoverable_endpoints", path_prefixes=( + "/oauth/client-metadata.json", "/.well-known/oauth-", "/.well-known/openid-configuration", "/.well-known/jwks.json", diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 9322fb77615..7dcca8f9ec0 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -11906,6 +11906,34 @@ "description": "When True (default), after sensitive data is detected and routed, all subsequent requests in the same session will continue routing to the same model.", "title": "Sticky Session Routing" }, + "stream_scope": { + "anyOf": [ + { + "enum": [ + "streaming", + "non_streaming", + "both" + ], + "type": "string" + }, + { + "additionalProperties": { + "enum": [ + "streaming", + "non_streaming", + "both" + ], + "type": "string" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "description": "Whether this guardrail runs on streaming requests, non-streaming requests, or both. A string applies to every configured mode. A map overrides named modes (pre_call, during_call, post_call, ...); omitted keys default to both. Unset means both, matching historical behavior.", + "title": "Stream Scope" + }, "template_id": { "anyOf": [ { @@ -14831,6 +14859,34 @@ "description": "When True (default), after sensitive data is detected and routed, all subsequent requests in the same session will continue routing to the same model.", "title": "Sticky Session Routing" }, + "stream_scope": { + "anyOf": [ + { + "enum": [ + "streaming", + "non_streaming", + "both" + ], + "type": "string" + }, + { + "additionalProperties": { + "enum": [ + "streaming", + "non_streaming", + "both" + ], + "type": "string" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "description": "Whether this guardrail runs on streaming requests, non-streaming requests, or both. A string applies to every configured mode. A map overrides named modes (pre_call, during_call, post_call, ...); omitted keys default to both. Unset means both, matching historical behavior.", + "title": "Stream Scope" + }, "template_id": { "anyOf": [ { diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index bc27b2352a9..bed2503d211 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, NamedTuple, Ty import httpx from pydantic import ( + AfterValidator, BaseModel, BeforeValidator, ConfigDict, @@ -2743,6 +2744,41 @@ class ScheduledJobStaggerSettings(LiteLLMPydanticObjectBase): ) +SPEND_LOGS_METADATA_ALWAYS_KEPT_FIELDS: Final = frozenset({"status", "cold_storage_object_key"}) + + +def _known_spend_logs_metadata_field(name: str) -> str: + if name not in SpendLogsMetadata.__annotations__: + raise ValueError(f"{name!r} is not a LiteLLM_SpendLogs.metadata field") + return name + + +SpendLogsMetadataFieldName: TypeAlias = Annotated[str, AfterValidator(_known_spend_logs_metadata_field)] + + +class SpendLogsMetadataFields(BaseModel): + model_config = ConfigDict(extra="forbid", frozen=True) + + include: tuple[SpendLogsMetadataFieldName, ...] | None = None + exclude: tuple[SpendLogsMetadataFieldName, ...] | None = None + + @model_validator(mode="after") + def _exactly_one_list(self) -> "SpendLogsMetadataFields": + if (self.include is None) == (self.exclude is None): + raise ValueError("set exactly one of 'include' or 'exclude'") + always_kept_excluded: Final = sorted(SPEND_LOGS_METADATA_ALWAYS_KEPT_FIELDS.intersection(self.exclude or ())) + if always_kept_excluded: + raise ValueError(f"{always_kept_excluded} are always kept and cannot be excluded") + return self + + def keeps(self, name: str) -> bool: + if name in SPEND_LOGS_METADATA_ALWAYS_KEPT_FIELDS: + return True + if self.include is not None: + return name in self.include + return name not in (self.exclude or ()) + + DEFAULT_RESPONSES_WEBSOCKET_SESSION_LIMIT_SECONDS: Final[float] = 3600.0 @@ -3091,6 +3127,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="If True, stores request messages and responses in spend logs. Default is False.", ) + spend_logs_metadata_fields: SpendLogsMetadataFields | None = Field( + None, + description="Which keys of LiteLLM_SpendLogs.metadata are written to the database. Set exactly one of 'include' (write only these keys) or 'exclude' (drop these keys). 'status' and 'cold_storage_object_key' are always written. Daily spend tables, budgets and logging callbacks still see every key. Unset writes every key", + ) disable_auto_add_proxy_admin_to_teams: bool | None = Field( None, description="By default, the user calling /team/new is automatically added to the new team as a team admin. If True, proxy admins are no longer auto-added; members explicitly listed in members_with_roles are unaffected. Default is False.", @@ -3185,6 +3225,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): ge=1, description="Number of trusted reverse proxies/load balancers in front of the gateway that append to X-Forwarded-For. When set (and mcp_trusted_proxy_ranges validates the direct peer), the client IP for MCP access control is read this many entries from the right of the chain instead of the spoofable leftmost value, defeating append-style X-Forwarded-For forgery.", ) + mcp_prefer_client_id_metadata_document: bool | None = Field( + None, + description="When true, a gateway-managed OAuth2 MCP server whose authorization server advertises Client ID Metadata Document support identifies itself with the gateway's public metadata document URL even when that authorization server also offers dynamic client registration. Requires a public HTTPS PROXY_BASE_URL the authorization server can fetch. Default false: dynamic client registration is used whenever the authorization server offers it, and the metadata document only when it does not.", + ) trusted_proxy_ranges: list[str] | None = Field( None, description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler, and whose X-Forwarded-For is used to attribute Admin UI sign-in attempts to a source address. Set it to an empty list when clients connect directly, so the peer address is the source. Left unset, or containing an entry that is not an address or CIDR range, the per-source sign-in limit is off.", diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index edd6028f07c..d31ce4dae03 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -60,6 +60,12 @@ def is_invalid_virtual_key_error(exception: BaseException | None) -> bool: return getattr(exception, INVALID_VIRTUAL_KEY_ERROR_MARKER, False) is True +def log_model_access_denial(exc: BaseException) -> None: + if not isinstance(exc, ModelAccessDeniedProxyException): + return + verbose_proxy_logger.warning(exc.sanitized_internal_message()) + + def mark_invalid_virtual_key_error(exception: ProxyException, is_invalid_virtual_key: bool) -> ProxyException: """Return an independently marked malformed-key exception after callback transformations.""" if not is_invalid_virtual_key or str(exception.code) != str(status.HTTP_401_UNAUTHORIZED): @@ -425,6 +431,9 @@ _BANNED_REQUEST_BODY_PARAMS: Final[tuple[str, ...]] = ( # so a caller-supplied value picks a transport and a callback surface the # admin did not choose. "rust", + # Deployment opt-in: a caller-supplied false would switch off identity + # forwarding and let the caller choose the `user` Fireworks sees. + "fireworks_forward_user_id", # SDK-only field; also rejected outright in is_request_body_safe. "model_list", "vertex_ai_credentials", diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index dda72928f26..69be1cebb3b 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -100,6 +100,7 @@ from litellm.proxy.auth.auth_utils import ( get_request_route_template, is_invalid_virtual_key_error, iter_request_fallback_targets, + log_model_access_denial, normalize_request_route, pre_db_read_auth_checks, request_dispatched_to_pass_through_endpoint, @@ -732,7 +733,7 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str except Exception as e: if is_invalid_virtual_key_error(e): raise WebSocketException(code=status.WS_1008_POLICY_VIOLATION) - verbose_proxy_logger.exception(e) + log_model_access_denial(e) await websocket.close(code=status.WS_1008_POLICY_VIOLATION) raise HTTPException(status_code=403, detail=str(e)) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 880e3fa2072..937cb6f40e5 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -72,10 +72,10 @@ from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attributio from litellm.proxy.route_llm_request import raise_if_required_body_param_missing from litellm.proxy.utils import PrismaClient, ProxyLogging, handle_exception_on_proxy, is_known_model from litellm.repositories.managed_batch_repository import ManagedBatchRepository -from litellm.repositories.table_repositories import ManagedFileRepository +from litellm.repositories.managed_file_repository import ManagedFileRepository from litellm.router import Router from litellm.types.llms.openai import LiteLLMBatchCreateRequest -from litellm.types.utils import LiteLLMBatch +from litellm.types.utils import LiteLLMBatch, LLMResponseTypes if TYPE_CHECKING: from prisma.models import LiteLLM_ManagedObjectTable @@ -96,6 +96,12 @@ def _request_tags(data: Mapping[str, object]) -> tuple[str, ...] | None: return request_tags_from_metadata(_METADATA_ADAPTER.validate_python(metadata)) +def _require_batch_response(response: LLMResponseTypes) -> LiteLLMBatch: + if not isinstance(response, LiteLLMBatch): + raise TypeError("Batch endpoint received a non-batch response") + return response + + def _litellm_executed_batch_runner(llm_router: Router, proxy_logging_obj: ProxyLogging) -> LiteLLMExecutedBatchRunner: from litellm.proxy.proxy_server import general_settings, prisma_client @@ -652,8 +658,9 @@ async def retrieve_batch( # The DB may store raw provider file IDs (before hooks translate them). # Register any missing managed-file rows and return unified IDs. if unified_batch_id: + terminal_batch_response: Final = _require_batch_response(response) await ensure_batch_response_managed_file_ids( - response=response, + response=terminal_batch_response, managed_files_obj=managed_files_obj, prisma_client=prisma_client, verbose_proxy_logger=verbose_proxy_logger, @@ -799,8 +806,9 @@ async def retrieve_batch( # Fix: bug_feb14_batch_retrieve_returns_raw_input_file_id # Register any missing managed-file rows and return unified IDs. if unified_batch_id: + retrieved_batch_response: Final = _require_batch_response(response) await ensure_batch_response_managed_file_ids( - response=response, + response=retrieved_batch_response, managed_files_obj=managed_files_obj, prisma_client=prisma_client, verbose_proxy_logger=verbose_proxy_logger, diff --git a/litellm/proxy/client/cli/commands/pi.py b/litellm/proxy/client/cli/commands/pi.py index 87c3d3e8db3..bd02dc43a83 100644 --- a/litellm/proxy/client/cli/commands/pi.py +++ b/litellm/proxy/client/cli/commands/pi.py @@ -186,10 +186,16 @@ def provider_block( Real contextWindow/maxTokens matter: pi otherwise assumes 128k/16384, which breaks compaction thresholds and over-asks models with smaller output caps. + pi sniffs compat from the base URL, and one gateway URL fronts models with + different capabilities, so both flags are pinned off. """ return { "baseUrl": base_url.rstrip("/") + "/v1", "api": "openai-completions", + "compat": { + "supportsStore": False, + "supportsLongCacheRetention": False, + }, "apiKey": f"${LITELLM_PROXY_API_KEY_ENV}", "models": [_model_entry(model_id, limits) for model_id in model_ids], } diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 8e8779ec193..e916f93ca36 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -7,6 +7,7 @@ from typing import Annotated, Any, Final, Literal, Union, get_args, get_origin import orjson from fastapi import Request, UploadFile, status from starlette._utils import get_route_path +from starlette.requests import HTTPConnection from typing_extensions import NotRequired, ReadOnly, Required, assert_never from litellm._logging import verbose_proxy_logger @@ -182,8 +183,10 @@ def _mark_body_received(byte_count: int | None) -> None: ) -def is_otlp_trace_request(request: Request) -> bool: - return request.method == "POST" and get_route_path(request.scope) in {"/v1/traces", "/v1/logs"} +def is_otlp_trace_request(request: HTTPConnection) -> bool: + if request.scope.get("type") != "http": + return False + return request.scope.get("method") == "POST" and get_route_path(request.scope) in {"/v1/traces", "/v1/logs"} async def read_request_body(request: Request | None) -> dict: diff --git a/litellm/proxy/db/baseline_accounting.py b/litellm/proxy/db/baseline_accounting.py index 486717f9b93..f7a2be654c8 100644 --- a/litellm/proxy/db/baseline_accounting.py +++ b/litellm/proxy/db/baseline_accounting.py @@ -243,18 +243,19 @@ WHERE scope = $1 AND revision > $2::bigint AND ($5::float8 IS NULL OR publication::jsonb->>'status' = 'estimated') ORDER BY started_at, request_id """ +_PUBLISHED_LOG_FIELDS: Final = ("autorouter_savings_estimate", "autorouter_savings") _UPDATE_LOGS: Final = """ WITH changes AS ( SELECT request_id, publication::jsonb AS publication FROM jsonb_to_recordset($1::jsonb) AS x(request_id text, publication jsonb) ) UPDATE "LiteLLM_SpendLogs" AS logs -SET metadata = (COALESCE(logs.metadata::jsonb, '{}'::jsonb) - 'autorouter_baseline_observation') || jsonb_build_object( +SET metadata = (COALESCE(logs.metadata::jsonb, '{}'::jsonb) - 'autorouter_baseline_observation') || (jsonb_build_object( 'autorouter_savings_estimate', changes.publication, 'autorouter_savings', CASE WHEN changes.publication->>'status' = 'estimated' THEN (changes.publication->>'baseline_spend')::float8 - (changes.publication->>'actual_spend')::float8 ELSE NULL END -) +) - ARRAY(SELECT jsonb_array_elements_text($2::jsonb))) FROM changes WHERE logs.request_id = changes.request_id """ _UPDATE_PUBLICATIONS: Final = """ @@ -397,7 +398,11 @@ async def _publish(db: SupportsRawQueries, changes: Sequence[_Change]) -> None: if not changes: return serialized: Final = json.dumps(tuple(change.model_dump(mode="json") for change in changes), separators=(",", ":")) - await db.execute_raw(_UPDATE_LOGS, serialized) + from litellm.proxy.spend_tracking.spend_tracking_utils import configured_spend_logs_metadata_fields + + fields: Final = configured_spend_logs_metadata_fields() + unstored: Final = tuple(name for name in _PUBLISHED_LOG_FIELDS if fields is not None and not fields.keeps(name)) + await db.execute_raw(_UPDATE_LOGS, serialized, json.dumps(unstored)) await db.execute_raw(_UPDATE_SESSIONS, serialized) if any(change.user_id for change in changes): await db.execute_raw(_UPDATE_USER_SESSIONS, serialized) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 370644675ce..94df9c1928c 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -654,7 +654,14 @@ class DBSpendUpdateWriter: from litellm.repositories.table_repositories import SpendLogsRepository request_id: Final = payload["request_id"] - row: Final = _batch_cost_row_to_write(payload, disable_spend_logs) + from litellm.proxy.spend_tracking.spend_tracking_utils import ( + configured_spend_logs_metadata_fields, + spend_log_row_with_retained_metadata, + ) + + row: Final = spend_log_row_with_retained_metadata( + _batch_cost_row_to_write(payload, disable_spend_logs), configured_spend_logs_metadata_fields() + ) spend_logs: Final = SpendLogsRepository(prisma_client).table try: claimed: Final = await spend_logs.create_many( @@ -2589,9 +2596,11 @@ class DBSpendUpdateWriter: PrismaDBExceptionHandler, ) - is_retryable = isinstance( - e, DB_RETRY_SAFE_ERROR_TYPES - ) or PrismaDBExceptionHandler.is_deadlock_error(e) + is_retryable = ( + isinstance(e, DB_RETRY_SAFE_ERROR_TYPES) + or PrismaDBExceptionHandler.is_deadlock_error(e) + or PrismaDBExceptionHandler.is_lock_timeout_error(e) + ) if not is_retryable: raise if i >= n_retry_times: diff --git a/litellm/proxy/google_endpoints/endpoints.py b/litellm/proxy/google_endpoints/endpoints.py index 87d2ee07818..8638bc57615 100644 --- a/litellm/proxy/google_endpoints/endpoints.py +++ b/litellm/proxy/google_endpoints/endpoints.py @@ -191,26 +191,12 @@ async def google_count_tokens(request: Request, model_name: str): request=token_request, call_endpoint=True, ) - if token_response is not None: - # cast the response to the well known format - original_response: Final[dict] = token_response.original_response or {} - if original_response: - return TokenCountDetailsResponse( - totalTokens=original_response.get("totalTokens", 0), - promptTokensDetails=original_response.get("promptTokensDetails", []), - ) - else: - return TokenCountDetailsResponse( - totalTokens=token_response.total_tokens or 0, - promptTokensDetails=[], - ) - - ######################################################### - # Return the response in the well known format - ######################################################### + if token_response is None: + return TokenCountDetailsResponse(totalTokens=0, promptTokensDetails=[]) + original_response: Final[dict] = token_response.original_response or {} return TokenCountDetailsResponse( - totalTokens=0, - promptTokensDetails=[], + totalTokens=original_response.get("totalTokens") or token_response.total_tokens or 0, + promptTokensDetails=original_response.get("promptTokensDetails", []), ) diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 8e1209401ab..6072d8bf22a 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -69,6 +69,7 @@ from litellm.types.guardrails import ( PresidioPresidioConfigModelUserInterface, SupportedGuardrailIntegrations, ToolPermissionGuardrailConfigModel, + with_tolerated_stream_scope, ) from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.proxy.guardrails.guardrail_hooks.hide_secrets import ( @@ -146,7 +147,7 @@ def _get_guardrails_list_response( GuardrailInfoResponse( guardrail_id=guardrail.get("guardrail_id"), guardrail_name=guardrail.get("guardrail_name"), - litellm_params=masked_params, + litellm_params=with_tolerated_stream_scope(masked_params), guardrail_info=guardrail.get("guardrail_info"), ) ) @@ -289,7 +290,7 @@ async def list_guardrails_v2( ) masked_litellm_params = ( parse_tolerant_litellm_params( - masked_litellm_params_dict, + with_tolerated_stream_scope(masked_litellm_params_dict), guardrail.get("guardrail_name") or "Unknown", params_model=BaseLitellmParams, ) @@ -336,7 +337,7 @@ async def list_guardrails_v2( ) masked_in_memory_litellm_params_typed = ( parse_tolerant_litellm_params( - masked_in_memory_litellm_params, + with_tolerated_stream_scope(masked_in_memory_litellm_params), guardrail.get("guardrail_name") or "Unknown", params_model=BaseLitellmParams, ) @@ -1263,7 +1264,7 @@ async def patch_guardrail( # Update litellm_params if default_on is provided or pii_entities_config is provided existing_litellm_params: Final = _as_str_object_mapping(dict(existing_guardrail.get("litellm_params", {}))) current_litellm_params: Final = parse_tolerant_litellm_params( - existing_litellm_params, + with_tolerated_stream_scope(existing_litellm_params), existing_guardrail.get("guardrail_name") or "Unknown", ) requested_litellm_params: Final[Mapping[str, object]] = ( @@ -1275,7 +1276,7 @@ async def patch_guardrail( MappingProxyType({**current_litellm_params.model_dump(exclude_unset=True), **requested_litellm_params}) ) try: - parsed_litellm_params: Final = LitellmParams(**merged_litellm_params) + parsed_litellm_params: Final = LitellmParams(**with_tolerated_stream_scope(merged_litellm_params)) except ValidationError as validation_error: raise HTTPException( status_code=422, @@ -1436,7 +1437,7 @@ async def get_guardrail_info(guardrail_id: str): ) masked_litellm_params = ( parse_tolerant_litellm_params( - masked_litellm_params_dict, + with_tolerated_stream_scope(masked_litellm_params_dict), result.get("guardrail_name") or "Unknown", params_model=BaseLitellmParams, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py index 06c8d6f390b..22d4406fa4b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py +++ b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py @@ -243,8 +243,12 @@ def first_value(request_data: Mapping[str, object], key: str) -> object: INPUT_HOOKS: Final = MappingProxyType( { - "request": frozenset((GuardrailEventHooks.pre_call, GuardrailEventHooks.pre_mcp_call)), - "response": frozenset((GuardrailEventHooks.post_call, GuardrailEventHooks.post_mcp_call)), + "request": frozenset( + (GuardrailEventHooks.pre_call, GuardrailEventHooks.pre_mcp_call, GuardrailEventHooks.logging_only) + ), + "response": frozenset( + (GuardrailEventHooks.post_call, GuardrailEventHooks.post_mcp_call, GuardrailEventHooks.logging_only) + ), } ) @@ -265,6 +269,7 @@ class AktoGuardrail(CustomGuardrail): GuardrailEventHooks.post_call, GuardrailEventHooks.pre_mcp_call, GuardrailEventHooks.post_mcp_call, + GuardrailEventHooks.logging_only, ] def __init__( diff --git a/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py b/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py index 3036e1eba8e..54d09ff0b1e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py +++ b/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py @@ -186,7 +186,10 @@ _UNSENDABLE: Final[_Classified] = (None, True) def request_attachments(request_data: Mapping[str, object]) -> RequestAttachments: # Both, so a decoy "messages" can't hide attachments in a Responses API "input" - containers: Final = (_parse(_ITEMS_ADAPTER, request_data.get(key)) or () for key in ("messages", "input")) + messages: Final = request_data.get("messages") + responses_input: Final = request_data.get("input") + sources: Final = (messages,) if responses_input is messages else (messages, responses_input) + containers: Final = (_parse(_ITEMS_ADAPTER, source) or () for source in sources) blocks: Final = tuple(chain.from_iterable(_message_blocks(message) for message in chain.from_iterable(containers))) classified: Final = tuple( chain.from_iterable(_block_attachments(block, index) for index, block in enumerate(blocks)) diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 1bfe0e5ed77..edbf1bb6dd2 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -54,6 +54,7 @@ from litellm.types.guardrails import ( LakeraCategoryThresholds, LitellmParams, SupportedGuardrailIntegrations, + with_tolerated_stream_scope, ) from .guardrail_hooks.llm_as_a_judge import ( @@ -634,6 +635,8 @@ def _configure_callback_scoping( "skip_tool_message_in_guardrail are enabled together, which excludes every message from " "scanning, so no request content would ever be scanned. Remove one of the two." ) + if isinstance(custom_guardrail_callback, CustomGuardrail): # pyright: ignore[reportUnnecessaryIsInstance] # module-path classes may only subclass CustomLogger + custom_guardrail_callback.apply_stream_scope(litellm_params.stream_scope) _apply_configured_bool_overrides(custom_guardrail_callback, litellm_params) @@ -718,9 +721,11 @@ class InMemoryGuardrailHandler: if isinstance(litellm_params_data, dict): if reject_invalid_logging_only_scope: - litellm_params = LitellmParams(**litellm_params_data) + litellm_params = LitellmParams(**with_tolerated_stream_scope(litellm_params_data)) else: - litellm_params = parse_tolerant_litellm_params(litellm_params_data, guardrail["guardrail_name"]) + litellm_params = parse_tolerant_litellm_params( + with_tolerated_stream_scope(litellm_params_data), guardrail["guardrail_name"] + ) else: litellm_params = litellm_params_data @@ -863,14 +868,17 @@ class InMemoryGuardrailHandler: # Extract additional params from litellm_params to pass to custom guardrail # This matches the behavior of other guardrail initializers (e.g., initialize_lakera) # and aligns with the documented behavior for custom guardrails - if hasattr(litellm_params, "model_dump"): - extra_params = litellm_params.model_dump(exclude_none=True) - else: - extra_params = dict(litellm_params) if litellm_params else {} - - # Remove params that are handled explicitly or are internal - for key in ["guardrail", "mode", "default_on"]: - extra_params.pop(key, None) + excluded_extra_param_keys: Final = frozenset(("guardrail", "mode", "default_on", "stream_scope")) + extra_params_items: Final = ( + litellm_params.model_dump(exclude_none=True).items() + if hasattr(litellm_params, "model_dump") + else iter(litellm_params) + if litellm_params + else () + ) + extra_params: Final = MappingProxyType( + {key: value for key, value in extra_params_items if key not in excluded_extra_param_keys} + ) _guardrail_callback: Final = _guardrail_class( guardrail_name=guardrail["guardrail_name"], @@ -1007,7 +1015,7 @@ class InMemoryGuardrailHandler: return params.model_dump() if isinstance(params, dict): try: - return parse_tolerant_litellm_params(params, guardrail_name).model_dump() + return parse_tolerant_litellm_params(with_tolerated_stream_scope(params), guardrail_name).model_dump() except ValidationError as e: verbose_proxy_logger.warning( "Could not normalize guardrail litellm_params for comparison; treating the guardrail as changed. Error: %s", diff --git a/litellm/proxy/lens/feedback_endpoints.py b/litellm/proxy/lens/feedback_endpoints.py index f9479873a84..59dac3f35ec 100644 --- a/litellm/proxy/lens/feedback_endpoints.py +++ b/litellm/proxy/lens/feedback_endpoints.py @@ -56,6 +56,8 @@ def write_scope(auth: UserAPIKeyAuth) -> Scope: return Scope(all_teams=True) if auth.user_role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY: raise HTTPException(403, "Admin viewers cannot write feedback") + if not auth.team_id and not auth.token: + raise HTTPException(403, "Feedback requires a team or API key") return Scope(team_id=auth.team_id or "", api_key_hash="" if auth.team_id else auth.token or "") diff --git a/litellm/proxy/lens/feedback_repository.py b/litellm/proxy/lens/feedback_repository.py index 9ee1af164e2..3228cd2792a 100644 --- a/litellm/proxy/lens/feedback_repository.py +++ b/litellm/proxy/lens/feedback_repository.py @@ -1,6 +1,7 @@ import hashlib from collections.abc import Mapping from datetime import datetime, timezone +from itertools import chain from typing import Final, Protocol from litellm.proxy.lens.feedback_models import Feedback, FeedbackInput, TraceFeedback, TraceFeedbackSummary @@ -165,7 +166,7 @@ class ClickHouseFeedbackStore: **access_parameters(scope).model_dump(), trace_ids=sorted({t.trace_id for t in traces}) ), ) - return tuple(summary for trace in traces for summary in _summaries(trace, rows)) + return tuple(chain.from_iterable(_summaries(trace, rows) for trace in traces)) def _summaries(trace: TraceIdentity, rows: tuple[FeedbackSummaryRow, ...]) -> tuple[TraceFeedbackSummary, ...]: diff --git a/litellm/proxy/lens/inference.py b/litellm/proxy/lens/inference.py index c31326908ce..65b2e2849f8 100644 --- a/litellm/proxy/lens/inference.py +++ b/litellm/proxy/lens/inference.py @@ -325,7 +325,7 @@ def renew_reservation(lens: Lens, reservation_id: str, now: datetime) -> Lens: async def wait_for_reservation( repo: LensRepository, lens_id: str, reservation_id: str, reserve: Callable[[Lens], Lens] ) -> None: - while (reserved := await repo.update(lens_id, reserve)) is not None: + while (reserved := await repo.update_locked(lens_id, reserve)) is not None: if any(held.id == reservation_id for held in reserved.reservations): return await asyncio.sleep(0.25) @@ -341,7 +341,7 @@ async def renew_budget_reservation( try: async with timeout(BUDGET_RENEW_INTERVAL): if ( - await repo.update( + await repo.update_locked( lens_id, lambda e: renew_reservation(e, reservation_id, datetime.now(timezone.utc)) ) is None diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 8efa878a08d..450e0368c5e 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -32,6 +32,7 @@ from litellm.constants import ( SESSION_ID_OMITTED_METADATA_KEY, X_LITELLM_DISABLE_CALLBACKS, ) +from litellm.integrations.custom_guardrail import without_server_streaming_classification from litellm.litellm_core_utils.core_helpers import is_codex_user_agent from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( @@ -208,6 +209,10 @@ def _sanitize_for_log(value: object) -> str: return text.replace("\r", "").replace("\n", "") +def sanitize_for_log(value: object) -> str: + return _sanitize_for_log(value) + + from litellm.router import Router from litellm.secret_managers.main import get_secret_bool from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS @@ -285,6 +290,9 @@ LITELLM_TRACE_CONTROL_METADATA_FIELDS: Final = frozenset( _UNTRUSTED_ROOT_CONTROL_FIELDS: Final = ( "weights", "_router_weights", + "fallback_depth", + "_target_order", + "attempted_targets", "proxy_server_request", "standard_logging_object", "secret_fields", @@ -2094,7 +2102,9 @@ def refresh_proxy_server_request_body_snapshot( | _TRANSPORT_ONLY_CREDENTIAL_KEYS | _CALLBACK_CREDENTIAL_KEYS ) - body: Final = {k: v for k, v in data.items() if k not in _body_snapshot_exclude} + body: Final = { + k: v for k, v in without_server_streaming_classification(data).items() if k not in _body_snapshot_exclude + } proxy_server_request["body"] = body if guardrails_applied and isinstance(logging_obj, Logging): metadata: Final = data.get(get_metadata_variable_name_from_kwargs(data)) diff --git a/litellm/proxy/management_endpoints/fallback_management_endpoints.py b/litellm/proxy/management_endpoints/fallback_management_endpoints.py index 543538f5edd..460b282a9f6 100644 --- a/litellm/proxy/management_endpoints/fallback_management_endpoints.py +++ b/litellm/proxy/management_endpoints/fallback_management_endpoints.py @@ -11,15 +11,21 @@ DELETE /fallback/{model} - Delete fallbacks for a specific model # pyright: reportMissingImports=false import json -from typing import TYPE_CHECKING, Final, Literal +from collections.abc import Mapping +from typing import TYPE_CHECKING, Annotated, Final, Literal + +from pydantic import Field, TypeAdapter from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.model_checks import get_all_fallbacks from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.utils import PrismaClient, evict_config_param, invalidate_config_param if TYPE_CHECKING: from fastapi import APIRouter, Depends, HTTPException, status + + from litellm.proxy.proxy_server import ProxyConfig else: try: from fastapi import APIRouter, Depends, HTTPException, status @@ -37,6 +43,33 @@ from litellm.types.management_endpoints.router_settings_endpoints import ( router: Final = APIRouter() +ROUTER_SETTINGS_PARAM: Final = "router_settings" +FallbackRule = dict[str, list[str]] +StoredFallback = Annotated[FallbackRule | dict[str, object] | str, Field(union_mode="left_to_right")] +_STORED_FALLBACKS: Final = TypeAdapter(list[StoredFallback]) + + +def _rule_covers(entry: StoredFallback, model: str) -> bool: + return isinstance(entry, dict) and model in entry + + +async def _router_settings_fresh_from_db(proxy_config: "ProxyConfig") -> dict[str, object]: + await evict_config_param(ROUTER_SETTINGS_PARAM) + config: Final = await proxy_config.get_config() + return config.get(ROUTER_SETTINGS_PARAM, {}) + + +async def _persist_router_settings(prisma_client: PrismaClient, router_settings: Mapping[str, object]) -> None: + router_settings_json: Final = json.dumps(router_settings) + await ConfigRepository(prisma_client).table.upsert( + where={"param_name": ROUTER_SETTINGS_PARAM}, + data={ + "create": {"param_name": ROUTER_SETTINGS_PARAM, "param_value": router_settings_json}, + "update": {"param_value": router_settings_json}, + }, + ) + await invalidate_config_param(ROUTER_SETTINGS_PARAM) + @router.post( "/fallback", @@ -85,24 +118,24 @@ async def create_fallback( ) # Validate that the model exists in the router - model_names: Final = llm_router.model_names - if data.model not in model_names: + known_model_names: Final = frozenset(llm_router.model_names) | llm_router.team_public_model_names + if data.model not in known_model_names: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail={ "error": f"Model '{data.model}' not found in router", - "available_models": list(model_names), + "available_models": sorted(known_model_names), }, ) # Validate that all fallback models exist in the router - invalid_fallback_models: Final = [m for m in data.fallback_models if m not in model_names] + invalid_fallback_models: Final = [m for m in data.fallback_models if m not in known_model_names] if invalid_fallback_models: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail={ "error": f"Invalid fallback models: {invalid_fallback_models}", - "available_models": list(model_names), + "available_models": sorted(known_model_names), }, ) @@ -122,9 +155,7 @@ async def create_fallback( }, ) - # Load existing config - config: Final = await proxy_config.get_config() - router_settings: Final = config.get("router_settings", {}) + router_settings: Final = await _router_settings_fresh_from_db(proxy_config) # Get the appropriate fallback list based on type fallback_key = "fallbacks" @@ -134,12 +165,12 @@ async def create_fallback( fallback_key = "content_policy_fallbacks" # Get existing fallbacks - existing_fallbacks: Final[list[dict[str, list[str]]]] = router_settings.get(fallback_key, []) + existing_fallbacks: Final = _STORED_FALLBACKS.validate_python(router_settings.get(fallback_key) or []) # Update or add the fallback configuration fallback_updated = False - for i, fallback_dict in enumerate(existing_fallbacks): - if data.model in fallback_dict: + for i, rule in enumerate(existing_fallbacks): + if _rule_covers(rule, data.model): # Update existing fallback existing_fallbacks[i] = {data.model: data.fallback_models} fallback_updated = True @@ -152,18 +183,7 @@ async def create_fallback( # Update router settings router_settings[fallback_key] = existing_fallbacks - # Save to database - convert router_settings to JSON string - router_settings_json: Final = json.dumps(router_settings) - await ConfigRepository(prisma_client).table.upsert( - where={"param_name": "router_settings"}, - data={ - "create": { - "param_name": "router_settings", - "param_value": router_settings_json, - }, - "update": {"param_value": router_settings_json}, - }, - ) + await _persist_router_settings(prisma_client, router_settings) # Update the in-memory router configuration setattr(llm_router, fallback_key, existing_fallbacks) @@ -291,9 +311,7 @@ async def delete_fallback( }, ) - # Load existing config - config: Final = await proxy_config.get_config() - router_settings: Final = config.get("router_settings", {}) + router_settings: Final = await _router_settings_fresh_from_db(proxy_config) # Get the appropriate fallback list based on type fallback_key = "fallbacks" @@ -303,14 +321,14 @@ async def delete_fallback( fallback_key = "content_policy_fallbacks" # Get existing fallbacks - existing_fallbacks: Final[list[dict[str, list[str]]]] = router_settings.get(fallback_key, []) + existing_fallbacks: Final = _STORED_FALLBACKS.validate_python(router_settings.get(fallback_key) or []) # Find and remove the fallback configuration fallback_found = False updated_fallbacks: Final = [] - for fallback_dict in existing_fallbacks: - if model not in fallback_dict: - updated_fallbacks.append(fallback_dict) + for rule in existing_fallbacks: + if not _rule_covers(rule, model): + updated_fallbacks.append(rule) else: fallback_found = True @@ -323,18 +341,7 @@ async def delete_fallback( # Update router settings router_settings[fallback_key] = updated_fallbacks - # Save to database - convert router_settings to JSON string - router_settings_json: Final = json.dumps(router_settings) - await ConfigRepository(prisma_client).table.upsert( - where={"param_name": "router_settings"}, - data={ - "create": { - "param_name": "router_settings", - "param_value": router_settings_json, - }, - "update": {"param_value": router_settings_json}, - }, - ) + await _persist_router_settings(prisma_client, router_settings) # Update the in-memory router configuration setattr(llm_router, fallback_key, updated_fallbacks) diff --git a/litellm/proxy/moyai_endpoints.py b/litellm/proxy/moyai_endpoints.py index a2335c0a49b..c4d8ad139f7 100644 --- a/litellm/proxy/moyai_endpoints.py +++ b/litellm/proxy/moyai_endpoints.py @@ -14,6 +14,8 @@ import json import os import secrets import time +from collections.abc import Mapping +from dataclasses import dataclass from typing import TYPE_CHECKING, Annotated, Final from urllib.parse import urlencode, urlparse @@ -63,6 +65,14 @@ class MoyaiConnectExchangeResponse(BaseModel): api_base: str +@dataclass(frozen=True, slots=True) +class _MoyaiConnectCode: + moyai_origin: str + nonce: str + exp: int + user_id: str | None + + def _b64url(data: bytes) -> str: return base64.urlsafe_b64encode(data).decode().rstrip("=") @@ -104,7 +114,7 @@ def _sign_connect_code(master_key: str, moyai_url: str, user_id: str | None) -> return f"{_b64url(payload)}.{_b64url(signature)}" -def _decode_connect_code(master_key: str, code: str) -> dict: +def _decode_connect_code(master_key: str, code: str) -> _MoyaiConnectCode: try: payload_b64, signature_b64 = code.split(".", 1) payload_raw: Final = _b64url_decode(payload_b64) @@ -122,7 +132,13 @@ def _decode_connect_code(master_key: str, code: str) -> dict: raise HTTPException(status_code=400, detail="Invalid Moyai connect code") if not isinstance(payload.get("moyai_origin"), str) or not isinstance(payload.get("nonce"), str): raise HTTPException(status_code=400, detail="Invalid Moyai connect code") - return payload + user_id: Final = payload.get("user_id") + return _MoyaiConnectCode( + moyai_origin=payload["moyai_origin"], + nonce=payload["nonce"], + exp=payload["exp"], + user_id=user_id if isinstance(user_id, str) else None, + ) @router.post( @@ -178,7 +194,7 @@ async def _claim_connect_nonce(prisma_client: "PrismaClient", nonce: str, exp: i raise HTTPException(status_code=400, detail="Invalid Moyai connect code") -async def _moyai_key_alias(prisma_client, moyai_url: str) -> str: +async def _moyai_key_alias(prisma_client: "PrismaClient", moyai_url: str) -> str: from litellm.repositories.verification_token_repository import ( VerificationTokenRepository, ) @@ -192,16 +208,14 @@ async def _moyai_key_alias(prisma_client, moyai_url: str) -> str: @with_service_target(CONFIG_PARAMS_TARGET) -async def _persist_moyai_url(prisma_client, moyai_url: str) -> None: +async def _persist_moyai_url(prisma_client: "PrismaClient", moyai_url: str) -> None: from litellm.proxy.proxy_server import user_api_key_cache - existing: dict = {} db_existing: Final = await _ui_settings_db(UISettingsRepository(prisma_client)).find_unique( where={"id": "ui_settings"} ) - if db_existing and db_existing.ui_settings: - raw: Final = db_existing.ui_settings - existing = json.loads(raw) if isinstance(raw, str) else dict(raw) + raw: Final = db_existing.ui_settings if db_existing else None + existing: Final[Mapping[str, object]] = (json.loads(raw) if isinstance(raw, str) else dict(raw)) if raw else {} ui_settings: Final = {**existing, "moyai_url": moyai_url} await _ui_settings_db(UISettingsRepository(prisma_client)).upsert( @@ -232,13 +246,13 @@ async def moyai_connect_exchange(request: Request, body: MoyaiConnectExchangeReq moyai_url: Final = normalize_moyai_url(body.moyai_url) except ValueError: raise HTTPException(status_code=400, detail="Invalid Moyai connect code") - if moyai_url is None or _origin(moyai_url) != payload["moyai_origin"]: + if moyai_url is None or _origin(moyai_url) != payload.moyai_origin: raise HTTPException(status_code=400, detail="Invalid Moyai connect code") if prisma_client is None: raise HTTPException(status_code=400, detail="Moyai quick connect needs a database connected to the proxy") - await _claim_connect_nonce(prisma_client, payload["nonce"], payload["exp"]) + await _claim_connect_nonce(prisma_client, payload.nonce, payload.exp) alias: Final = await _moyai_key_alias(prisma_client, moyai_url) key_response: Final = await generate_key_helper_fn( @@ -248,7 +262,7 @@ async def moyai_connect_exchange(request: Request, body: MoyaiConnectExchangeReq metadata={ "created_via": "moyai_quick_connect", "moyai_url": moyai_url, - "connected_by": payload.get("user_id"), + "connected_by": payload.user_id, }, table_name="key", llm_router=llm_router, diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index fb0ae09382e..ad498c4a4e5 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -1,4 +1,5 @@ import base64 +import logging import mimetypes import re from collections.abc import Mapping, Sequence @@ -18,8 +19,8 @@ from typing import ( from litellm.batches.batch_utils import batch_cost_is_final from litellm.constants import MAX_FILE_LIST_LIMIT from litellm.proxy._types import ProxyException +from litellm.repositories.managed_file_repository import ManagedFileRepository from litellm.repositories.table_repositories import ( - ManagedFileRepository, ManagedObjectRepository, ) from litellm.types.llms.openai import OpenAIFilesPurpose @@ -27,6 +28,7 @@ from litellm.types.utils import SpecialEnums if TYPE_CHECKING: from fastapi import Request + from opentelemetry.trace import Span from prisma.models import LiteLLM_ManagedObjectTable from litellm.proxy._types import UserAPIKeyAuth @@ -104,6 +106,31 @@ class ManagedFileIdResolver(Protocol): ) -> Mapping[str, str]: ... +@runtime_checkable +class ManagedBatchOutputFileWriter(Protocol): + def get_unified_output_file_id( + self, + output_file_id: str, + model_id: str, + model_name: str | None = None, + ) -> str: ... + + async def store_batch_output_file( + self, + *, + unified_file_id: str, + provider_file_id: str, + model_id: str | None, + model_name: str | None = None, + owner: "UserAPIKeyAuth", + litellm_parent_otel_span: "Span | None", + size_bytes: int | None = None, + fetch_provider_details: bool = True, + ) -> None: + """Register file metadata, optionally fetching provider details.""" + ... + + def is_base64_encoded_unified_file_id(b64_uid: object) -> str | Literal[False]: # Ensure b64_uid is a string and not a mock object if not isinstance(b64_uid, str): @@ -1282,19 +1309,21 @@ def apply_unified_file_ids(response: "LiteLLMBatch", unified_id_by_raw_id: Mappi async def ensure_batch_response_managed_file_ids( - response, - managed_files_obj, - prisma_client, - verbose_proxy_logger, - user_api_key_dict=None, + response: "LiteLLMBatch", + managed_files_obj: object | None, + prisma_client: "PrismaClient | None", + verbose_proxy_logger: logging.Logger, + user_api_key_dict: "UserAPIKeyAuth | None" = None, db_batch_object: object | None = None, unified_batch_id: str | Literal[False] | None = None, + *, + fetch_provider_details: bool = True, ) -> None: - """Normalize batch file IDs to managed unified IDs before DB persistence.""" + """Normalize batch file IDs and register output and error file metadata.""" await resolve_input_file_id_to_unified(response, prisma_client) await resolve_output_file_ids_to_unified(response, prisma_client) - if managed_files_obj is None: + if not isinstance(managed_files_obj, ManagedBatchOutputFileWriter): return model_id: Final = _model_id_for_batch_response(response, unified_batch_id) @@ -1308,8 +1337,10 @@ async def ensure_batch_response_managed_file_ids( if effective_auth is None: return - for file_attr in ("output_file_id", "error_file_id"): - raw_file_id = getattr(response, file_attr, None) + for file_attr, raw_file_id in ( + ("output_file_id", response.output_file_id), + ("error_file_id", response.error_file_id), + ): if not raw_file_id or is_base64_encoded_unified_file_id(raw_file_id): continue try: @@ -1318,12 +1349,15 @@ async def ensure_batch_response_managed_file_ids( model_id=model_id, model_name=model_name, ) - await managed_files_obj.store_unified_file_id( - file_id=new_unified_file_id, - file_object=None, - litellm_parent_otel_span=getattr(effective_auth, "parent_otel_span", None), - model_mappings={model_id: raw_file_id}, - user_api_key_dict=effective_auth, + await managed_files_obj.store_batch_output_file( + unified_file_id=new_unified_file_id, + provider_file_id=raw_file_id, + model_id=model_id, + model_name=model_name, + owner=effective_auth, + litellm_parent_otel_span=effective_auth.parent_otel_span, + size_bytes=None, + fetch_provider_details=fetch_provider_details, ) setattr(response, file_attr, new_unified_file_id) verbose_proxy_logger.debug("Converted batch %s %r to managed ID before DB write", file_attr, raw_file_id) diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 4e31f9713d0..bdc71908c7c 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -102,7 +102,7 @@ from litellm.proxy.openai_files_endpoints.general_upload_validation import ( raise_upload_validation_failure, ) from litellm.proxy.utils import PrismaClient, ProxyLogging, is_known_model -from litellm.repositories.table_repositories import ManagedFileRepository +from litellm.repositories.managed_file_repository import ManagedFileRepository from litellm.router import Router from litellm.types.llms.openai import ( CREATE_FILE_REQUESTS_PURPOSE, diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index edc9d96b298..865ed8df921 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -44,12 +44,14 @@ from litellm.constants import ( AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER, BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES, ) +from litellm.integrations.custom_guardrail import guardrail_request_data_with_streaming from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix from litellm.llms.anthropic.common_utils import AnthropicModelInfo, merge_anthropic_beta_headers from litellm.llms.azure.passthrough.transformation import ( foreign_azure_deployment, is_azure_body_model_inference_endpoint, ) +from litellm.llms.bedrock.passthrough.transformation import is_bedrock_streaming_endpoint from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.llms.deepgram.common_utils import ( deepgram_listen_callback_params, @@ -946,8 +948,6 @@ BEDROCK_ENDPOINT_ACTIONS: Final = { "count-tokens", } -BEDROCK_STREAMING_ACTIONS: Final = {"invoke-with-response-stream", "converse-stream"} - def is_bedrock_count_tokens_endpoint(endpoint: str) -> bool: return "count_tokens" in endpoint or "count-tokens" in endpoint @@ -1077,7 +1077,7 @@ async def handle_bedrock_passthrough_router_model( from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing # Detect streaming based on endpoint - is_streaming: Final = any(action in endpoint for action in BEDROCK_STREAMING_ACTIONS) + is_streaming: Final = is_bedrock_streaming_endpoint(endpoint) verbose_proxy_logger.debug( "Bedrock router passthrough: model='%s', endpoint='%s', streaming=%s", model, endpoint, is_streaming @@ -1085,15 +1085,19 @@ async def handle_bedrock_passthrough_router_model( # Use the common processing path (same as non-router models) # This ensures all metadata, hooks, and logging are properly initialized - data: Final[dict[str, object]] = {} + bedrock_payload: Final[dict[str, object]] = { + "model": model, + "method": request.method, + "endpoint": endpoint, + "data": request_body, + "custom_llm_provider": "bedrock", + } + data: Final[dict[str, object]] = guardrail_request_data_with_streaming( + MappingProxyType(bedrock_payload), + is_streaming=is_streaming, + ) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) - data["model"] = model - data["method"] = request.method - data["endpoint"] = endpoint - data["data"] = request_body - data["custom_llm_provider"] = "bedrock" - # Use the common passthrough processing to handle metadata and hooks # This also handles all response formatting (streaming/non-streaming) and exceptions try: @@ -1285,14 +1289,19 @@ async def bedrock_llm_proxy_route( "Bedrock passthrough: Using direct Bedrock model '%s' for endpoint '%s'", model, endpoint ) - data: Final[dict[str, object]] = {} + is_streaming: Final = is_bedrock_streaming_endpoint(endpoint) + passthrough_payload: Final[dict[str, object]] = { + "method": request.method, + "endpoint": endpoint, + "data": request_body, + "custom_llm_provider": "bedrock", + } + data: Final[dict[str, object]] = guardrail_request_data_with_streaming( + MappingProxyType(passthrough_payload), + is_streaming=is_streaming, + ) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) - data["method"] = request.method - data["endpoint"] = endpoint - data["data"] = request_body - data["custom_llm_provider"] = "bedrock" - try: result: Final = await base_llm_response_processor.base_passthrough_process_llm_request( request=request, diff --git a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py index 94a75a9802e..a5a04fb68d2 100644 --- a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py +++ b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py @@ -54,10 +54,8 @@ from litellm.llms.base_llm.managed_resources.isolation import ( from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.batches_endpoints.common_utils import validate_batch_list_limit from litellm.proxy.common_utils.sse_keepalive import split_complete_sse_frames -from litellm.repositories.table_repositories import ( - ManagedFileRepository, - ManagedObjectRepository, -) +from litellm.repositories.managed_file_repository import ManagedFileRepository +from litellm.repositories.table_repositories import ManagedObjectRepository from litellm.types.llms.openai import BATCH_GUARDRAIL_RESPONSE_FIELD, OpenAIFileObject from litellm.types.passthrough_endpoints.managed_id_rewriter import ( ManagedFileIdReader, diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e80d3ed189b..755a41aff1d 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -48,7 +48,11 @@ from litellm.constants import ( SESSION_ID_OMITTED_METADATA_KEY, WEBSOCKET_CLOSE_REASON_MAX_BYTES, ) -from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + guardrail_request_data_with_streaming, + without_server_streaming_classification, +) from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import ( bind_budget_reservation_to_callbacks, @@ -603,6 +607,12 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): from litellm.proxy.proxy_server import llm_router _parsed_body = _parsed_body or {} + parsed_body_typed: Final[Mapping[str, object]] = cast(Mapping[str, object], _parsed_body) # cast-ok: json + server_marker_free_body: Final = without_server_streaming_classification(parsed_body_typed) + # The marker-free body must propagate through the caller's request dict, so + # downstream guardrail scans and snapshots never observe the server streaming marker. + _parsed_body.clear() + _parsed_body.update(server_marker_free_body) managed_model: Final = get_model_from_request( request_data=_parsed_body, route=get_request_route(request), @@ -828,7 +838,7 @@ def _build_passthrough_failure_request_payload( error response. Spend tracking only attributes a recovered cost when it comes paired with a usage object, so both keys are written together. """ - request_payload: Final[dict] = dict(parsed_body or {}) + request_payload: Final[dict] = dict(cast(Mapping[str, object], parsed_body or {})) # cast-ok: json body if kwargs: request_payload.update(kwargs) if logging_obj is not None: @@ -1163,7 +1173,7 @@ async def pass_through_request( is_multipart: Final = HttpPassThroughEndpointHelpers.is_multipart(request) and not custom_body if custom_body: - _parsed_body = custom_body + _parsed_body = dict(custom_body) elif is_multipart: # Don't parse multipart body here - it will be handled by make_multipart_http_request _parsed_body = {} @@ -1232,6 +1242,14 @@ async def pass_through_request( if _parsed_body is None: _parsed_body = {} _parsed_body["litellm_logging_obj"] = logging_obj + is_streaming_pass_through: Final = bool( + HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body( + parsed_body=_parsed_body, + stream=stream, + ) + ) + typed_body: Final[Mapping[str, object]] = cast(Mapping[str, object], _parsed_body) # cast-ok: json + _parsed_body = guardrail_request_data_with_streaming(typed_body, is_streaming=is_streaming_pass_through) ### CALL HOOKS ### - modify incoming data / reject request before calling the model _parsed_body = await proxy_logging_obj.pre_call_hook( @@ -2518,10 +2536,10 @@ async def websocket_passthrough_request( ) ### CALL HOOKS ### - modify incoming data / reject request before calling the model - websocket_data: dict[str, object] = {} - websocket_data = await proxy_logging_obj.pre_call_hook( + websocket_hook_data: Final = guardrail_request_data_with_streaming(MappingProxyType({}), is_streaming=True) + await proxy_logging_obj.pre_call_hook( user_api_key_dict=user_api_key_dict, - data=websocket_data, + data=websocket_hook_data, call_type="pass_through_endpoint", ) diff --git a/litellm/proxy/policy_engine/pipeline_executor.py b/litellm/proxy/policy_engine/pipeline_executor.py index e0e17feb88c..45f7cddcbc1 100644 --- a/litellm/proxy/policy_engine/pipeline_executor.py +++ b/litellm/proxy/policy_engine/pipeline_executor.py @@ -8,7 +8,8 @@ pass/fail actions (allow, block, next, modify_response) and data forwarding. import copy import time from collections.abc import Callable, Mapping, Sequence -from typing import TYPE_CHECKING, Final, Literal, TypeVar +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal, TypeVar, cast from pydantic import BaseModel @@ -29,6 +30,7 @@ from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_g from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( UnifiedLLMGuardrails, ) +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.proxy.policy_engine.pipeline_types import ( PipelineExecutionResult, PipelineStep, @@ -265,6 +267,30 @@ class _LegacyHookStreamAdapter(CustomGuardrail): return recorder.inputs +_PIPELINE_EVENT_HOOKS: Final = MappingProxyType( + { + "pre_call": GuardrailEventHooks.pre_call, + "post_call": GuardrailEventHooks.post_call, + "during_call": GuardrailEventHooks.during_call, + } +) + + +def _pipeline_stream_scope_allows( + callback: CustomGuardrail, + hook_input: Mapping[str, object], + mode: str, + streaming_chunks: list[object] | None, +) -> bool: + event_type: Final = _PIPELINE_EVENT_HOOKS.get(mode) + if event_type is None: + return True + return callback.stream_scope_allows( + hook_input if streaming_chunks is None else {**hook_input, "stream": True}, + event_type, + ) + + def _prepare_hook_input( step: PipelineStep, callback: CustomGuardrail, @@ -494,7 +520,7 @@ class PipelineExecutor: streaming_chunks: list[object] | None = None, # mutable-ok: shared buffered-stream chunks, read per step endpoint_translation: "BaseTranslation | None" = None, ) -> tuple[ - Literal["pass", "fail", "error"], + Literal["pass", "fail", "error", "skip"], dict | None, str | None, Exception | None, @@ -504,7 +530,7 @@ class PipelineExecutor: Returns: Tuple of (outcome, modified_data, error_detail, original_exception): - - outcome: "pass", "fail", or "error" + - outcome: "pass", "fail", "error", or "skip" - modified_data: dict if guardrail returned modified data, else None - error_detail: error message string if fail/error, else None - original_exception: the exception the guardrail raised, so the @@ -516,6 +542,10 @@ class PipelineExecutor: verbose_proxy_logger.warning("Pipeline: guardrail '%s' not found in callbacks", step.guardrail) return ("error", None, f"Guardrail '{step.guardrail}' not found", None) + hook_data: Final[Mapping[str, object]] = cast(Mapping[str, object], data) # cast-ok: payload + if not _pipeline_stream_scope_allows(callback, hook_data, mode, streaming_chunks): + return ("skip", None, None, None) + hook_input, scans_raw_request = _prepare_hook_input(step, callback, data, raw_request_snapshot) snapshot_entries_before: Final = len(_recorded_guardrail_information(hook_input)) @@ -702,10 +732,13 @@ def _pipeline_action_for_outcome(step: PipelineStep, outcome: str) -> str: """ Map pipeline step outcome to the configured action. + - skip -> next (stream_scope mismatch; do not apply on_pass/on_fail) - pass -> on_pass - fail -> on_fail (content/policy intervention) - error -> on_error if set, else on_fail (backward compatible) """ + if outcome == "skip": + return "next" if outcome == "pass": return step.on_pass if outcome == "fail": diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 06db63350e2..8627565321c 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -114,7 +114,6 @@ from litellm.proxy._types import ( LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, LitellmUserRoles, - ModelAccessDeniedProxyException, PassThroughGenericEndpoint, ProxyErrorTypes, ProxyException, @@ -356,6 +355,7 @@ from litellm.proxy.auth.auth_object_prefetch import AUTH_OBJECTS_TARGET from litellm.proxy.auth.auth_utils import ( check_response_size_is_safe, is_request_body_safe, + log_model_access_denial, log_once_if_budget_reservation_disabled, warn_once_if_custom_auth_skips_common_checks, ) @@ -768,6 +768,7 @@ except ImportError: shutdown_billing_metrics_recorder = None from fastapi.exception_handlers import http_exception_handler from starlette.exceptions import HTTPException as StarletteHTTPException +from starlette.websockets import WebSocketState from litellm.proxy import tracing_endpoints from litellm.proxy.middleware.admission_control_middleware import ( @@ -886,7 +887,7 @@ from litellm.proxy.utils import ( # noqa: F401, RUF100 # legacy module exports hash_password, hash_token, invalidate_config_param, - is_projected_spend_over_limit, + is_projected_spend_over_limit, # pyright: ignore[reportUnusedImport] # backwards-compatible package export is_valid_team_configs, litellm_config_cache, migrate_passwords_to_scrypt_async, @@ -2057,7 +2058,7 @@ class UserAPIKeyCacheTTLEnum(enum.Enum): @app.exception_handler(ProxyException) async def openai_exception_handler(request: Request, exc: ProxyException): # NOTE: DO NOT MODIFY THIS, its crucial to map to Openai exceptions - _log_model_access_denial(exc) + log_model_access_denial(exc) headers: Final = exc.headers error_dict: Final = with_call_id( JSON_OBJECT.validate_python(exc.to_dict()), @@ -2076,18 +2077,32 @@ async def openai_exception_handler(request: Request, exc: ProxyException): @app.exception_handler(StarletteHTTPException) -async def otlp_http_exception_handler(request: Request, exc: StarletteHTTPException) -> Response: - response: Final = tracing_endpoints.otlp_error_response(request, exc.status_code, exc.headers) +async def otlp_http_exception_handler(connection: Request | WebSocket, exc: StarletteHTTPException) -> Response | None: + if isinstance(connection, WebSocket): + return await _websocket_http_exception_response(connection, exc) + response: Final = tracing_endpoints.otlp_error_response(connection, exc.status_code, exc.headers) if response is not None: - _close_dangling_otel_server_span(request, exc.status_code, exc=exc) + _close_dangling_otel_server_span(connection, exc.status_code, exc=exc) return response - return await http_exception_handler(request, exc) + return await http_exception_handler(connection, exc) -def _log_model_access_denial(exc: ProxyException) -> None: - if not isinstance(exc, ModelAccessDeniedProxyException): - return - verbose_proxy_logger.warning(exc.sanitized_internal_message()) +async def _websocket_http_exception_response(websocket: WebSocket, exc: StarletteHTTPException) -> Response | None: + state: Final = websocket.application_state + match state: + case WebSocketState.CONNECTING: + return JSONResponse({"detail": exc.detail}, status_code=exc.status_code, headers=exc.headers) + case WebSocketState.CONNECTED: + await websocket.close(code=_websocket_close_code(exc.status_code)) + return None + case WebSocketState.RESPONSE | WebSocketState.DISCONNECTED: + return None + case _: + assert_never(state) + + +def _websocket_close_code(status_code: int) -> int: + return status.WS_1011_INTERNAL_ERROR if status_code >= 500 else status.WS_1008_POLICY_VIOLATION def _close_dangling_otel_server_span(request: Request, status_code: int, exc: Exception | None = None) -> None: @@ -4150,21 +4165,18 @@ async def update_cache( new_spend: Final = existing_spend + response_cost ## CHECK IF USER PROJECTED SPEND > SOFT LIMIT - if ( - existing_spend_obj.soft_budget_cooldown is False - and existing_spend_obj.soft_budget is not None - and ( - is_projected_spend_over_limit( - current_spend=new_spend, - soft_budget_limit=existing_spend_obj.soft_budget, - ) - is True - ) - ): - projected_spend, projected_exceeded_date = get_projected_spend_over_limit( + projection: Final = ( + get_projected_spend_over_limit( current_spend=new_spend, soft_budget_limit=existing_spend_obj.soft_budget, + budget_duration=existing_spend_obj.budget_duration, + budget_reset_at=existing_spend_obj.budget_reset_at, ) + if existing_spend_obj.soft_budget_cooldown is False and existing_spend_obj.soft_budget is not None + else None + ) + if projection is not None: + projected_spend, projected_exceeded_date = projection soft_limit: Final = existing_spend_obj.soft_budget call_info: Final = CallInfo( token=existing_spend_obj.token or "", @@ -6938,6 +6950,13 @@ class ProxyConfig: ).user_api_key_cache_max_size ) + if "spend_logs_metadata_fields" in general_settings: + _ = ConfigGeneralSettings.model_validate( + MappingProxyType( + {"spend_logs_metadata_fields": typed_general_settings["spend_logs_metadata_fields"]} + ) + ) + ### PKCE MULTI-INSTANCE PREREQUISITE CHECK ### # PKCE verifiers are stored in redis_usage_cache when available so they can # be read back by any instance (not just the one that started the auth flow). @@ -13247,7 +13266,7 @@ async def realtime_websocket_endpoint( llm_router=llm_router, ) except ProxyException as e: - _log_model_access_denial(e) + log_model_access_denial(e) await _reject_realtime_session(websocket, user_api_key_dict, code=1008, reason=e.message[:120]) return await websocket.accept(**accept_kwargs) @@ -18523,6 +18542,7 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro "mcp_client_id_header": "String", "mcp_trusted_proxy_ranges": "List", "mcp_xff_num_trusted_hops": "Integer", + "mcp_prefer_client_id_metadata_document": "Boolean", "always_include_stream_usage": "Boolean", "forward_client_headers_to_llm_api": "Boolean", "mcp_required_fields": "List", diff --git a/litellm/proxy/public_endpoints/provider_create_fields.json b/litellm/proxy/public_endpoints/provider_create_fields.json index 2dd9f50bec8..965d8dae4d2 100644 --- a/litellm/proxy/public_endpoints/provider_create_fields.json +++ b/litellm/proxy/public_endpoints/provider_create_fields.json @@ -3189,6 +3189,34 @@ ], "default_model_placeholder": "soniox/stt-async-v5" }, + { + "provider": "StrandsDecider", + "provider_display_name": "Strands Decider", + "litellm_provider": "strands_decider", + "credential_fields": [ + { + "key": "api_base", + "label": "API Base", + "placeholder": "http://localhost:8000", + "tooltip": "URL of your self-hosted Strands Decider server", + "required": true, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "api_key", + "label": "API Key", + "placeholder": null, + "tooltip": null, + "required": false, + "field_type": "password", + "options": null, + "default_value": null + } + ], + "default_model_placeholder": "strands_decider/strands-decider-2B-hobson-v19" + }, { "provider": "Tencent", "provider_display_name": "Tencent", @@ -3319,6 +3347,34 @@ ], "default_model_placeholder": "gpt-3.5-turbo" }, + { + "provider": "TypeSafe", + "provider_display_name": "TypeSafe", + "litellm_provider": "typesafe", + "credential_fields": [ + { + "key": "api_base", + "label": "API Base", + "placeholder": "https://api.typesafe.ai", + "tooltip": null, + "required": false, + "field_type": "text", + "options": null, + "default_value": null + }, + { + "key": "api_key", + "label": "API Key", + "placeholder": null, + "tooltip": null, + "required": true, + "field_type": "password", + "options": null, + "default_value": null + } + ], + "default_model_placeholder": "typesafe/jev-latest" + }, { "provider": "V0", "provider_display_name": "V0", @@ -3527,7 +3583,7 @@ }, { "provider": "Voyage", - "provider_display_name": "Voyage AI", + "provider_display_name": "VoyageAI by MongoDB", "litellm_provider": "voyage", "credential_fields": [ { diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index 6b6ee7a9b67..1cccc6b2cfb 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -18,6 +18,7 @@ from litellm.litellm_core_utils.get_blog_posts import ( GetBlogPosts, get_blog_posts, ) +from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap from litellm.proxy._types import ( CommonProxyErrors, ) @@ -453,15 +454,16 @@ async def get_public_fuse_presets() -> FusePresetCatalog: "/public/litellm_model_cost_map", tags=["public", "model management"], ) -async def get_litellm_model_cost_map(): +async def get_litellm_model_cost_map(catalog_only: bool = False): """ Public endpoint to get the LiteLLM model cost map. Returns pricing information for all supported models. + With catalog_only=true, returns the catalog as loaded, without entries registered at runtime for proxy deployments. """ import litellm try: - _model_cost_map: Final = litellm.model_cost + _model_cost_map: Final = GetModelCostMap.loaded_model_cost_map() if catalog_only else litellm.model_cost return _model_cost_map except Exception as e: raise HTTPException( diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 09c483c47f2..31ea40e2c24 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -678,9 +678,8 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr "enable_tag_filtering", ] - # Merge override settings into data (only if not already set in request) for key in per_request_settings: - if key in override_settings and key not in data: + if override_settings.get(key) is not None and key not in data: data[key] = override_settings[key] # Use main router with overridden kwargs diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 9d52f7d71bb..99bd2fbc480 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -9,7 +9,7 @@ from functools import reduce from types import MappingProxyType from typing import TYPE_CHECKING, Final, Literal, Protocol, cast, runtime_checkable -from pydantic import BaseModel, JsonValue +from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError import litellm from litellm._logging import verbose_proxy_logger @@ -47,7 +47,12 @@ from litellm.litellm_core_utils.ptu_pricing import azure_spillover from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.llms.anthropic.common_utils import resolve_used_client_oauth_token -from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload, SpendLogsRouterMetadata +from litellm.proxy._types import ( + SpendLogsMetadata, + SpendLogsMetadataFields, + SpendLogsPayload, + SpendLogsRouterMetadata, +) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error from litellm.proxy.utils import PrismaClient, hash_token @@ -1726,6 +1731,35 @@ def should_store_prompts_and_responses_in_spend_logs() -> bool: return get_secret_bool("STORE_PROMPTS_IN_SPEND_LOGS") is True +_SPEND_LOGS_METADATA_FIELDS_ADAPTER: Final[TypeAdapter[SpendLogsMetadataFields | None]] = TypeAdapter( + SpendLogsMetadataFields | None +) +_SPEND_LOGS_METADATA_ADAPTER: Final = TypeAdapter(dict[str, JsonValue]) + + +def configured_spend_logs_metadata_fields() -> SpendLogsMetadataFields | None: + from litellm.proxy.proxy_server import general_settings_view + + try: + return _SPEND_LOGS_METADATA_FIELDS_ADAPTER.validate_python( + general_settings_view().get("spend_logs_metadata_fields") + ) + except ValidationError as e: + verbose_proxy_logger.error("Ignoring invalid general_settings.spend_logs_metadata_fields: %s", e) + return None + + +def spend_log_row_with_retained_metadata( + row: Mapping[str, object], fields: SpendLogsMetadataFields | None +) -> Mapping[str, object]: + metadata_json: Final = row.get("metadata") + if fields is None or not isinstance(metadata_json, str): + return row + metadata: Final = _SPEND_LOGS_METADATA_ADAPTER.validate_json(metadata_json) + retained: Final = {name: value for name, value in metadata.items() if fields.keeps(name)} + return {**row, "metadata": safe_dumps(retained)} + + def _get_status_for_spend_log( metadata: dict, ) -> Literal["success", "failure"]: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 614a9b19d8a..101e6ef553b 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -147,6 +147,7 @@ from litellm.litellm_core_utils.core_helpers import ( independent_snapshot, is_expected_client_error, ) +from litellm.litellm_core_utils.duration_parser import get_budget_window_start from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.safe_json_loads import safe_json_loads @@ -4391,9 +4392,13 @@ async def get_config_param(prisma_client: "PrismaClient", param_name: str) -> An return row -async def evict_config_param(param_name: str) -> None: - with service_target(CONFIG_PARAMS_TARGET): - await litellm_config_cache.async_delete_cache(_config_cache_key(param_name)) +async def evict_config_param(param_name: str, cache: DualCache | None = None) -> None: + target: Final = cache if cache is not None else litellm_config_cache + try: + with service_target(CONFIG_PARAMS_TARGET): + await target.async_delete_cache(_config_cache_key(param_name)) + except Exception as e: # noqa: BLE001 # best-effort eviction; config writes must never fail on redis errors + verbose_proxy_logger.warning("config cache eviction of %s failed: %s", param_name, e) async def invalidate_config_param(param_name: str) -> None: @@ -7441,6 +7446,13 @@ class ProxyUpdateSpend: "Spend tracking - processing %d spend logs for DB write", len(logs_to_process), ) + from litellm.proxy.spend_tracking.spend_tracking_utils import ( + configured_spend_logs_metadata_fields, + spend_log_row_with_retained_metadata, + ) + + retention: Final = configured_spend_logs_metadata_fields() + rows_to_write: Final = [spend_log_row_with_retained_metadata(row, retention) for row in logs_to_process] start_time: Final = time.time() try: for i in range(n_retry_times + 1): @@ -7450,7 +7462,7 @@ class ProxyUpdateSpend: if not base_url.endswith("/"): base_url += "/" verbose_proxy_logger.debug("base_url: %s", base_url) - json_data = json.dumps(logs_to_process) + json_data = json.dumps(rows_to_write) response = await db_writer_client.post( url=base_url + "spend/update", data=json_data, @@ -7461,8 +7473,8 @@ class ProxyUpdateSpend: # Items already removed from queue at start of function pass else: - for j in range(0, len(logs_to_process), BATCH_SIZE): - batch = logs_to_process[j : j + BATCH_SIZE] + for j in range(0, len(rows_to_write), BATCH_SIZE): + batch = rows_to_write[j : j + BATCH_SIZE] batch_with_dates = [prisma_client.jsonify_object({**entry}) for entry in batch] isolation_budget = MAX_SPEND_LOG_ISOLATION_FAILURES_PER_BATCH for statement_rows in spend_log_write_batches( @@ -8097,66 +8109,94 @@ def _get_month_end_date(today: date) -> date: return date(today.year, today.month + 1, 1) - timedelta(days=1) -def is_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None) -> bool: - if soft_budget_limit is None: - # If there's no limit, we can't exceed it. - return False +MIN_ELAPSED_WINDOW_FRACTION: Final = 1 / 24 + +def _as_aware(moment: datetime) -> datetime: + return moment if moment.tzinfo is not None else moment.replace(tzinfo=timezone.utc) + + +def _reset_window(budget_duration: str, budget_reset_at: datetime) -> tuple[datetime, datetime] | None: + window_end: Final = _as_aware(budget_reset_at) + window_start: Final = get_budget_window_start(budget_duration, window_end) + return (window_start, window_end) if window_start < window_end else None + + +def _project_within_window( + current_spend: float, + soft_budget_limit: float, + window: tuple[datetime, datetime], + now: datetime | None, +) -> tuple[float, date] | None: + window_start, window_end = window + moment: Final = (_as_aware(now) if now is not None else datetime.now(timezone.utc)).astimezone(window_end.tzinfo) + if moment < window_start: + return None + elapsed: Final = max(moment - window_start, (window_end - window_start) * MIN_ELAPSED_WINDOW_FRACTION) + remaining: Final = max(window_end - moment, timedelta(0)) + spend_per_second: Final = current_spend / elapsed.total_seconds() + projected_spend: Final = current_spend + spend_per_second * remaining.total_seconds() + if projected_spend <= soft_budget_limit: + return None + remaining_budget: Final = soft_budget_limit - current_spend + if spend_per_second <= 0 or remaining_budget <= 0: + return projected_spend, moment.date() + exceed_at: Final = min(moment + timedelta(seconds=remaining_budget / spend_per_second), window_end) + return projected_spend, exceed_at.date() + + +def _project_to_month_end(current_spend: float, soft_budget_limit: float) -> tuple[float, date] | None: today: Final = date.today() + remaining_days: Final = (_get_month_end_date(today) - today).days + daily_spend: Final = current_spend / max(today.day - 1, 1) + projected_spend: Final = current_spend + daily_spend * remaining_days + if projected_spend <= soft_budget_limit: + return None + remaining_budget: Final = soft_budget_limit - current_spend + if daily_spend <= 0 or remaining_budget <= 0: + return projected_spend, today + return projected_spend, today + timedelta(days=remaining_budget / daily_spend) - # Finding the first day of the next month, then subtracting one day to get the end of the current month. - end_month: Final = _get_month_end_date(today) - remaining_days: Final = (end_month - today).days - - # Check for the start of the month to avoid division by zero - if today.day == 1: - daily_spend_estimate = current_spend - else: - daily_spend_estimate = current_spend / (today.day - 1) - - # Total projected spend for the month - projected_spend: Final = current_spend + (daily_spend_estimate * remaining_days) - - if projected_spend > soft_budget_limit: - print_verbose("Projected spend exceeds soft budget limit!") - return True - return False +def is_projected_spend_over_limit( + current_spend: float, + soft_budget_limit: float | None, + budget_duration: str | None = None, + budget_reset_at: datetime | None = None, + now: datetime | None = None, +) -> bool: + return ( + get_projected_spend_over_limit( + current_spend=current_spend, + soft_budget_limit=soft_budget_limit, + budget_duration=budget_duration, + budget_reset_at=budget_reset_at, + now=now, + ) + is not None + ) _is_projected_spend_over_limit: Final = is_projected_spend_over_limit -def get_projected_spend_over_limit(current_spend: float, soft_budget_limit: float | None) -> tuple | None: +def get_projected_spend_over_limit( + current_spend: float, + soft_budget_limit: float | None, + budget_duration: str | None = None, + budget_reset_at: datetime | None = None, + now: datetime | None = None, +) -> tuple[float, date] | None: if soft_budget_limit is None: return None - - today: Final = date.today() - end_month: Final = _get_month_end_date(today) - remaining_days: Final = (end_month - today).days - - # assuming the current spend till today (not including today) - if today.day == 1: - daily_spend = current_spend - else: - daily_spend = current_spend / (today.day - 1) - projected_spend: Final = current_spend + (daily_spend * remaining_days) - - if projected_spend > soft_budget_limit: - if daily_spend <= 0: - limit_exceed_date = today - else: - remaining_budget: Final = soft_budget_limit - current_spend - if remaining_budget <= 0: - limit_exceed_date = today - else: - approx_days: Final = remaining_budget / daily_spend - limit_exceed_date = today + timedelta(days=approx_days) - - # return the projected spend and the date it will exceeded - return projected_spend, limit_exceed_date - - return None + window: Final = ( + _reset_window(budget_duration, budget_reset_at) + if budget_duration is not None and budget_reset_at is not None + else None + ) + if window is None: + return _project_to_month_end(current_spend, soft_budget_limit) + return _project_within_window(current_spend, soft_budget_limit, window, now) _get_projected_spend_over_limit: Final = get_projected_spend_over_limit diff --git a/litellm/repositories/__init__.py b/litellm/repositories/__init__.py index 7ffdcfa5ce6..f090911c416 100644 --- a/litellm/repositories/__init__.py +++ b/litellm/repositories/__init__.py @@ -6,6 +6,7 @@ from litellm.repositories.autorouter_session_repository import AutoRouterSession from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.credentials_repository import CredentialsRepository +from litellm.repositories.managed_file_repository import ManagedFileRepository from litellm.repositories.model_repository import ModelRepository from litellm.repositories.object_permission_repository import ( ObjectPermissionRepository, @@ -42,7 +43,6 @@ from litellm.repositories.table_repositories import ( HealthCheckRepository, InvitationLinkRepository, JWTKeyMappingRepository, - ManagedFileRepository, ManagedObjectRepository, ManagedVectorStoreIndexRepository, ManagedVectorStoresRepository, diff --git a/litellm/repositories/managed_file_repository.py b/litellm/repositories/managed_file_repository.py new file mode 100644 index 00000000000..e3cc4c1eac9 --- /dev/null +++ b/litellm/repositories/managed_file_repository.py @@ -0,0 +1,18 @@ +from typing import TYPE_CHECKING, Final + +from litellm.repositories.table_repositories import PrismaTableRepository +from litellm.types.llms.openai import OpenAIFileObject + +if TYPE_CHECKING: + from prisma import models as prisma_models # noqa: F401 # used by quoted base-class subscripts + + +class ManagedFileRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedFileTable"]): + table_name = "litellm_managedfiletable" + + async def update_file_object(self, unified_file_id: str, file_object: OpenAIFileObject) -> bool: + updated_rows: Final = await self.table.update_many( + where={"unified_file_id": unified_file_id}, + data={"file_object": file_object.model_dump_json()}, + ) + return updated_rows > 0 diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index 0f85818cba3..95b174f845f 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -131,10 +131,6 @@ class JWTKeyMappingRepository(PrismaTableRepository["prisma_models.LiteLLM_JWTKe table_name = "litellm_jwtkeymapping" -class ManagedFileRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedFileTable"]): - table_name = "litellm_managedfiletable" - - class MemoryRepository(PrismaTableRepository["prisma_models.LiteLLM_MemoryTable"]): table_name = "litellm_memorytable" diff --git a/litellm/responses/dispatch.py b/litellm/responses/dispatch.py index c728dd0ea13..ad642fb427c 100644 --- a/litellm/responses/dispatch.py +++ b/litellm/responses/dispatch.py @@ -56,14 +56,14 @@ def _public_request( model: Final = fields.get("model") if not isinstance(model, str): return None - return native_call(args, kwargs, fields) + return native_call(legacy, args, kwargs) def _context(request: NativeCall) -> RouteContext: return RouteContext( Route.RESPONSES, - provider=optional_str(request.bound.get("custom_llm_provider")), - model=str(request.bound["model"]), + provider=optional_str(request.resolved.get("custom_llm_provider")), + model=str(request.resolved["model"]), ) diff --git a/litellm/responses/litellm_completion_transformation/handler.py b/litellm/responses/litellm_completion_transformation/handler.py index 505b5b09433..477d3a64339 100644 --- a/litellm/responses/litellm_completion_transformation/handler.py +++ b/litellm/responses/litellm_completion_transformation/handler.py @@ -106,6 +106,7 @@ class LiteLLMCompletionTransformationHandler: litellm_completion_request = await LiteLLMCompletionResponsesConfig.async_responses_api_session_handler( previous_response_id=previous_response_id, litellm_completion_request=litellm_completion_request, + instructions=responses_api_request.get("instructions"), ) acompletion_args: Final = {} diff --git a/litellm/responses/litellm_completion_transformation/session_handler.py b/litellm/responses/litellm_completion_transformation/session_handler.py index e60a71c7494..c20eac80085 100644 --- a/litellm/responses/litellm_completion_transformation/session_handler.py +++ b/litellm/responses/litellm_completion_transformation/session_handler.py @@ -1,5 +1,6 @@ import asyncio import json +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, cast import litellm @@ -41,6 +42,11 @@ def _normalize_redacted_tool_call_arguments(message: Message) -> None: function_call.arguments = REDACTED_TOOL_CALL_ARGUMENTS_PLACEHOLDER +def _stored_instructions(proxy_server_request: Mapping[str, object] | None) -> str | None: + instructions: Final = None if proxy_server_request is None else proxy_server_request.get("instructions") + return instructions if isinstance(instructions, str) and instructions else None + + class ResponsesSessionHandler: @staticmethod async def get_chat_completion_message_history_for_previous_response_id( @@ -70,10 +76,15 @@ class ResponsesSessionHandler: | ChatCompletionResponseMessage | Message ] = [] - for spend_log in all_spend_logs: + proxy_server_requests: Final = [ + await ResponsesSessionHandler.get_proxy_server_request_from_spend_log(spend_log=spend_log) + for spend_log in all_spend_logs + ] + for spend_log, proxy_server_request_dict in zip(all_spend_logs, proxy_server_requests): chat_completion_message_history = ( - await ResponsesSessionHandler.extend_chat_completion_message_with_spend_log_payload( + ResponsesSessionHandler.extend_chat_completion_message_with_spend_log_payload( spend_log=spend_log, + proxy_server_request_dict=proxy_server_request_dict, chat_completion_message_history=chat_completion_message_history, ) ) @@ -85,11 +96,20 @@ class ResponsesSessionHandler: return ChatCompletionSession( messages=chat_completion_message_history, litellm_session_id=litellm_session_id, + instructions=next( + ( + instructions + for instructions in map(_stored_instructions, reversed(proxy_server_requests)) + if instructions + ), + None, + ), ) @staticmethod - async def extend_chat_completion_message_with_spend_log_payload( + def extend_chat_completion_message_with_spend_log_payload( spend_log: "SpendLogsPayload", + proxy_server_request_dict: Mapping[str, object] | None, chat_completion_message_history: list[ AllMessageValues | GenericChatCompletionMessage @@ -105,18 +125,15 @@ class ResponsesSessionHandler: LiteLLMCompletionResponsesConfig, ) - proxy_server_request_dict: Final = await ResponsesSessionHandler.get_proxy_server_request_from_spend_log( - spend_log=spend_log, - ) response_input_param: str | ResponseInputParam | None = None - _messages: str | ResponseInputParam | None = None ############################################################ # Add Input messages for this Spend Log ############################################################ if proxy_server_request_dict: - _response_input_param: Final = proxy_server_request_dict.get("input", None) - _messages = proxy_server_request_dict.get("messages", None) + _response_input_param: Final = proxy_server_request_dict.get("input") or proxy_server_request_dict.get( + "messages" + ) if isinstance(_response_input_param, (str, list)): response_input_param = _response_input_param elif isinstance(_response_input_param, dict): @@ -126,25 +143,13 @@ class ResponsesSessionHandler: ) if response_input_param: - chat_completion_messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( - input=response_input_param, - responses_api_request=proxy_server_request_dict or {}, - replay_reasoning=True, + chat_completion_message_history.extend( + LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( + input=response_input_param, + responses_api_request={}, + replay_reasoning=True, + ) ) - chat_completion_message_history.extend(chat_completion_messages) - - ############################################################ - # Check if `messages` field is present in the proxy server request dict - ############################################################ - elif _messages: - # ensure all messages are /chat/completions/messages - # certain requests can be stored as Responses API format - this ensures they are transformed to /chat/completions/messages - chat_completion_messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( - input=_messages, - responses_api_request=proxy_server_request_dict or {}, - replay_reasoning=True, - ) - chat_completion_message_history.extend(chat_completion_messages) ############################################################ # Add Output messages for this Spend Log @@ -162,7 +167,7 @@ class ResponsesSessionHandler: @staticmethod async def get_proxy_server_request_from_spend_log( spend_log: "SpendLogsPayload", - ) -> dict | None: + ) -> dict[str, object] | None: """ Get the parsed proxy server request from the spend log """ diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 5fa2b5fdabe..5f3ae5a7b2a 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -239,6 +239,7 @@ class ChatCompletionSession(TypedDict, total=False): | Message ] litellm_session_id: str | None + instructions: ReadOnly[str | None] ########### End of Initialize Classes used for Responses API ########### @@ -581,6 +582,7 @@ class LiteLLMCompletionResponsesConfig: async def async_responses_api_session_handler( previous_response_id: str, litellm_completion_request: dict, + instructions: str | None, ) -> dict: """ Async hook to get the chain of previous input and output pairs and return a list of Chat Completion messages @@ -589,18 +591,25 @@ class LiteLLMCompletionResponsesConfig: if previous_response_id: chat_completion_session = ( await ResponsesSessionHandler.get_chat_completion_message_history_for_previous_response_id( - previous_response_id=previous_response_id + previous_response_id=previous_response_id, ) ) _messages: Final = litellm_completion_request.get("messages") or [] session_messages: Final = chat_completion_session.get("messages") or [] + instructions_end: Final = 1 if instructions else 0 + carried_instructions: Final = None if instructions else chat_completion_session.get("instructions") + leading_system_messages: Final = ( + [LiteLLMCompletionResponsesConfig.transform_instructions_to_system_message(carried_instructions)] + if carried_instructions + else _messages[:instructions_end] + ) # If session messages are empty (e.g., no database in test environment), # we still need to process the new input messages # Store original _messages before combining for safety check original_new_messages: Final = _messages.copy() if _messages else [] - combined_messages = session_messages + _messages + combined_messages = leading_system_messages + session_messages + _messages[instructions_end:] # Fix: Ensure tool_results have corresponding tool_calls in previous assistant message # Pass tools parameter to help reconstruct tool_calls if not in cache diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 6610b19a8ee..3df029bb01f 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -18,7 +18,11 @@ from litellm.completion_extras.litellm_responses_transformation.transformation i LiteLLMResponsesTransformationHandler, ) from litellm.constants import DEFAULT_CHAT_COMPLETION_PARAM_VALUES, request_timeout -from litellm.integrations.anthropic_cache_control_hook import CARRY_UNMATCHED_MESSAGE_POINTS +from litellm.integrations.anthropic_cache_control_hook import ( + CARRY_UNMATCHED_MESSAGE_POINTS, + AnthropicCacheControlHook, + configured_injection_points, +) from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.core_helpers import normalize_drop_params from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -526,6 +530,13 @@ def _api_base_kwarg(kwargs: Mapping[str, object]) -> str | None: return api_base if isinstance(api_base, str) else None +def _dispatched_model_name(model: str, custom_llm_provider: str, api_base: str | None) -> str: + provider_model, _, _, _ = litellm.get_llm_provider( + model=model, custom_llm_provider=custom_llm_provider, api_base=api_base + ) + return _strip_responses_routing_prefix(provider_model) + + def _will_bridge_to_chat_completions( model: str, custom_llm_provider: str | None, @@ -536,20 +547,50 @@ def _will_bridge_to_chat_completions( """``_bridges_to_chat_completions`` for callers running before the provider config is resolved. Resolving the config is a pure lookup, so this asks the same question the dispatch - asks rather than restating its condition. Both callers resolve the provider before - this runs, so the only way to be wrong is a prompt manager that moves the model - across the bridge boundary, which would leave the deferred points to a pass that - never comes. + asks rather than restating its condition, with the model name the dispatch hands the + lookup: a provider whose config is keyed by model name (Bedrock Mantle reads the + price map) answers nothing for ``bedrock_mantle/openai.gpt-5.6-sol`` and would read + as bridged. Both callers resolve the provider before this runs, so the only way to be + wrong is a prompt manager that moves the model across the bridge boundary, which + would leave the deferred points to a pass that never comes. """ normalized_model: Final = _normalize_openai_chat_completions_responses_model(model) if custom_llm_provider is None: return True return _bridges_to_chat_completions( - _resolve_responses_api_provider_config(normalized_model[0], custom_llm_provider, model_info, api_base), + _resolve_responses_api_provider_config( + _dispatched_model_name(normalized_model[0], custom_llm_provider, api_base), + custom_llm_provider, + model_info, + api_base, + ), use_chat_completions_api or normalized_model[1], ) +def _stamp_injection_points_with_dialect( + kwargs: dict[str, object], + model: str, + custom_llm_provider: str | None, +) -> None: + """Carry the provider this layer resolved onto the points. + + The hook reads ``custom_llm_provider`` from the request kwargs, which never hold the one + resolved here, and resolving the model name alone reads a Foundry deployment of an OpenAI + model (``azure_ai/gpt-6-astra``) as Azure OpenAI, which left it on the Anthropic dialect. + """ + points: Final = configured_injection_points(kwargs.get("cache_control_injection_points")) + if not points: + return + kwargs["cache_control_injection_points"] = AnthropicCacheControlHook._stamped_with_dialect( + points, + model, + custom_llm_provider, + kwargs.get("api_base") or kwargs.get("base_url"), + kwargs.get("prompt_cache_options"), + ) + + @contextmanager def _prompt_management_sees_a_provisional_message_list( kwargs: dict[str, object], @@ -659,6 +700,7 @@ async def aresponses( _api_base_kwarg(kwargs), ), ): + _stamp_injection_points_with_dialect(kwargs, model, custom_llm_provider) ( model, merged_input, @@ -829,6 +871,7 @@ def _apply_prompt_management_to_responses_call( _api_base_kwarg(kwargs), ), ): + _stamp_injection_points_with_dialect(kwargs, model, custom_llm_provider) ( model, merged_input, diff --git a/litellm/router.py b/litellm/router.py index d03e085cf1e..c8640356f7e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -199,6 +199,7 @@ from litellm.router_utils.common_utils import ( resolve_model_group_alias, truncate_fallback_error_detail, warn_on_provider_credential_mismatch, + without_router_only_kwargs, ) from litellm.router_utils.cooldown_cache import CooldownCache from litellm.router_utils.cooldown_handlers import ( @@ -227,6 +228,7 @@ from litellm.router_utils.fallback_event_handlers import ( get_pre_routing_selection, has_unattempted_fallback_target, mid_stream_fallback_hop_kwargs, + mid_stream_fallback_snapshot_kwargs, mid_stream_retry_kwargs, per_request_fallback_controls, record_disable_fallbacks, @@ -498,19 +500,33 @@ _MODEL_INFO_ADAPTER: Final = TypeAdapter(Mapping[str, object]) _SILENT_MODEL_ADAPTER: Final = TypeAdapter(str | list[str]) _RESOLVED_RETRY_POLICY_ADAPTER: Final = TypeAdapter(RetryPolicy | None) _ROUTING_KWARGS_ADAPTER: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None) +_FALLBACK_HOP_ADAPTER: Final = TypeAdapter(Mapping[str, object]) _DEPLOYMENT_SELECTED_EVENT: Final = "litellm.request.deployment_selected" +def _is_fallback_hop(request_kwargs: Mapping[str, object]) -> bool: + fallback_depth: Final = request_kwargs.get("fallback_depth") + return isinstance(fallback_depth, int) and fallback_depth > 0 + + +def _deployment_that_just_failed(request_metadata: object) -> str | None: + try: + model_info: Final = _FALLBACK_HOP_ADAPTER.validate_python( + _FALLBACK_HOP_ADAPTER.validate_python(request_metadata).get("model_info") + ) + except ValidationError: + return None + model_id: Final = model_info.get("id") + return model_id if isinstance(model_id, str) else None + + def _deployment_pick_attributes(model: str, request_kwargs: Mapping[str, object] | None) -> Mapping[str, str | int]: """Bounded attributes for one deployment pick; attempt is 1-based within the current model group.""" kwargs: Final = request_kwargs or {} metadata: Final = kwargs.get("litellm_metadata", kwargs.get("metadata")) attempted_retries: Final = metadata.get("attempted_retries") if isinstance(metadata, Mapping) else None retries: Final = attempted_retries if isinstance(attempted_retries, int) else 0 - fallback_depth: Final = kwargs.get("fallback_depth") - reason: Final = ( - "retry" if retries > 0 else "fallback" if isinstance(fallback_depth, int) and fallback_depth > 0 else "initial" - ) + reason: Final = "retry" if retries > 0 else "fallback" if _is_fallback_hop(kwargs) else "initial" return MappingProxyType( { "litellm.deployment.attempt": retries + 1, @@ -2660,6 +2676,8 @@ class Router: kwargs["model"] = model kwargs["messages"] = messages kwargs["original_function"] = self._completion + controls: Final = per_request_fallback_controls(kwargs) + kwargs[MID_STREAM_FALLBACK_CONTROLS_KEY] = controls # rebind-ok: forwarded to every hop self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) response: Final = self.function_with_fallbacks(**kwargs) @@ -2671,10 +2689,10 @@ class Router: model_name = None deployment = None try: - # Capture kwargs before deployment selection so the streaming - # fallback iterator can re-dispatch with the original model group. - input_kwargs_for_streaming_fallback: Final = kwargs.copy() - input_kwargs_for_streaming_fallback["model"] = model + controls: Final = kwargs.pop(MID_STREAM_FALLBACK_CONTROLS_KEY, None) + input_kwargs_for_streaming_fallback: Final = mid_stream_fallback_snapshot_kwargs( + model=model, controls=controls, kwargs=kwargs + ) # pick the one that is available (lowest TPM/RPM) deployment = self.get_available_deployment( @@ -2728,13 +2746,15 @@ class Router: if model in self.model_names or not self.has_model_id(model): self.routing_strategy_pre_call_checks(deployment=deployment) - input_kwargs: Final = { - **litellm_params, - "messages": messages, - "caching": self.cache_responses, - "client": model_client, - **kwargs, - } + input_kwargs: Final = without_router_only_kwargs( + { + **litellm_params, + "messages": messages, + "caching": self.cache_responses, + "client": model_client, + **kwargs, + } + ) response: Final = litellm.completion(**input_kwargs) verbose_router_logger.info("litellm.completion(model=%s)\x1b[32m 200 OK\x1b[0m", model_name) @@ -2920,6 +2940,8 @@ class Router: messages=messages, kwargs=kwargs, ) + controls: Final = per_request_fallback_controls(kwargs) + kwargs[MID_STREAM_FALLBACK_CONTROLS_KEY] = controls # rebind-ok: forwarded to every hop if request_priority is not None and isinstance(request_priority, int): response = await self.schedule_acompletion(**kwargs) else: @@ -3815,8 +3837,10 @@ class Router: deployment = None _timeout_debug_deployment_dict = {} # this is a temporary dict to debug timeout issues try: - input_kwargs_for_streaming_fallback: Final = kwargs.copy() - input_kwargs_for_streaming_fallback["model"] = model + controls: Final = kwargs.pop(MID_STREAM_FALLBACK_CONTROLS_KEY, None) + input_kwargs_for_streaming_fallback: Final = mid_stream_fallback_snapshot_kwargs( + model=model, controls=controls, kwargs=kwargs + ) parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) start_time: Final = time.time() @@ -3880,15 +3904,15 @@ class Router: ) self.total_calls[model_name] += 1 - input_kwargs: Final = { - **litellm_params, - "messages": messages, - "caching": self.cache_responses, - "client": model_client, - **kwargs, - } - input_kwargs.pop("silent_model", None) - input_kwargs.pop("include_fallback_errors", None) + input_kwargs: Final = without_router_only_kwargs( + { + **litellm_params, + "messages": messages, + "caching": self.cache_responses, + "client": model_client, + **kwargs, + } + ) logging_obj: Final[LiteLLMLogging | None] = kwargs.get("litellm_logging_obj", None) @@ -4129,10 +4153,11 @@ class Router: function_name: str | None = None, ) -> None: """ - 3 jobs: + 4 jobs: - Adds selected deployment, model_info and api_base to kwargs["metadata"] (used for logging) - Adds default litellm params to kwargs, if set. - Merges tools from deployment with request (proxy-configured tools + request tools). + - On a fallback hop, drops the encrypted reasoning this deployment cannot decrypt, keeping its summary. """ for key in self._forwarded_alias_marker_keys_the_deployment_sets( deployment=deployment, forwarded_keys=kwargs.pop(_ALIAS_MARKER_FORWARDED_PARAMS_KWARG, ()) @@ -4155,6 +4180,9 @@ class Router: metadata_variable_name: Final = get_router_metadata_variable_name( function_name=function_name, ) + deployment_that_just_failed: Final = _deployment_that_just_failed( + _FALLBACK_HOP_ADAPTER.validate_python(kwargs).get(metadata_variable_name) + ) kwargs.setdefault(metadata_variable_name, {}).update( { @@ -4227,6 +4255,15 @@ class Router: kwargs["timeout"] = self._get_timeout(kwargs=kwargs, data=deployment["litellm_params"]) self._update_kwargs_with_default_litellm_params(kwargs=kwargs, metadata_variable_name=metadata_variable_name) + hop_kwargs: Final = _FALLBACK_HOP_ADAPTER.validate_python(kwargs) + if _is_fallback_hop(hop_kwargs): + EncryptedContentAffinityCheck.strip_reasoning_the_targets_cannot_decrypt( + self, + hop_kwargs.get("input"), + hop_kwargs.get("messages"), + (_FALLBACK_HOP_ADAPTER.validate_python(deployment),), + unmarked_origin=deployment_that_just_failed, + ) def _get_async_openai_model_client(self, deployment: dict, kwargs: dict): """ @@ -5522,6 +5559,7 @@ class Router: model=model, original_generic_function=original_generic_function, **kwargs ) carry_over_pre_routing_selection(live_kwargs=kwargs, snapshot=hop_kwargs) + carry_over_routed_deployment(live_kwargs=kwargs, snapshot=hop_kwargs) if kwargs.get("stream") and isinstance(response, BaseResponsesAPIStreamingIterator): return await self._aresponses_streaming_iterator(response=response, initial_kwargs=hop_kwargs) return response diff --git a/litellm/router_utils/common_utils.py b/litellm/router_utils/common_utils.py index 6e400738b8e..4d6af710bb9 100644 --- a/litellm/router_utils/common_utils.py +++ b/litellm/router_utils/common_utils.py @@ -2,7 +2,7 @@ import hashlib import json from collections.abc import Mapping, Sequence from types import MappingProxyType -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, TypeVar if TYPE_CHECKING: from litellm.types.llms.openai import OpenAIFileObject @@ -16,6 +16,14 @@ from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_stru from litellm.types.router import CredentialLiteLLMParams from litellm.types.utils import LlmProviders +_V = TypeVar("_V") + +ROUTER_ONLY_CALL_KWARGS: Final = frozenset({"silent_model", "include_fallback_errors"}) + + +def without_router_only_kwargs(kwargs: Mapping[str, _V]) -> dict[str, _V]: + return {key: value for key, value in kwargs.items() if key not in ROUTER_ONLY_CALL_KWARGS} + def is_proxy_admin_request(request_kwargs: Mapping[str, object] | None) -> bool: if request_kwargs is None: diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index ffa85c000c5..c7e292e09c9 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -335,6 +335,28 @@ def per_request_fallback_controls(kwargs: Mapping[str, object]) -> MidStreamFall ) +def mid_stream_fallback_snapshot_kwargs( + model: str, + controls: object, + kwargs: Mapping[str, object], +) -> dict[str, object]: + """ + The kwargs a completion attempt's stream re-enters the fallback chain with if it fails. + + async_function_with_retries popped the per-request fallback lists before the attempt ran, so + the carrier restores them here and rides along into every hop this re-entry opens. A shallow + copy keeps the metadata buckets shared with the live kwargs, the way the attempt's own + in-place bucket writes expect. + """ + hop_controls: Final = controls if isinstance(controls, MidStreamFallbackControls) else _NO_FALLBACK_CONTROLS + return { + **kwargs, + **hop_controls.overrides, + MID_STREAM_FALLBACK_CONTROLS_KEY: hop_controls, + "model": model, + } + + def mid_stream_fallback_hop_kwargs( model: str, original_generic_function: Callable[..., object], @@ -342,22 +364,18 @@ def mid_stream_fallback_hop_kwargs( kwargs: Mapping[str, object], ) -> dict[str, object]: """ - The kwargs one streaming attempt re-enters the fallback chain with if its stream fails. + The kwargs one generic-endpoint streaming attempt re-enters the fallback chain with if its stream fails. A shallow copy keeps ``attempted_targets`` shared with the outer chain, so entries this request already tried are never retried; the metadata buckets are copied key by key because the attempt writes deployment-specific fields into them in place. """ - hop_controls: Final = controls if isinstance(controls, MidStreamFallbackControls) else _NO_FALLBACK_CONTROLS copied_buckets: Final = MappingProxyType( {name: safe_deep_copy(kwargs[name]) for name in _ROUTER_METADATA_BUCKETS if isinstance(kwargs.get(name), dict)} ) return { - **kwargs, + **mid_stream_fallback_snapshot_kwargs(model=model, controls=controls, kwargs=kwargs), **copied_buckets, - **hop_controls.overrides, - MID_STREAM_FALLBACK_CONTROLS_KEY: hop_controls, - "model": model, "original_generic_function": original_generic_function, } diff --git a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py index 5d7241db379..5d47bf63a38 100644 --- a/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/encrypted_content_affinity_check.py @@ -40,6 +40,8 @@ from collections.abc import Iterator, Mapping, Sequence from functools import cache from typing import TYPE_CHECKING, Final, Optional, cast +from pydantic import TypeAdapter + from litellm._logging import verbose_router_logger from litellm.integrations.custom_logger import CustomLogger, Span from litellm.litellm_core_utils.credential_accessor import CredentialAccessor @@ -51,10 +53,13 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.llms.openai import AllMessageValues from litellm.types.router import Deployment +from litellm.utils import get_order_filtered_deployments if TYPE_CHECKING: from litellm.router import Router +_REQUEST_KWARGS_ADAPTER: Final = TypeAdapter(Mapping[str, object]) + class EncryptedContentAffinityCheck(CustomLogger): """ @@ -253,12 +258,22 @@ class EncryptedContentAffinityCheck(CustomLogger): ] return matches, originating - def _strip_reasoning_the_target_cannot_decrypt( - self, + @staticmethod + def strip_reasoning_the_targets_cannot_decrypt( + router: "Router | None", request_input: object, anthropic_messages: object, target_deployments: Sequence[Mapping[str, object]], + *, + unmarked_origin: str | None, ) -> None: + """ + Drop the encrypted reasoning that none of ``target_deployments`` minted or shares an + encryption boundary with, keeping each item's readable summary. Encrypted reasoning that + carries no litellm origin marker is attributed to ``unmarked_origin``: the affinity pin + names the deployment its marker decoded to, a fallback hop names the deployment that just + failed, and ``None`` drops it, since no deployment is known to have minted it. + """ target_ids: Final = frozenset( str(model_info["id"]) for target in target_deployments @@ -267,30 +282,34 @@ class EncryptedContentAffinityCheck(CustomLogger): target_boundaries: Final = frozenset( boundary for target in target_deployments - if (boundary := self._encryption_boundary_key(target.get("litellm_params"))) is not None + if (boundary := EncryptedContentAffinityCheck._encryption_boundary_key(target.get("litellm_params"))) + is not None ) @cache - def target_can_decrypt(origin_model_id: str) -> bool: + def target_can_decrypt(marked_origin: str | None) -> bool: + origin_model_id: Final = marked_origin if marked_origin is not None else unmarked_origin + if origin_model_id is None: + return False if origin_model_id in target_ids: return True - if self.router is None: + if router is None: return False - origin: Final = self.router.get_deployment(model_id=origin_model_id) + origin: Final = router.get_deployment(model_id=origin_model_id) origin_boundary: Final = ( - self._encryption_boundary_key(origin.litellm_params.model_dump(exclude_none=True)) + EncryptedContentAffinityCheck._encryption_boundary_key( + origin.litellm_params.model_dump(exclude_none=True) + ) if origin is not None else None ) return origin_boundary is not None and origin_boundary in target_boundaries def should_strip_input_item(item: Mapping[str, object]) -> bool: - origin_model_id: Final = self._model_id_of_input_item(item) - return origin_model_id is not None and not target_can_decrypt(origin_model_id) + return not target_can_decrypt(EncryptedContentAffinityCheck._model_id_of_input_item(item)) def should_strip_anthropic_block(block: Mapping[str, object]) -> bool: - origin_model_id: Final = self._model_id_of_anthropic_block(block) - return origin_model_id is not None and not target_can_decrypt(origin_model_id) + return not target_can_decrypt(EncryptedContentAffinityCheck._model_id_of_anthropic_block(block)) ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input( request_input, should_strip=should_strip_input_item @@ -317,12 +336,21 @@ class EncryptedContentAffinityCheck(CustomLogger): unhealthy, the request was routed to a different group by an auto-router tier change or model switch, or the marker is removed/unknown/forged), the encrypted reasoning is stripped and the request dispatches to the healthy - pool with its readable history instead of failing. + pool with its readable history instead of failing. An order-based fallback hop + carries ``_target_order``, and the pin only considers deployments of that order, + so the hop reaches the next order with the origin's reasoning stripped instead + of replaying it to a deployment that cannot decrypt it. """ request_kwargs = request_kwargs or {} - typed_healthy_deployments: Final = cast(list[dict], healthy_deployments) + typed_healthy_deployments: Final = cast(list[dict[str, object]], healthy_deployments) if not self._is_enabled_for_model_group(model): return typed_healthy_deployments + target_order: Final = _REQUEST_KWARGS_ADAPTER.validate_python(request_kwargs).get("_target_order") + candidates: Final = ( + get_order_filtered_deployments(typed_healthy_deployments, target_order=target_order) + if isinstance(target_order, int) + else typed_healthy_deployments + ) # Signal to the response post-processor that encrypted item IDs should be # encoded in the output of this request. Only set the flag when @@ -348,7 +376,7 @@ class EncryptedContentAffinityCheck(CustomLogger): ) deployment: Final = self._find_deployment_by_model_id( - healthy_deployments=typed_healthy_deployments, + healthy_deployments=candidates, model_id=model_id, ) if deployment is not None: @@ -357,12 +385,14 @@ class EncryptedContentAffinityCheck(CustomLogger): model_id, ) request_kwargs["_encrypted_content_affinity_pinned"] = True - self._strip_reasoning_the_target_cannot_decrypt(request_input, anthropic_messages, (deployment,)) + self.strip_reasoning_the_targets_cannot_decrypt( + self.router, request_input, anthropic_messages, (deployment,), unmarked_origin=model_id + ) return [deployment] # Follow-up switched model_name (LIT-2531): pin by Azure resource instead. boundary_matches, _originating = self._find_deployments_on_same_encryption_boundary( - healthy_deployments=typed_healthy_deployments, + healthy_deployments=candidates, model_id=model_id, ) if boundary_matches: @@ -373,7 +403,9 @@ class EncryptedContentAffinityCheck(CustomLogger): len(boundary_matches), ) request_kwargs["_encrypted_content_affinity_pinned"] = True - self._strip_reasoning_the_target_cannot_decrypt(request_input, anthropic_messages, boundary_matches) + self.strip_reasoning_the_targets_cannot_decrypt( + self.router, request_input, anthropic_messages, boundary_matches, unmarked_origin=model_id + ) return boundary_matches # The origin cannot serve this turn and no peer shares its encryption boundary, so its @@ -389,4 +421,4 @@ class EncryptedContentAffinityCheck(CustomLogger): ) ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input) strip_encrypted_reasoning_from_messages(anthropic_messages) - return typed_healthy_deployments + return candidates diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 5ca7b79d127..c83e4e5318c 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -110,12 +110,6 @@ def messages( def amessages( call: NativeCall, ) -> Coroutine[object, object, AnthropicMessagesResponse | AsyncIterator[bytes]]: ... -def chat_completions( - call: NativeCall, -) -> dict[str, object]: ... -def achat_completions( - call: NativeCall, -) -> Future[dict[str, object]]: ... @final class ResponsesWebSocketConnection: @@ -251,14 +245,12 @@ __all__ = [ "RustUpstreamError", "TokenCounter", "Tokenizer", - "achat_completions", "acompletion", "aembedding", "amessages", "aocr", "aresponses", "atranscription", - "chat_completions", "completion", "embedding", "gil_stats", diff --git a/litellm/rust_bridge/chat_completions/route_host.py b/litellm/rust_bridge/chat_completions/route_host.py index cb06ea8e213..d600495c919 100644 --- a/litellm/rust_bridge/chat_completions/route_host.py +++ b/litellm/rust_bridge/chat_completions/route_host.py @@ -37,16 +37,12 @@ def response(value: Mapping[str, object]) -> ModelResponse: return ModelResponse(**value) -def arguments(request: Mapping[str, object]) -> Mapping[str, object]: - return request - - def map_failure(error: Exception, request: Mapping[str, object]) -> Exception: provider: Final = optional_str(request.get("custom_llm_provider")) or str(request["model"]).partition("/")[0] return failures.map_native_failure( error, str(request["model"]), provider, - arguments(request), + request, optional_str(request.get("api_base")) or optional_str(request.get("base_url")), ) diff --git a/litellm/rust_bridge/messages/route_host.py b/litellm/rust_bridge/messages/route_host.py index 88a81fdae4a..69e27e78223 100644 --- a/litellm/rust_bridge/messages/route_host.py +++ b/litellm/rust_bridge/messages/route_host.py @@ -39,10 +39,6 @@ def stream_hidden_params(headers: Sequence[tuple[str, str]]) -> Mapping[str, obj return anthropic_messages_stream_hidden_params(httpx.Headers(list(headers))) -def arguments(request: Mapping[str, object]) -> Mapping[str, object]: - return request - - def map_failure(error: Exception, request: Mapping[str, object], request_provider: str) -> Exception: if getattr(error, "messages_request_error", False): return litellm.BadRequestError( @@ -51,7 +47,7 @@ def map_failure(error: Exception, request: Mapping[str, object], request_provide llm_provider=request_provider, ) return failures.map_native_failure( - error, str(request["model"]), request_provider, arguments(request), optional_str(request.get("api_base")) + error, str(request["model"]), request_provider, request, optional_str(request.get("api_base")) ) diff --git a/litellm/rust_bridge/ocr/route_host.py b/litellm/rust_bridge/ocr/route_host.py index a7eb0829f5c..8f563a381ef 100644 --- a/litellm/rust_bridge/ocr/route_host.py +++ b/litellm/rust_bridge/ocr/route_host.py @@ -12,7 +12,7 @@ from litellm.rust_bridge import failures from litellm.rust_bridge.failures import UpstreamFailure from litellm.rust_bridge.public_call import optional_str -__all__ = ("UpstreamFailure", "arguments", "map_failure", "response") +__all__ = ("UpstreamFailure", "map_failure", "response") _RESPONSE_ADAPTER: Final = TypeAdapter(dict[str, object]) @@ -27,10 +27,6 @@ def response(value: Mapping[str, object]) -> OCRResponse: return normalized -def arguments(request: Mapping[str, object]) -> Mapping[str, object]: - return request - - def map_failure(error: Exception, request: Mapping[str, object], request_provider: str) -> Exception: if getattr(error, "ocr_request_format_error", False): return litellm.UnsupportedParamsError( @@ -39,5 +35,5 @@ def map_failure(error: Exception, request: Mapping[str, object], request_provide llm_provider=request_provider, ) return failures.map_native_failure( - error, str(request["model"]), request_provider, arguments(request), optional_str(request.get("api_base")) + error, str(request["model"]), request_provider, request, optional_str(request.get("api_base")) ) diff --git a/litellm/rust_bridge/public_call.py b/litellm/rust_bridge/public_call.py index a25e9802593..e5700be89ad 100644 --- a/litellm/rust_bridge/public_call.py +++ b/litellm/rust_bridge/public_call.py @@ -84,17 +84,35 @@ def inference_decline_reason(parameters: tuple[str, ...], kwargs: Mapping[str, o return None +_BAGS: Final = frozenset({inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD}) + + @dataclass(frozen=True, slots=True) class NativeCall: + """A public call as Python bound it. + + ``base`` is the positional arguments by name plus the signature defaults, with no caller + keyword in it. ``kwargs`` is the caller's keyword dict, which callbacks may rewrite before + the request is decoded. The call the public function sees is ``kwargs`` laid over ``base``. + """ + args: tuple[object, ...] kwargs: Mapping[str, object] - bound: Mapping[str, object] + base: Mapping[str, object] + + @property + def resolved(self) -> Mapping[str, object]: + return MappingProxyType({**self.base, **self.kwargs}) -def native_call(args: tuple[object, ...], kwargs: Mapping[str, object], fields: Mapping[str, object]) -> NativeCall: - extra: Final = optional_mapping(fields.get("kwargs")) or MappingProxyType({}) - named: Final = {name: value for name, value in fields.items() if name != "kwargs"} - return NativeCall(args=args, kwargs=kwargs, bound=MappingProxyType({**named, **extra})) +def _without_bags(legacy: inspect.Signature, named: Mapping[str, object]) -> Mapping[str, object]: + return MappingProxyType({name: value for name, value in named.items() if legacy.parameters[name].kind not in _BAGS}) + + +def native_call(legacy: inspect.Signature, args: tuple[object, ...], kwargs: Mapping[str, object]) -> NativeCall: + positional: Final = legacy.bind_partial(*args) + positional.apply_defaults() + return NativeCall(args=args, kwargs=kwargs, base=_without_bags(legacy, positional.arguments)) NativeResultT: Final = TypeVar("NativeResultT") diff --git a/litellm/rust_bridge/responses/route_host.py b/litellm/rust_bridge/responses/route_host.py index 00f24244f49..d8ad2a94b75 100644 --- a/litellm/rust_bridge/responses/route_host.py +++ b/litellm/rust_bridge/responses/route_host.py @@ -20,17 +20,13 @@ def response(value: Mapping[str, object]) -> ResponsesAPIResponse: return ResponsesAPIResponse.model_validate(value) -def arguments(request: Mapping[str, object]) -> Mapping[str, object]: - return request - - def map_failure(error: Exception, request: Mapping[str, object]) -> Exception: provider: Final = optional_str(request.get("custom_llm_provider")) or "openai" return failures.map_native_failure( error, str(request["model"]), provider, - arguments(request), + request, optional_str(request.get("api_base")) or optional_str(request.get("base_url")), ) diff --git a/litellm/rust_bridge/tokenizer.py b/litellm/rust_bridge/tokenizer.py index 5ed89813620..1ae7f136a43 100644 --- a/litellm/rust_bridge/tokenizer.py +++ b/litellm/rust_bridge/tokenizer.py @@ -4,14 +4,16 @@ from functools import lru_cache from typing import TYPE_CHECKING, Final, cast # noqa: TID251 # native class is validated at the binding boundary import tiktoken -from tokenizers import Tokenizer as PythonHuggingFaceTokenizer -from litellm.litellm_core_utils.tokenizer import Encoding, HuggingFace, HuggingFaceTokenizer, OpenAIEncoding +from litellm.litellm_core_utils.tokenizer import HuggingFaceTokenizer, OpenAIEncoding from litellm.rust_bridge import runtime from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.catalog import Route, RouteContext if TYPE_CHECKING: + from tokenizers import Tokenizer as PythonHuggingFaceTokenizer + + from litellm.litellm_core_utils.tokenizer import Encoding, HuggingFace from litellm.rust_bridge._native import Tokenizer as NativeTokenizer @@ -76,6 +78,16 @@ def get_encoding(name: str) -> Encoding: ) +def _python_tokenizer() -> type[PythonHuggingFaceTokenizer]: + try: + from tokenizers import Tokenizer + except ModuleNotFoundError as error: + if error.name != "tokenizers": + raise + raise ImportError("Python tokenization requires tokenizers. Run 'pip install tokenizers'.") from error + return Tokenizer + + def anthropic() -> HuggingFace: """The packaged Anthropic tokenizer on the selected backend.""" from litellm.utils import claude_json_str @@ -84,7 +96,7 @@ def anthropic() -> HuggingFace: HUGGINGFACE_CONTEXT, binding=TOKENIZER, native=lambda factory: HuggingFaceTokenizer(_native_anthropic(factory)), - python=lambda: PythonHuggingFaceTokenizer.from_str(claude_json_str), + python=lambda: _python_tokenizer().from_str(claude_json_str), ) @@ -93,7 +105,7 @@ def from_str(json: str) -> HuggingFace: HUGGINGFACE_CONTEXT, binding=TOKENIZER, native=lambda factory: HuggingFaceTokenizer(factory.from_json(json)), - python=lambda: PythonHuggingFaceTokenizer.from_str(json), + python=lambda: _python_tokenizer().from_str(json), ) @@ -104,5 +116,5 @@ def from_pretrained(identifier: str, revision: str = "main", token: str | None = native=lambda factory: HuggingFaceTokenizer( factory.from_pretrained(identifier, revision=revision, token=token) ), - python=lambda: PythonHuggingFaceTokenizer.from_pretrained(identifier, revision=revision, token=token), + python=lambda: _python_tokenizer().from_pretrained(identifier, revision=revision, token=token), ) diff --git a/litellm/secret_managers/aws_secret_manager_v2.py b/litellm/secret_managers/aws_secret_manager_v2.py index a702c206cfd..98d12e497da 100644 --- a/litellm/secret_managers/aws_secret_manager_v2.py +++ b/litellm/secret_managers/aws_secret_manager_v2.py @@ -25,6 +25,7 @@ import httpx import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix +from litellm.litellm_core_utils.optional_imports import ensure_optional_import from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, @@ -627,11 +628,10 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): request_data: dict | None = None, ) -> tuple[str, "HTTPHeaders", bytes]: """Prepare the AWS Secrets Manager request""" - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + ensure_optional_import("botocore") + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest + optional_params = optional_params or {} # Build optional_params from instance settings if not provided diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 201b92f8107..ecc3d0d747c 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -2,11 +2,12 @@ from collections.abc import Mapping from datetime import datetime from enum import Enum from types import MappingProxyType -from typing import Final, Literal +from typing import Final, Literal, cast from pydantic import ConfigDict, Field, field_validator, model_validator from typing_extensions import ReadOnly, Required, TypedDict +from litellm._logging import verbose_logger from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import ( @@ -901,6 +902,93 @@ class ContentFilterConfigModel(LiteLLMBaseModel): MCP_SECURITY_ON_VIOLATION: Final = frozenset({"block", "alert"}) +GuardrailStreamScope = Literal["streaming", "non_streaming", "both"] +DEFAULT_GUARDRAIL_STREAM_SCOPE: Final[GuardrailStreamScope] = "both" + + +class GuardrailEventHooks(str, Enum): + pre_call = "pre_call" + post_call = "post_call" + during_call = "during_call" + logging_only = "logging_only" + pre_mcp_call = "pre_mcp_call" + during_mcp_call = "during_mcp_call" + post_mcp_call = "post_mcp_call" + realtime_input_transcription = "realtime_input_transcription" + + +GUARDRAIL_EVENT_HOOK_VALUES: Final = frozenset(member.value for member in GuardrailEventHooks) + +_GUARDRAIL_STREAM_SCOPES: Final[Mapping[str, GuardrailStreamScope]] = MappingProxyType( + { + "streaming": "streaming", + "non_streaming": "non_streaming", + "both": "both", + } +) + + +def _as_guardrail_stream_scope(value: object) -> GuardrailStreamScope: + if not isinstance(value, str): + raise ValueError(f"stream_scope values must be strings, got {type(value).__name__}") + scope: Final = _GUARDRAIL_STREAM_SCOPES.get(value.lower()) + if scope is None: + raise ValueError(f"stream_scope must be one of both, streaming, non_streaming, got {value!r}") + return scope + + +def _validated_stream_scope_hook(key: object) -> str: + if not isinstance(key, str): + raise ValueError(f"stream_scope keys must be strings, got {type(key).__name__}") + hook: Final = key.lower() + if hook not in GUARDRAIL_EVENT_HOOK_VALUES: + raise ValueError( + f"stream_scope keys must be guardrail modes ({sorted(GUARDRAIL_EVENT_HOOK_VALUES)}), got {key!r}" + ) + return hook + + +def coerce_stream_scope(value: object) -> GuardrailStreamScope | dict[str, GuardrailStreamScope] | None: + if value is None: + return None + if isinstance(value, str): + return _as_guardrail_stream_scope(value) + if isinstance(value, Mapping): + scope_map: Final[Mapping[str, object]] = cast(Mapping[str, object], value) # cast-ok: keys validated below + return { + _validated_stream_scope_hook(key): _as_guardrail_stream_scope(scope) for key, scope in scope_map.items() + } + raise ValueError(f"stream_scope must be a string or mapping, got {type(value).__name__}") + + +def stored_stream_scope(value: object) -> GuardrailStreamScope | dict[str, GuardrailStreamScope] | None: + try: + return coerce_stream_scope(value) + except ValueError: + verbose_logger.warning("Ignoring invalid stored stream_scope value of type %s", type(value).__name__) + return None + + +def with_tolerated_stream_scope(params: Mapping[str, object]) -> dict[str, object]: + if "stream_scope" not in params: + return dict(params) + return { + **params, + "stream_scope": stored_stream_scope(params["stream_scope"]), + } + + +def runtime_stream_scope( + stream_scope: object, +) -> tuple[GuardrailStreamScope, MappingProxyType[str, GuardrailStreamScope]]: + coerced: Final = coerce_stream_scope(stream_scope) + if coerced is None: + return DEFAULT_GUARDRAIL_STREAM_SCOPE, MappingProxyType({}) + if isinstance(coerced, str): + return coerced, MappingProxyType({}) + return DEFAULT_GUARDRAIL_STREAM_SCOPE, MappingProxyType(coerced) + + LoggingOnlyScope = Literal["input", "output", "both"] @@ -1144,6 +1232,21 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up ), ) + stream_scope: GuardrailStreamScope | dict[str, GuardrailStreamScope] | None = Field( + default=None, + description=( + "Whether this guardrail runs on streaming requests, non-streaming requests, or both. " + "A string applies to every configured mode. A map overrides named modes " + "(pre_call, during_call, post_call, ...); omitted keys default to both. " + "Unset means both, matching historical behavior." + ), + ) + + @field_validator("stream_scope", mode="before") + @classmethod + def normalize_stream_scope(cls, v: object) -> GuardrailStreamScope | dict[str, GuardrailStreamScope] | None: + return coerce_stream_scope(v) + logging_only_scope: LoggingOnlyScope | None = Field( default=None, description=( @@ -1278,17 +1381,6 @@ class guardrailConfig(TypedDict): guardrails: list[Guardrail] -class GuardrailEventHooks(str, Enum): - pre_call = "pre_call" - post_call = "post_call" - during_call = "during_call" - logging_only = "logging_only" - pre_mcp_call = "pre_mcp_call" - during_mcp_call = "during_mcp_call" - post_mcp_call = "post_mcp_call" - realtime_input_transcription = "realtime_input_transcription" - - class DynamicGuardrailParams(TypedDict): extra_body: ReadOnly[dict[str, object]] diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index a951381a2c9..9f888185779 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -204,6 +204,8 @@ class UserAPIKeyLabelNames(Enum): API_KEY_ALIAS = "api_key_alias" TEAM = "team" TEAM_ALIAS = "team_alias" + PROJECT_ID = "project_id" + PROJECT_ALIAS = "project_alias" REQUESTED_MODEL = REQUESTED_MODEL v1_LITELLM_MODEL_NAME = "model" v2_LITELLM_MODEL_NAME = "litellm_model_name" @@ -303,6 +305,8 @@ DEFINED_PROMETHEUS_METRICS = Literal[ "litellm_api_key_rate_limit_used_metric", "litellm_team_rate_limit_allowed_metric", "litellm_team_rate_limit_used_metric", + "litellm_project_model_rate_limit_allowed_metric", + "litellm_project_model_rate_limit_used_metric", "litellm_llm_api_failed_requests_metric", "litellm_callback_logging_failures_metric", "litellm_in_flight_requests", @@ -847,6 +851,15 @@ class PrometheusMetricLabels: litellm_team_rate_limit_used_metric = litellm_team_rate_limit_allowed_metric + litellm_project_model_rate_limit_allowed_metric: ClassVar[tuple[str, ...]] = ( + UserAPIKeyLabelNames.PROJECT_ID.value, + UserAPIKeyLabelNames.PROJECT_ALIAS.value, + UserAPIKeyLabelNames.REQUESTED_MODEL.value, + UserAPIKeyLabelNames.RATE_LIMIT_TYPE.value, + ) + + litellm_project_model_rate_limit_used_metric = litellm_project_model_rate_limit_allowed_metric + litellm_llm_api_failed_requests_metric = [ UserAPIKeyLabelNames.END_USER.value, UserAPIKeyLabelNames.API_KEY_HASH.value, @@ -1058,6 +1071,8 @@ class UserAPIKeyLabelValues: api_key_alias: str | None = None team: str | None = None team_alias: str | None = None + project_id: str | None = None + project_alias: str | None = None model_group: str | None = None requested_model: str | None = None model: str | None = None diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py index 17ab5ce223b..62029250269 100644 --- a/litellm/types/litellm_params.py +++ b/litellm/types/litellm_params.py @@ -87,6 +87,7 @@ class ProviderConnection: litellm_credential_name: str | None = None configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None use_xai_oauth: bool | None = None + fireworks_forward_user_id: bool | None = None @dataclass(frozen=True, slots=True, kw_only=True) diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index cc420fb9509..78d03ecc6e1 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -1,5 +1,6 @@ import json -from collections.abc import Sequence +from collections.abc import Mapping, Sequence +from dataclasses import dataclass from enum import Enum from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias @@ -1009,13 +1010,21 @@ else: AWSPreparedRequest = Any +@dataclass(frozen=True) +class BearerPreparedRequest: + method: str + url: str + headers: Mapping[str, str] + body: bytes + + class BedrockPreparedRequest(TypedDict): """ Internal/Helper class for preparing the request for bedrock image generation """ endpoint_url: str - prepped: AWSPreparedRequest + prepped: AWSPreparedRequest | BearerPreparedRequest body: bytes data: dict diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 9dc53cf2adf..492157f89fc 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -365,6 +365,7 @@ _JsonValue: TypeAlias = object BATCH_GUARDRAIL_RESPONSE_FIELD: Final = "litellm_batch_guardrail" +LITELLM_DETAILS_FALLBACK_RESPONSE_FIELD: Final = "litellm_details_fallback" class OpenAIFileObject(LiteLLMBaseModel): @@ -413,6 +414,9 @@ class OpenAIFileObject(LiteLLMBaseModel): Absent on every other upload, so OpenAI-shaped clients see an unchanged response. """ + litellm_details_fallback: bool | None = None + """Set by the LiteLLM proxy on a saved batch output file entry built without provider metadata; stripped from API responses.""" + _hidden_params: dict = PrivateAttr(default={"response_cost": 0.0}) # no cost for writing a file @property @@ -424,13 +428,19 @@ class OpenAIFileObject(LiteLLMBaseModel): self._hidden_params = hidden_params @model_serializer(mode="wrap") - def _omit_absent_batch_guardrail( # noqa: ANN202 # annotating it replaces the model's serialization schema + def _omit_absent_proxy_only_fields( # noqa: ANN202 # annotating it replaces the model's serialization schema self, handler: SerializerFunctionWrapHandler ): serialized: Final[Mapping[str, object]] = handler(self) - if self.litellm_batch_guardrail is not None: - return serialized - return {key: value for key, value in serialized.items() if key != BATCH_GUARDRAIL_RESPONSE_FIELD} + fields_to_omit: Final = tuple( + field_name + for field_name, value in ( + (BATCH_GUARDRAIL_RESPONSE_FIELD, self.litellm_batch_guardrail), + (LITELLM_DETAILS_FALLBACK_RESPONSE_FIELD, self.litellm_details_fallback), + ) + if value is None + ) + return {key: value for key, value in serialized.items() if key not in fields_to_omit} def __contains__(self, key) -> bool: # Define custom behavior for the 'in' operator diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index ce51e46ef15..116fdda0b5d 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -2,16 +2,20 @@ from enum import Enum from typing import Any, Final, Literal, Protocol from typing_extensions import ( + ReadOnly, Required, TypedDict, ) -from litellm.types.llms.openai import EmbeddingInput +from litellm.types.llms.openai import ChatCompletionFileObject, EmbeddingInput # Gemini supports nested-list inputs (e.g. [["text", "image"]]) as an explicit # opt-in for combined embeddings — a provider-specific extension of the # OpenAI-faithful EmbeddingInput shape. -GeminiEmbeddingInput = EmbeddingInput | list[list[str]] +GeminiEmbeddingElement = str | ChatCompletionFileObject +GeminiEmbeddingInput = ( + EmbeddingInput | list[GeminiEmbeddingElement] | list[list[str]] | list[list[GeminiEmbeddingElement]] +) class FunctionResponse(TypedDict, total=False): @@ -46,6 +50,12 @@ class FunctionResponsePartType(TypedDict, total=False): file_data: FileDataType +class VideoMetadataType(TypedDict, total=False): + fps: ReadOnly[float] + startOffset: ReadOnly[str] + endOffset: ReadOnly[str] + + class PartType(TypedDict, total=False): text: str inline_data: BlobType @@ -55,6 +65,7 @@ class PartType(TypedDict, total=False): thought: bool thoughtSignature: str media_resolution: Literal["low", "medium", "high"] + video_metadata: ReadOnly[VideoMetadataType] class HttpxFunctionCall(TypedDict, total=False): diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 2f5d983cdc0..a7027102250 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -39,6 +39,7 @@ class MCPOAuthMetadata(LiteLLMBaseModel): authorization_url: str | None = None token_url: str | None = None registration_url: str | None = None + client_id_metadata_document_supported: bool | None = None authorization_response_iss_parameter_supported: bool = False discovered_issuer: str | None = None """The ``issuer`` the authorization-server metadata document self-attests (RFC 8414). Persisted @@ -120,6 +121,7 @@ class MCPServer(LiteLLMBaseModel): client_secret: str | None = None issuer: str | None = None issuer_is_anchored: bool = False + client_id_metadata_document_supported: bool | None = None authorization_response_iss_parameter_supported: bool = False dcr_issuer: str | None = None dcr_server_url: str | None = None diff --git a/litellm/types/openai_decisions.py b/litellm/types/openai_decisions.py deleted file mode 100644 index 074c0a803f3..00000000000 --- a/litellm/types/openai_decisions.py +++ /dev/null @@ -1,162 +0,0 @@ -from collections.abc import Mapping, Sequence -from dataclasses import dataclass -from typing import Annotated, Literal, TypeAlias - -from pydantic import ConfigDict, Field, PrivateAttr, StrictBool, StrictStr - -from litellm.types.llms.base import LiteLLMPydanticObjectBase - -ChoiceValue: TypeAlias = StrictStr | StrictBool - - -class DecisionsObjectBase(LiteLLMPydanticObjectBase): - model_config = ConfigDict(extra="allow", frozen=True) - - -class DecisionInputText(DecisionsObjectBase): - type: Literal["input_text"] - text: str - - -class DecisionInputImage(DecisionsObjectBase): - type: Literal["input_image"] - image_url: str - detail: Literal["low", "high", "auto", "original"] | None = None - - -DecisionInputPart: TypeAlias = Annotated[DecisionInputText | DecisionInputImage, Field(discriminator="type")] - - -class DecisionInputMessage(DecisionsObjectBase): - role: Literal["user"] - content: str | Sequence[DecisionInputPart] - type: Literal["message"] | None = None - - -DecisionInput: TypeAlias = str | Sequence[DecisionInputMessage] - - -class DecisionChoice(DecisionsObjectBase): - value: ChoiceValue - description: str | None = None - - -class DecisionLevel(DecisionsObjectBase): - label: str - description: str | None = None - - -class PredicateQuestion(DecisionsObjectBase): - type: Literal["predicate"] - instructions: str - name: str | None = None - - -class ChoiceQuestion(DecisionsObjectBase): - type: Literal["choice"] - instructions: str - choices: Sequence[DecisionChoice] - name: str | None = None - - -class ScoreQuestion(DecisionsObjectBase): - type: Literal["score"] - instructions: str - levels: Sequence[DecisionLevel] - name: str | None = None - - -DecisionQuestion: TypeAlias = Annotated[ - PredicateQuestion | ChoiceQuestion | ScoreQuestion, - Field(discriminator="type"), -] - -DecisionQuestions: TypeAlias = Sequence[DecisionQuestion] - - -class DecisionsRequestBody(DecisionsObjectBase): - input: DecisionInput - questions: DecisionQuestions - safety_identifier: str | None = None - - -@dataclass(frozen=True, slots=True) -class DecisionsRequest: - model: str - body: DecisionsRequestBody - - -class PredicateAnswer(DecisionsObjectBase): - type: Literal["predicate"] - name: str | None = None - probability: float - - -class ChoiceProbability(DecisionsObjectBase): - value: ChoiceValue - probability: float - - -class ChoiceAnswer(DecisionsObjectBase): - type: Literal["choice"] - name: str | None = None - choice: ChoiceValue - probabilities: Sequence[ChoiceProbability] - confidence: float - - -class ScoreProbability(DecisionsObjectBase): - value: int - label: str - probability: float - - -class ScoreAnswer(DecisionsObjectBase): - type: Literal["score"] - name: str | None = None - score: float - probabilities: Sequence[ScoreProbability] - confidence: float - - -class RefusalAnswer(DecisionsObjectBase): - type: Literal["refusal"] - name: str | None = None - - -DecisionAnswer: TypeAlias = Annotated[ - PredicateAnswer | ChoiceAnswer | ScoreAnswer | RefusalAnswer, - Field(discriminator="type"), -] - - -class DecisionInputTokensDetails(DecisionsObjectBase): - cached_tokens: int - cache_write_tokens: int - - -class DecisionOutputTokensDetails(DecisionsObjectBase): - reasoning_tokens: int - - -class DecisionUsage(DecisionsObjectBase): - input_tokens: int - input_tokens_details: DecisionInputTokensDetails - output_tokens: int - output_tokens_details: DecisionOutputTokensDetails - total_tokens: int - - -class DecisionsResponse(DecisionsObjectBase): - model: str - answers: Sequence[DecisionAnswer] - usage: DecisionUsage - - _hidden_params: dict[str, object] = PrivateAttr(default_factory=dict) - - @property - def hidden_params(self) -> dict[str, object]: # mutable-ok: API requires mutation - return self._hidden_params - - def set_hidden_params(self, params: Mapping[str, object]) -> None: - self._hidden_params.update(params) diff --git a/litellm/types/proxy/policy_engine/pipeline_types.py b/litellm/types/proxy/policy_engine/pipeline_types.py index 5278754701d..0a089d2fe9b 100644 --- a/litellm/types/proxy/policy_engine/pipeline_types.py +++ b/litellm/types/proxy/policy_engine/pipeline_types.py @@ -87,7 +87,7 @@ class PipelineStepResult(LiteLLMBaseModel): """Result of executing a single pipeline step.""" guardrail_name: str - outcome: Literal["pass", "fail", "error"] + outcome: Literal["pass", "fail", "error", "skip"] action_taken: str modified_data: dict[str, Any] | None = None error_detail: str | None = None diff --git a/litellm/types/router.py b/litellm/types/router.py index a66c4571b39..67683e400ed 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -504,6 +504,10 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): default=False, description="Use stored xAI OAuth credentials when no xAI API key is configured.", ) + fireworks_forward_user_id: bool | None = Field( + default=None, + description="Send the LiteLLM user id of the calling key as the `user` field on Fireworks AI chat, responses and messages requests.", + ) model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True) merge_reasoning_content_in_choices: bool | None = False model_info: dict | None = None diff --git a/litellm/utils.py b/litellm/utils.py index 036b8c5b20a..a2992bda8b6 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -59,6 +59,7 @@ from litellm._lazy_imports import ( get_messages_reach_token_count, get_token_counter_new, ) +from litellm._logging import redact_secrets from litellm._uuid import uuid from litellm.constants import ( DEFAULT_CHAT_COMPLETION_PARAM_VALUES, @@ -90,7 +91,7 @@ from litellm.litellm_core_utils.fallback_generalizations import ( match_fill_missing_generalizations, ) from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload -from litellm.litellm_core_utils.tokenizer import Encoding, HuggingFace, strip_special_tokens +from litellm.litellm_core_utils.tokenizer import strip_special_tokens from litellm.rust_bridge import tokenizer as tokenizer_dispatch from litellm.rust_bridge.catalog import decision from litellm.rust_bridge.configuration import Decision @@ -369,6 +370,7 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.rules import Rules from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.litellm_core_utils.thread_pool_executor import BoundedLoggingThreadPoolExecutor + from litellm.litellm_core_utils.tokenizer import Encoding, HuggingFace from litellm.llms.base_llm.anthropic_messages.transformation import ( BaseAnthropicMessagesConfig, ) @@ -1895,17 +1897,18 @@ def client(original_function): # LOG SUCCESS - handle streaming success logging in the _next_ object, remove `handle_success` once it's deprecated verbose_logger.info("Wrapper: Completed Call, calling success_handler") - # Copy the current context to propagate it to the background thread - # This is essential for OpenTelemetry span context propagation - ctx: Final = contextvars.copy_context() - executor: Final[BoundedLoggingThreadPoolExecutor] = getattr(sys.modules[__name__], "executor") - executor.submit( - ctx.run, - logging_obj.success_handler, - result, - start_time, - end_time, - ) + if not is_internal_call.get(): + # Copy the current context to propagate it to the background thread + # This is essential for OpenTelemetry span context propagation + ctx: Final = contextvars.copy_context() + executor: Final[BoundedLoggingThreadPoolExecutor] = getattr(sys.modules[__name__], "executor") + executor.submit( + ctx.run, + logging_obj.success_handler, + result, + start_time, + end_time, + ) # RETURN RESULT return result except Exception as e: @@ -2238,7 +2241,7 @@ def client(original_function): is_acompletion_litellm_router_call: Final = _is_litellm_router_call(kwargs, is_async=True) if ( - num_retries and not is_acompletion_litellm_router_call + num_retries and not is_acompletion_litellm_router_call and not isinstance(e, ImportError) ): # only enter this if call is not from litellm router/proxy. router has it's own logic for retrying try: litellm.num_retries = None # set retries to None to prevent infinite loops @@ -2269,7 +2272,7 @@ def client(original_function): is_aresponses_litellm_router_call: Final = _is_litellm_router_call(kwargs, is_async=True) if ( - num_retries and not is_aresponses_litellm_router_call + num_retries and not is_aresponses_litellm_router_call and not isinstance(e, ImportError) ): # only enter this if call is not from litellm router/proxy. router has it's own logic for retrying try: litellm.num_retries = None # set retries to None to prevent infinite loops @@ -2422,7 +2425,12 @@ def _select_tokenizer_helper(model: str) -> SelectTokenizerResponse: if isinstance(e, (ForkedAfterNativeRuntimeStarted, ProcessReservedForForking)): raise - verbose_logger.debug("Error selecting tokenizer: %s", e) + verbose_logger.warning( + "Falling back to tiktoken for %s; token counts may be approximate. " + "For Python Hugging Face tokenization, install tokenizers and huggingface-hub. Error: %s", + json.dumps(redact_secrets(model)), + json.dumps(redact_secrets(str(e))), + ) # default - tiktoken return _return_openai_tokenizer(model) @@ -2493,7 +2501,7 @@ def encode(model="", text="", custom_tokenizer: dict | None = None): tokenizer_json: Final = custom_tokenizer or select_tokenizer(model=model) if tokenizer_json["type"] == "openai_tokenizer": openai_tokenizer: Final = cast( # cast-ok: [LIT006] caller's explicit type tag selects this interface - Encoding, tokenizer_json["tokenizer"] + "Encoding", tokenizer_json["tokenizer"] ) return openai_tokenizer.encode(text, disallowed_special=()) encoded: Final = tokenizer_json["tokenizer"].encode(text) @@ -2518,7 +2526,7 @@ def decode( if tokenizer_json["type"] == "huggingface_tokenizer": ids: Final = strip_special_tokens(tokenizer_json["tokenizer"], tokens) if skip_special_tokens else tokens hf_tokenizer: Final = cast( # cast-ok: [LIT006] caller's explicit type tag selects this interface - HuggingFace, tokenizer_json["tokenizer"] + "HuggingFace", tokenizer_json["tokenizer"] ) return hf_tokenizer.decode(ids, skip_special_tokens=skip_special_tokens) return tokenizer_json["tokenizer"].decode(tokens) @@ -2962,7 +2970,7 @@ def supports_pdf_input(model: str, custom_llm_provider: str | None = None) -> bo def supports_audio_output(model: str, custom_llm_provider: str | None = None) -> bool: """Check if a given model supports audio output in a chat completion call""" - return supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_audio_input") + return supports_factory(model=model, custom_llm_provider=custom_llm_provider, key="supports_audio_output") def supports_prompt_caching(model: str, custom_llm_provider: str | None = None) -> bool: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index f34606ce127..7930a823e07 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -342,12 +342,14 @@ }, "writer.palmyra-vision-7b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 4096, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_vision": true }, @@ -486,13 +488,16 @@ }, "us.amazon.nova-2-lite-v1:0": { "cache_read_input_token_cost": 8.25e-08, + "cache_read_input_token_cost_flex": 4.125e-08, "input_cost_per_token": 3.3e-07, + "input_cost_per_token_flex": 1.65e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 1000000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 2.75e-06, + "output_cost_per_token_flex": 1.375e-06, "supports_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -556,13 +561,16 @@ }, "amazon.nova-pro-v1:0": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_flex": 1e-07, "input_cost_per_token": 8e-07, + "input_cost_per_token_flex": 4e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 300000, "max_output_tokens": 10000, "max_tokens": 10000, "mode": "chat", "output_cost_per_token": 3.2e-06, + "output_cost_per_token_flex": 1.6e-06, "supports_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -3188,13 +3196,16 @@ }, "apac.amazon.nova-pro-v1:0": { "cache_read_input_token_cost": 2.1e-07, + "cache_read_input_token_cost_flex": 1.05e-07, "input_cost_per_token": 8.4e-07, + "input_cost_per_token_flex": 4.2e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 300000, "max_output_tokens": 10000, "max_tokens": 10000, "mode": "chat", "output_cost_per_token": 3.36e-06, + "output_cost_per_token_flex": 1.68e-06, "supports_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -13076,12 +13087,14 @@ }, "bedrock/ap-northeast-1/deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -13089,32 +13102,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/ap-northeast-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/ap-northeast-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13122,43 +13139,52 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/ap-northeast-1/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.3e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.03e-06, + "output_cost_per_token_flex": 1.52e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/ap-northeast-1/moonshotai.kimi-k2.5": { "input_cost_per_token": 7.2e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.6e-06, + "output_cost_per_token_flex": 1.8e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/ap-northeast-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13168,13 +13194,17 @@ "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-coder-next.html" }, "bedrock/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.3e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.03e-06, + "output_cost_per_token_flex": 1.52e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, @@ -13182,11 +13212,11 @@ "input_cost_per_token": 6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.03e-06, - "source": "https://platform.moonshot.ai/docs/guide/kimi-k2-5-quickstart", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13216,12 +13246,14 @@ }, "bedrock/ap-south-1/deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -13229,32 +13261,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/ap-south-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/ap-south-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13262,43 +13298,52 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/ap-south-1/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.1e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.94e-06, + "output_cost_per_token_flex": 1.47e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/ap-south-1/moonshotai.kimi-k2.5": { "input_cost_per_token": 7.2e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.6e-06, + "output_cost_per_token_flex": 1.8e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/ap-south-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13309,12 +13354,13 @@ }, "bedrock/ap-southeast-2/minimax.minimax-m2.5": { "input_cost_per_token": 3.1e-07, + "input_cost_per_token_flex": 1.55e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13322,16 +13368,19 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.24e-06 + "output_cost_per_token": 1.24e-06, + "output_cost_per_token_flex": 6.2e-07 }, "bedrock/ap-southeast-3/deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -13339,32 +13388,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/ap-southeast-3/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/ap-southeast-3/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13372,32 +13425,37 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/ap-southeast-3/moonshotai.kimi-k2.5": { "input_cost_per_token": 7.2e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.6e-06, + "output_cost_per_token_flex": 1.8e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/ap-southeast-3/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13426,12 +13484,14 @@ }, "bedrock/eu-north-1/deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -13439,32 +13499,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/eu-north-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/eu-north-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13472,23 +13536,26 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/eu-north-1/moonshotai.kimi-k2.5": { "input_cost_per_token": 7.2e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.6e-06, + "output_cost_per_token_flex": 1.8e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/eu-central-1/1-month-commitment/anthropic.claude-instant-v1": { "cost_per_second": 0.01635, @@ -13578,29 +13645,33 @@ "supports_tool_choice": true }, "bedrock/eu-central-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/eu-central-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13608,16 +13679,19 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/eu-central-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13645,29 +13719,33 @@ "output_cost_per_token": 6.5e-07 }, "bedrock/eu-west-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/eu-west-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13675,16 +13753,19 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/eu-west-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13712,29 +13793,33 @@ "output_cost_per_token": 7.8e-07 }, "bedrock/eu-west-2/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 4.7e-07, + "input_cost_per_token_flex": 2.35e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.86e-06, + "output_cost_per_token_flex": 9.3e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/eu-west-2/minimax.minimax-m2.5": { "input_cost_per_token": 4.7e-07, + "input_cost_per_token_flex": 2.35e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13742,16 +13827,19 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.86e-06 + "output_cost_per_token": 1.86e-06, + "output_cost_per_token_flex": 9.3e-07 }, "bedrock/eu-west-2/nvidia.nemotron-super-3-120b": { "input_cost_per_token": 2.3e-07, + "input_cost_per_token_flex": 1.15e-07, "litellm_provider": "bedrock", "max_input_tokens": 256000, "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 1.01e-06, + "output_cost_per_token_flex": 5.05e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -13763,12 +13851,14 @@ }, "bedrock/eu-west-2/qwen.qwen3-coder-next": { "input_cost_per_token": 7.8e-07, + "input_cost_per_token_flex": 3.9e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.86e-06, + "output_cost_per_token_flex": 9.3e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13809,29 +13899,33 @@ "supports_tool_choice": true }, "bedrock/eu-south-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/eu-south-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13839,16 +13933,19 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/eu-south-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -13895,12 +13992,14 @@ }, "bedrock/sa-east-1/deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -13908,32 +14007,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/sa-east-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/sa-east-1/minimax.minimax-m2.5": { "input_cost_per_token": 3.6e-07, + "input_cost_per_token_flex": 1.8e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -13941,43 +14044,52 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.44e-06 + "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07 }, "bedrock/sa-east-1/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.3e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.03e-06, + "output_cost_per_token_flex": 1.52e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/sa-east-1/moonshotai.kimi-k2.5": { "input_cost_per_token": 7.2e-07, + "input_cost_per_token_flex": 3.6e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3.6e-06, + "output_cost_per_token_flex": 1.8e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/sa-east-1/qwen.qwen3-coder-next": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.44e-06, + "output_cost_per_token_flex": 7.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -14125,12 +14237,14 @@ }, "bedrock/us-east-1/deepseek.v3.2": { "input_cost_per_token": 6.2e-07, + "input_cost_per_token_flex": 3.1e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.85e-06, + "output_cost_per_token_flex": 9.25e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -14138,29 +14252,34 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/us-east-1/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/us-east-1/minimax.minimax-m2.5": { "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", @@ -14170,43 +14289,51 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html" }, "bedrock/us-east-1/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_flex": 1.25e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/us-east-1/moonshotai.kimi-k2.5": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3e-06, + "output_cost_per_token_flex": 1.5e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/us-east-1/qwen.qwen3-coder-next": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -14217,12 +14344,14 @@ }, "bedrock/us-east-2/deepseek.v3.2": { "input_cost_per_token": 6.2e-07, + "input_cost_per_token_flex": 3.1e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.85e-06, + "output_cost_per_token_flex": 9.25e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -14230,32 +14359,36 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/us-east-2/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/us-east-2/minimax.minimax-m2.5": { "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html", "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, @@ -14263,43 +14396,52 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "output_cost_per_token": 1.2e-06 + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07 }, "bedrock/us-east-2/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_flex": 1.25e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/us-east-2/moonshotai.kimi-k2.5": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3e-06, + "output_cost_per_token_flex": 1.5e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/us-east-2/qwen.qwen3-coder-next": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -14755,12 +14897,14 @@ }, "bedrock/us-west-2/deepseek.v3.2": { "input_cost_per_token": 6.2e-07, + "input_cost_per_token_flex": 3.1e-07, "litellm_provider": "bedrock", "max_input_tokens": 163840, - "max_output_tokens": 163840, - "max_tokens": 163840, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.85e-06, + "output_cost_per_token_flex": 9.25e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -14768,29 +14912,34 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-2.html" }, "bedrock/us-west-2/minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html" }, "bedrock/us-west-2/minimax.minimax-m2.5": { "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "litellm_provider": "bedrock", - "max_input_tokens": 1000000, + "max_input_tokens": 196000, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", @@ -14800,43 +14949,51 @@ "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-5.html" }, "bedrock/us-west-2/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_flex": 1.25e-06, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_function_calling": true, "supports_reasoning": true }, "bedrock/us-west-2/moonshotai.kimi-k2.5": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3e-06, + "output_cost_per_token_flex": 1.5e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, "supports_audio_input": false, "supports_response_schema": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-5.html" }, "bedrock/us-west-2/qwen.qwen3-coder-next": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock", "max_input_tokens": 262144, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -15039,6 +15196,7 @@ "source": "https://developers.openai.com/api/docs/pricing" }, "claude-haiku-4-5-20251001": { + "supports_web_search": true, "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, "cache_creation_input_token_cost_batches": 6.25e-07, @@ -15181,6 +15339,7 @@ "source": "https://docs.anthropic.com/en/docs/about-claude/pricing" }, "claude-sonnet-5": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, @@ -15271,6 +15430,7 @@ "supports_web_search": true }, "claude-sonnet-4-6": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 3.75e-06, "cache_creation_input_token_cost_above_1hr": 6e-06, @@ -15344,6 +15504,7 @@ "output_cost_per_token_batches": 7.5e-06 }, "claude-opus-4-5-20251101": { + "supports_web_search": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, "cache_creation_input_token_cost_batches": 3.125e-06, @@ -15413,6 +15574,7 @@ "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-6": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, @@ -15498,6 +15660,7 @@ "prompt_cache_min_tokens": 4096 }, "claude-opus-4-7": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, @@ -15585,6 +15748,7 @@ "prompt_cache_min_tokens": 2048 }, "claude-fable-5": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, @@ -15629,6 +15793,7 @@ "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-fable-5-1": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, @@ -15674,6 +15839,7 @@ "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-5-5": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 5e-06, "cache_creation_input_token_cost_above_1hr": 8e-06, @@ -15719,6 +15885,7 @@ "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-5": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, @@ -15766,6 +15933,7 @@ "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, "claude-opus-4-8": { + "supports_web_search": true, "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, @@ -22487,14 +22655,15 @@ "supports_tool_choice": true }, "deepseek.v3-v1:0": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 5.8e-07, "litellm_provider": "bedrock_converse", - "max_input_tokens": 163840, - "max_output_tokens": 81920, - "max_tokens": 81920, + "max_input_tokens": 128000, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.68e-06, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-1.html", "supports_function_calling": true, "supports_reasoning": true, "supports_tool_choice": true, @@ -22502,12 +22671,14 @@ }, "deepseek.v3.2": { "input_cost_per_token": 6.2e-07, + "input_cost_per_token_flex": 3.1e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 164000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.85e-06, + "output_cost_per_token_flex": 9.25e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -23183,13 +23354,16 @@ }, "eu.amazon.nova-pro-v1:0": { "cache_read_input_token_cost": 2.625e-07, + "cache_read_input_token_cost_flex": 1.3125e-07, "input_cost_per_token": 1.05e-06, + "input_cost_per_token_flex": 5.25e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 300000, "max_output_tokens": 10000, "max_tokens": 10000, "mode": "chat", "output_cost_per_token": 4.2e-06, + "output_cost_per_token_flex": 2.1e-06, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_function_calling": true, "supports_pdf_input": true, @@ -29253,7 +29427,7 @@ }, "gemini/gemini-2.5-flash-image": { "cache_read_input_token_cost": 3e-08, - "deprecation_date": "2026-10-02", + "deprecation_date": "2027-03-15", "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 3e-07, "input_cost_per_token_batches": 1.5e-07, @@ -29703,6 +29877,7 @@ "web_search_billing_unit": "per_query" }, "gemini/gemini-2.5-flash-preview-tts": { + "deprecation_date": "2026-11-17", "input_cost_per_token": 5e-07, "input_cost_per_token_batches": 2.5e-07, "litellm_provider": "gemini", @@ -30276,7 +30451,7 @@ "supports_video_input": true, "supports_vision": true, "tpm": 800000, - "deprecation_date": "2026-09-30" + "deprecation_date": "2026-10-22" }, "gemini/gemini-3.1-pro-preview": { "prompt_cache_min_tokens": 4096, @@ -30783,6 +30958,7 @@ }, "gemini/gemini-2.5-pro-preview-tts": { "cache_read_input_token_cost": 1.25e-07, + "deprecation_date": "2026-11-17", "input_cost_per_audio_token": 7e-07, "input_cost_per_token": 1e-06, "input_cost_per_token_batches": 5e-07, @@ -32129,14 +32305,17 @@ "output_cost_per_token": 2e-06 }, "google.gemma-3-12b-it": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 9e-08, + "input_cost_per_token_flex": 5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.9e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 1.5e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-google-gemma-3-12b-it.html", "supports_audio_input": false, "supports_function_calling": true, "supports_response_schema": true, @@ -32144,14 +32323,17 @@ "supports_vision": true }, "google.gemma-3-27b-it": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 2.3e-07, + "input_cost_per_token_flex": 1.2e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 3.8e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 1.9e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-google-gemma-3-27b-pt.html", "supports_audio_input": false, "supports_function_calling": true, "supports_response_schema": true, @@ -32159,14 +32341,17 @@ "supports_vision": true }, "google.gemma-3-4b-it": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 4e-08, + "input_cost_per_token_flex": 2e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 8e-08, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 4e-08, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-google-gemma-3-4b-it.html", "supports_audio_input": false, "supports_function_calling": true, "supports_system_messages": true, @@ -38434,14 +38619,17 @@ "supports_tool_choice": true }, "minimax.minimax-m2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 1000000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.2e-06, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 6e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2.html", "supports_audio_input": false, "supports_function_calling": true, "supports_system_messages": true, @@ -38450,24 +38638,29 @@ "supports_vision": false }, "minimax.minimax-m2.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 196000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-minimax-minimax-m2-1.html", "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false }, "minimax.minimax-m2.5": { "input_cost_per_token": 3e-07, + "input_cost_per_token_flex": 1.5e-07, "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 196000, "max_output_tokens": 8000, @@ -38607,29 +38800,34 @@ "max_output_tokens": 128000 }, "mistral.devstral-2-123b": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 4e-07, + "input_cost_per_token_flex": 2e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 2e-06, + "output_cost_per_token_flex": 1e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-mistral-ai-devstral-2-123b.html", "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false }, "mistral.magistral-small-2509": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 40000, "max_tokens": 40000, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_flex": 7.5e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -38640,12 +38838,14 @@ }, "mistral.ministral-3-14b-instruct": { "input_cost_per_token": 2e-07, + "input_cost_per_token_flex": 1e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2e-07, + "output_cost_per_token_flex": 1e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -38656,12 +38856,14 @@ }, "mistral.ministral-3-3b-instruct": { "input_cost_per_token": 1e-07, + "input_cost_per_token_flex": 5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1e-07, + "output_cost_per_token_flex": 5e-08, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -38672,12 +38874,14 @@ }, "mistral.ministral-3-8b-instruct": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.5e-07, + "output_cost_per_token_flex": 7e-08, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -38722,12 +38926,14 @@ }, "mistral.mistral-large-3-675b-instruct": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_flex": 7.5e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -38761,27 +38967,33 @@ "supports_tool_choice": true }, "mistral.voxtral-mini-3b-2507": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 4e-08, + "input_cost_per_token_flex": 2e-08, "litellm_provider": "bedrock_converse", - "max_input_tokens": 128000, + "max_input_tokens": 32768, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 4e-08, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 2e-08, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-mistral-ai-voxtral-mini-3b-2507.html", "supports_audio_input": true, "supports_system_messages": true, "supports_native_structured_output": true }, "mistral.voxtral-small-24b-2507": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 1e-07, + "input_cost_per_token_flex": 5e-08, "litellm_provider": "bedrock_converse", - "max_input_tokens": 128000, + "max_input_tokens": 32768, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 3e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 1.5e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-mistral-ai-voxtral-small-24b-2507.html", "supports_audio_input": true, "supports_system_messages": true, "supports_native_structured_output": true @@ -39262,8 +39474,8 @@ "cache_read_input_token_cost": 6.8e-08, "input_cost_per_token": 6.8e-07, "litellm_provider": "mistral", - "max_input_tokens": 524288, - "max_tokens": 524288, + "max_input_tokens": 1048576, + "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 2.09e-06, "reasoning_effort_levels": [ @@ -39283,8 +39495,8 @@ "cache_read_input_token_cost": 6.8e-08, "input_cost_per_token": 6.8e-07, "litellm_provider": "mistral", - "max_input_tokens": 524288, - "max_tokens": 524288, + "max_input_tokens": 1048576, + "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 2.09e-06, "reasoning_effort_levels": [ @@ -39590,6 +39802,7 @@ "supports_vision": true }, "moonshot.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, @@ -39597,7 +39810,7 @@ "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.5e-06, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html", "supports_audio_input": false, "supports_function_calling": true, "supports_reasoning": true, @@ -39608,12 +39821,14 @@ }, "moonshotai.kimi-k2.5": { "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 3e-06, + "output_cost_per_token_flex": 1.5e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -40708,14 +40923,17 @@ "source": "https://tokenfactory.nebius.com/models/catalog/embedding/Qwen%2FQwen3-Embedding-8B" }, "nvidia.nemotron-nano-12b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 2e-07, + "input_cost_per_token_flex": 1e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 3e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-12b-v2-vl-bf16.html", "supports_audio_input": false, "supports_function_calling": true, "supports_response_schema": true, @@ -40723,14 +40941,17 @@ "supports_vision": true }, "nvidia.nemotron-nano-9b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-08, + "input_cost_per_token_flex": 3e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.3e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "output_cost_per_token_flex": 1.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-9b-v2.html", "supports_audio_input": false, "supports_function_calling": true, "supports_response_schema": true, @@ -40739,12 +40960,14 @@ }, "nvidia.nemotron-nano-3-30b": { "input_cost_per_token": 6e-08, + "input_cost_per_token_flex": 3e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 2.4e-07, + "output_cost_per_token_flex": 1.2e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -40756,12 +40979,14 @@ }, "nvidia.nemotron-super-3-120b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", "output_cost_per_token": 6.5e-07, + "output_cost_per_token_flex": 3.25e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, @@ -42095,14 +42320,15 @@ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "input_cost_per_token": 1.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -42113,14 +42339,15 @@ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "input_cost_per_token": 7e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 3e-07, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -42128,14 +42355,17 @@ }, "openai.gpt-oss-safeguard-120b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_response_schema": true, "supports_system_messages": true, @@ -42143,14 +42373,17 @@ }, "openai.gpt-oss-safeguard-20b": { "input_cost_per_token": 7e-08, + "input_cost_per_token_flex": 3e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2e-07, + "output_cost_per_token_flex": 1e-07, "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_response_schema": true, "supports_system_messages": true, @@ -42415,6 +42648,33 @@ "prompt_cache_min_tokens": 512, "supports_sampling_params": false }, + "openrouter/anthropic/claude-haiku-5.5:batch": { + "cache_creation_input_token_cost": 6.25e-08, + "cache_creation_input_token_cost_above_100k_tokens": 3.125e-07, + "cache_creation_input_token_cost_above_1hr": 1e-07, + "cache_creation_input_token_cost_above_1hr_above_100k_tokens": 5e-07, + "cache_read_input_token_cost": 5e-09, + "cache_read_input_token_cost_above_100k_tokens": 2.5e-08, + "input_cost_per_token": 5e-08, + "input_cost_per_token_above_100k_tokens": 2.5e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 2.5e-07, + "output_cost_per_token_above_100k_tokens": 1.25e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true + }, "openrouter/anthropic/claude-haiku-4.5": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, @@ -42600,14 +42860,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v3.2": { - "cache_read_input_token_cost": 2.8e-08, + "cache_read_input_token_cost": 1.35e-07, "deprecation_date": "2026-09-28", - "input_cost_per_token": 2.8e-07, - "input_cost_per_token_cache_hit": 2.8e-08, + "input_cost_per_token": 2.59e-07, + "input_cost_per_token_cache_hit": 1.35e-07, "litellm_provider": "openrouter", "max_input_tokens": 163840, - "max_output_tokens": 65536, - "max_tokens": 65536, + "max_output_tokens": 147456, + "max_tokens": 147456, "mode": "chat", "output_cost_per_token": 4.2e-07, "supports_assistant_prefill": true, @@ -42690,14 +42950,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "cache_read_input_token_cost": 1.74e-08, - "input_cost_per_token": 2.088e-07, + "cache_read_input_token_cost": 2.39395e-08, + "input_cost_per_token": 2.87274e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 4.176e-07, + "output_cost_per_token": 5.74548e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42710,14 +42970,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "cache_read_input_token_cost": 1.5e-08, - "input_cost_per_token": 4.4e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 3e-07, + "output_cost_per_token": 1.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -43304,14 +43564,14 @@ "supports_web_search": false }, "openrouter/nvidia/nemotron-3.5-lightning": { - "cache_read_input_token_cost": 3e-08, - "input_cost_per_token": 6e-08, + "cache_read_input_token_cost": 2.45e-08, + "input_cost_per_token": 4.9e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.6e-07, + "output_cost_per_token": 1.4e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -44055,13 +44315,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.5-35b-a3b": { - "input_cost_per_token": 1.5e-07, + "input_cost_per_token": 1.625e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 235929, - "max_tokens": 235929, + "max_output_tokens": 65536, + "max_tokens": 65536, "mode": "chat", - "output_cost_per_token": 1e-06, + "output_cost_per_token": 1.3e-06, "source": "https://openrouter.ai/api/v1/models", "supports_function_calling": true, "supports_reasoning": true, @@ -44075,13 +44335,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.5-27b": { - "input_cost_per_token": 1.95e-07, + "input_cost_per_token": 2.6e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 65536, - "max_tokens": 65536, + "max_output_tokens": 81920, + "max_tokens": 81920, "mode": "chat", - "output_cost_per_token": 1.56e-06, + "output_cost_per_token": 2.6e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -44153,14 +44413,14 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.5-397b-a17b": { - "cache_read_input_token_cost": 2.2e-07, - "input_cost_per_token": 4.5e-07, + "cache_read_input_token_cost": 2.25e-07, + "input_cost_per_token": 5.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 81920, - "max_tokens": 81920, + "max_output_tokens": 235929, + "max_tokens": 235929, "mode": "chat", - "output_cost_per_token": 3e-06, + "output_cost_per_token": 3.5e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -45287,12 +45547,14 @@ }, "qwen.qwen3-235b-a22b-2507-v1:0": { "input_cost_per_token": 2.2e-07, + "input_cost_per_token_flex": 1.1e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 262144, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 8.8e-07, + "output_cost_per_token_flex": 4.4e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_tool_choice": true, @@ -45318,12 +45580,14 @@ }, "qwen.qwen3-32b-v1:0": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 32768, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-32b.html", "supports_function_calling": true, "supports_reasoning": true, @@ -45460,12 +45724,14 @@ }, "qwen.qwen3-coder-next": { "input_cost_per_token": 5e-07, + "input_cost_per_token_flex": 2.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 256000, "max_output_tokens": 16000, "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -47124,7 +47390,6 @@ }, "together_ai/deepseek-ai/DeepSeek-V4-Flash-0731": { "cache_read_input_token_cost": 3e-08, - "deprecation_date": "2026-09-29", "input_cost_per_token": 1.4e-07, "litellm_provider": "together_ai", "max_input_tokens": 1048576, @@ -47154,7 +47419,6 @@ }, "together_ai/deepseek-ai/DeepSeek-V4-Pro-0813": { "cache_read_input_token_cost": 1.3e-07, - "deprecation_date": "2026-09-29", "input_cost_per_token": 1.32e-06, "litellm_provider": "together_ai", "max_input_tokens": 1048576, @@ -47366,13 +47630,16 @@ }, "us.amazon.nova-pro-v1:0": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_flex": 1e-07, "input_cost_per_token": 8e-07, + "input_cost_per_token_flex": 4e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 300000, "max_output_tokens": 10000, "max_tokens": 10000, "mode": "chat", "output_cost_per_token": 3.2e-06, + "output_cost_per_token_flex": 1.6e-06, "supports_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -47781,6 +48048,7 @@ "supports_vision": false }, "us-gov.nvidia.nemotron-nano-12b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 2.4e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, @@ -47788,6 +48056,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-12b-v2-vl-bf16.html", "supports_system_messages": true, "supports_vision": true, "supports_audio_input": false, @@ -47795,6 +48064,7 @@ "supports_response_schema": true }, "us-gov.nvidia.nemotron-nano-9b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.2e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, @@ -47802,6 +48072,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 2.76e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-9b-v2.html", "supports_system_messages": true, "supports_audio_input": false, "supports_function_calling": true, @@ -47829,14 +48100,16 @@ "input_cost_per_token": 8.4e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 3.6e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -47846,14 +48119,16 @@ "input_cost_per_token": 1.8e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -48082,12 +48357,14 @@ }, "us.deepseek.v3.2": { "input_cost_per_token": 6.2e-07, + "input_cost_per_token_flex": 3.1e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 163840, "max_output_tokens": 163840, "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 1.85e-06, + "output_cost_per_token_flex": 9.25e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -48098,12 +48375,14 @@ }, "eu.deepseek.v3.2": { "input_cost_per_token": 7.4e-07, + "input_cost_per_token_flex": 3.7e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 163840, "max_output_tokens": 163840, "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 2.22e-06, + "output_cost_per_token_flex": 1.11e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_native_structured_output": true, @@ -52898,18 +53177,21 @@ "supports_web_search": true }, "zai.glm-4.7": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 203000, "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 2.2e-06, + "output_cost_per_token_flex": 1.1e-06, "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-zai-glm-4-7.html", "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false @@ -52933,18 +53215,21 @@ "supports_vision": false }, "zai.glm-4.7-flash": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7e-08, + "input_cost_per_token_flex": 3.5e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 203000, "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", "output_cost_per_token": 4e-07, + "output_cost_per_token_flex": 2e-07, "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-zai-glm-4-7-flash.html", "supports_audio_input": false, "supports_response_schema": true, "supports_vision": false @@ -53201,7 +53486,7 @@ ] }, "azure/sora-2": { - "deprecation_date": "2026-10-15", + "deprecation_date": "2026-11-02", "litellm_provider": "azure", "mode": "video_generation", "output_cost_per_video_per_second": 0.1, @@ -57519,6 +57804,7 @@ "gemini_native_audio": true }, "gemini-2.5-flash-native-audio-preview-12-2025": { + "deprecation_date": "2026-11-17", "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", @@ -57549,6 +57835,7 @@ "supports_web_search": true }, "gemini-3.1-flash-live-preview": { + "deprecation_date": "2026-11-17", "input_cost_per_audio_token": 3e-06, "input_cost_per_image_token": 1e-06, "input_cost_per_token": 7.5e-07, @@ -57710,6 +57997,7 @@ "gemini_native_audio": true }, "gemini/gemini-2.5-flash-native-audio-preview-12-2025": { + "deprecation_date": "2026-11-17", "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", @@ -57743,6 +58031,7 @@ "supports_reasoning": true }, "gemini/gemini-3.1-flash-live-preview": { + "deprecation_date": "2026-11-17", "input_cost_per_audio_token": 3e-06, "input_cost_per_image_token": 1e-06, "input_cost_per_token": 7.5e-07, @@ -57782,6 +58071,7 @@ "supports_reasoning": true }, "gemini/gemini-3.1-flash-tts-preview": { + "deprecation_date": "2026-11-17", "input_cost_per_token": 1e-06, "input_cost_per_token_batches": 5e-07, "litellm_provider": "gemini", @@ -57853,6 +58143,7 @@ "supports_prompt_caching": true }, "gemini-2.5-flash-preview-tts": { + "deprecation_date": "2026-11-17", "input_cost_per_token": 5e-07, "input_cost_per_token_batches": 2.5e-07, "litellm_provider": "gemini", @@ -58243,12 +58534,15 @@ }, "bedrock_mantle/openai.gpt-oss-120b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -58261,12 +58555,15 @@ }, "bedrock_mantle/openai.gpt-oss-20b": { "input_cost_per_token": 7e-08, + "input_cost_per_token_flex": 3.5e-08, "output_cost_per_token": 3e-07, + "output_cost_per_token_flex": 1.5e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -58279,12 +58576,15 @@ }, "bedrock_mantle/openai.gpt-oss-safeguard-120b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7e-08, "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 65536, - "max_tokens": 65536, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-safeguard-120b.html", "supported_endpoints": [ "/v1/chat/completions" ], @@ -58295,12 +58595,15 @@ }, "bedrock_mantle/openai.gpt-oss-safeguard-20b": { "input_cost_per_token": 7e-08, + "input_cost_per_token_flex": 3e-08, "output_cost_per_token": 2e-07, + "output_cost_per_token_flex": 1e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 65536, - "max_tokens": 65536, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-safeguard-20b.html", "supported_endpoints": [ "/v1/chat/completions" ], @@ -58344,6 +58647,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58387,6 +58691,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58419,6 +58724,7 @@ "supports_function_calling": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58495,6 +58801,7 @@ "supports_max_reasoning_effort": true, "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58882,6 +59189,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58920,6 +59228,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -58958,6 +59267,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -59357,7 +59667,9 @@ }, "bedrock_mantle/google.gemma-4-31b": { "input_cost_per_token": 1.4e-07, + "input_cost_per_token_flex": 7e-08, "output_cost_per_token": 4e-07, + "output_cost_per_token_flex": 2e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 256000, @@ -59376,7 +59688,9 @@ }, "bedrock_mantle/google.gemma-4-26b-a4b": { "input_cost_per_token": 1.3e-07, + "input_cost_per_token_flex": 6.5e-08, "output_cost_per_token": 4e-07, + "output_cost_per_token_flex": 2e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 256000, @@ -59395,7 +59709,9 @@ }, "bedrock_mantle/google.gemma-4-e2b": { "input_cost_per_token": 4e-08, + "input_cost_per_token_flex": 2e-08, "output_cost_per_token": 8e-08, + "output_cost_per_token_flex": 4e-08, "litellm_provider": "bedrock_mantle", "max_input_tokens": 128000, "max_output_tokens": 128000, @@ -65580,6 +65896,7 @@ "supports_vision": false }, "bedrock/us-gov-west-1/nvidia.nemotron-nano-12b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 2.4e-07, "litellm_provider": "bedrock", "max_input_tokens": 128000, @@ -65587,6 +65904,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-12b-v2-vl-bf16.html", "supports_system_messages": true, "supports_vision": true, "supports_audio_input": false, @@ -65594,6 +65912,7 @@ "supports_response_schema": true }, "bedrock/us-gov-west-1/nvidia.nemotron-nano-9b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.2e-08, "litellm_provider": "bedrock", "max_input_tokens": 128000, @@ -65601,6 +65920,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 2.76e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-9b-v2.html", "supports_system_messages": true, "supports_audio_input": false, "supports_function_calling": true, @@ -65628,14 +65948,16 @@ "input_cost_per_token": 8.4e-08, "litellm_provider": "bedrock", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 3.6e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -65645,14 +65967,16 @@ "input_cost_per_token": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -65842,6 +66166,7 @@ "supports_vision": false }, "bedrock/us-gov-east-1/nvidia.nemotron-nano-12b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 2.4e-07, "litellm_provider": "bedrock", "max_input_tokens": 128000, @@ -65849,6 +66174,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-12b-v2-vl-bf16.html", "supports_system_messages": true, "supports_vision": true, "supports_audio_input": false, @@ -65856,6 +66182,7 @@ "supports_response_schema": true }, "bedrock/us-gov-east-1/nvidia.nemotron-nano-9b-v2": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 7.2e-08, "litellm_provider": "bedrock", "max_input_tokens": 128000, @@ -65863,6 +66190,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 2.76e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-nvidia-nvidia-nemotron-nano-9b-v2.html", "supports_system_messages": true, "supports_audio_input": false, "supports_function_calling": true, @@ -65890,14 +66218,16 @@ "input_cost_per_token": 8.4e-08, "litellm_provider": "bedrock", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 3.6e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -65907,14 +66237,16 @@ "input_cost_per_token": 1.8e-07, "litellm_provider": "bedrock", "max_input_tokens": 128000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 7.2e-07, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions" ], "supports_bedrock_runtime_chat_completions_tools_with_reasoning": true, + "supports_bedrock_runtime_chat_completions_inline_reasoning": true, "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, @@ -66292,9 +66624,10 @@ "output_cost_per_token": 3.6e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -66310,9 +66643,10 @@ "output_cost_per_token": 7.2e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -66449,9 +66783,10 @@ "output_cost_per_token": 3.6e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-20b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -66467,9 +66802,10 @@ "output_cost_per_token": 7.2e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 131072, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-oss-120b.html", "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -66481,8 +66817,11 @@ "supports_tool_choice": true }, "bedrock_mantle/deepseek.v3.1": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 5.8e-07, + "input_cost_per_token_flex": 2.9e-07, "output_cost_per_token": 1.68e-06, + "output_cost_per_token_flex": 8.4e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 128000, "max_output_tokens": 8000, @@ -66496,8 +66835,11 @@ "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-1.html" }, "bedrock_mantle/moonshotai.kimi-k2-thinking": { + "deprecation_date": "2027-03-30", "input_cost_per_token": 6e-07, + "input_cost_per_token_flex": 3e-07, "output_cost_per_token": 2.5e-06, + "output_cost_per_token_flex": 1.25e-06, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 16000, @@ -66513,7 +66855,9 @@ }, "bedrock_mantle/qwen.qwen3-235b-a22b-2507": { "input_cost_per_token": 2.2e-07, + "input_cost_per_token_flex": 1.1e-07, "output_cost_per_token": 8.8e-07, + "output_cost_per_token_flex": 4.4e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 8000, @@ -66529,7 +66873,9 @@ }, "bedrock_mantle/qwen.qwen3-32b": { "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 32000, "max_output_tokens": 8000, @@ -66546,7 +66892,9 @@ "bedrock_mantle/qwen.qwen3-coder-30b-a3b-instruct": { "deprecation_date": "2027-03-30", "input_cost_per_token": 1.5e-07, + "input_cost_per_token_flex": 7.5e-08, "output_cost_per_token": 6e-07, + "output_cost_per_token_flex": 3e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 16000, @@ -66562,7 +66910,9 @@ "bedrock_mantle/qwen.qwen3-coder-480b-a35b-instruct": { "deprecation_date": "2027-03-30", "input_cost_per_token": 4.5e-07, + "input_cost_per_token_flex": 2.25e-07, "output_cost_per_token": 1.8e-06, + "output_cost_per_token_flex": 9e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 128000, "max_output_tokens": 16000, @@ -66577,7 +66927,9 @@ }, "bedrock_mantle/qwen.qwen3-next-80b-a3b-instruct": { "input_cost_per_token": 1.4e-07, + "input_cost_per_token_flex": 7e-08, "output_cost_per_token": 1.2e-06, + "output_cost_per_token_flex": 6e-07, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 8000, @@ -66593,7 +66945,9 @@ }, "bedrock_mantle/qwen.qwen3-vl-235b-a22b-instruct": { "input_cost_per_token": 5.3e-07, + "input_cost_per_token_flex": 2.6e-07, "output_cost_per_token": 2.66e-06, + "output_cost_per_token_flex": 1.33e-06, "litellm_provider": "bedrock_mantle", "max_input_tokens": 256000, "max_output_tokens": 8000, @@ -67997,14 +68351,14 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "cache_read_input_token_cost": 6.5e-08, - "input_cost_per_token": 7e-08, + "cache_read_input_token_cost": 3.8e-08, + "input_cost_per_token": 3.9e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 7e-06, + "output_cost_per_token": 3.39e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68135,8 +68489,8 @@ "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-flash-0731": { - "cache_read_input_token_cost": 1.63e-08, - "input_cost_per_token": 1.63e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 6e-09, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -68226,14 +68580,14 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k3": { - "cache_read_input_token_cost": 2.6e-07, - "input_cost_per_token": 7.9e-07, + "cache_read_input_token_cost": 4e-07, + "input_cost_per_token": 7.6e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 1.5e-05, + "output_cost_per_token": 1.3e-05, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68351,14 +68705,14 @@ "supports_web_search": true }, "openrouter/z-ai/glm-5.2": { - "cache_read_input_token_cost": 1.5e-07, - "input_cost_per_token": 1.52e-07, + "cache_read_input_token_cost": 6.8e-08, + "input_cost_per_token": 6.9e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 1.2e-05, + "output_cost_per_token": 4.3e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68654,13 +69008,13 @@ }, "openrouter/qwen/qwen3.6-27b": { "cache_read_input_token_cost": 3e-08, - "input_cost_per_token": 3e-07, + "input_cost_per_token": 4.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", - "output_cost_per_token": 2e-06, + "output_cost_per_token": 2.7e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68714,8 +69068,8 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "cache_read_input_token_cost": 3e-08, - "input_cost_per_token": 3e-08, + "cache_read_input_token_cost": 7.5e-09, + "input_cost_per_token": 7.5e-09, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -68734,14 +69088,14 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k2.6": { - "cache_read_input_token_cost": 1.6e-07, - "input_cost_per_token": 9.5e-07, + "cache_read_input_token_cost": 9.75e-08, + "input_cost_per_token": 4.65e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, "max_tokens": 235929, "mode": "chat", - "output_cost_per_token": 4e-06, + "output_cost_per_token": 2.45e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68968,6 +69322,7 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.5-9b": { + "deprecation_date": "2026-10-21", "input_cost_per_token": 1e-07, "output_cost_per_token": 1.5e-07, "litellm_provider": "openrouter", @@ -69161,8 +69516,8 @@ "supports_web_search": false }, "openrouter/nvidia/nemotron-3-nano-30b-a3b": { - "input_cost_per_token": 5e-08, - "output_cost_per_token": 2e-07, + "input_cost_per_token": 6e-08, + "output_cost_per_token": 2.4e-07, "cache_read_input_token_cost": 3e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, @@ -69183,7 +69538,7 @@ "openrouter/z-ai/glm-4.6v": { "input_cost_per_token": 3e-07, "output_cost_per_token": 9e-07, - "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost": 5.5e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, @@ -69609,13 +69964,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-next-80b-a3b-instruct": { - "input_cost_per_token": 1e-07, + "input_cost_per_token": 9e-08, "output_cost_per_token": 1.1e-06, "cache_read_input_token_cost": 7e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 235929, - "max_tokens": 235929, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -69771,13 +70126,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-30b-a3b-instruct-2507": { - "input_cost_per_token": 4.815e-08, + "input_cost_per_token": 1e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 32000, - "max_tokens": 32000, + "max_output_tokens": 235929, + "max_tokens": 235929, "mode": "chat", - "output_cost_per_token": 1.9305e-07, + "output_cost_per_token": 3e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -70530,7 +70885,7 @@ "supports_web_search": false }, "openrouter/mistralai/mistral-nemo": { - "input_cost_per_token": 1.9e-08, + "input_cost_per_token": 2.9e-08, "output_cost_per_token": 3e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, @@ -73416,16 +73771,21 @@ "supports_web_search": true }, "openrouter/~anthropic/claude-haiku-latest": { - "cache_creation_input_token_cost": 1.25e-06, - "cache_creation_input_token_cost_above_1hr": 2e-06, - "cache_read_input_token_cost": 1e-07, - "input_cost_per_token": 1e-06, + "input_cost_per_token_above_100k_tokens": 5e-07, + "output_cost_per_token_above_100k_tokens": 2.5e-06, + "cache_read_input_token_cost_above_100k_tokens": 5e-08, + "cache_creation_input_token_cost_above_100k_tokens": 6.25e-07, + "cache_creation_input_token_cost_above_1hr_above_100k_tokens": 1e-06, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_1hr": 2e-07, + "cache_read_input_token_cost": 1e-08, + "input_cost_per_token": 1e-07, "litellm_provider": "openrouter", - "max_input_tokens": 200000, - "max_output_tokens": 64000, - "max_tokens": 64000, + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 5e-06, + "output_cost_per_token": 5e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -73462,7 +73822,7 @@ "openrouter/~anthropic/claude-sonnet-latest": { "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, - "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost": 1e-07, "input_cost_per_token": 2e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, @@ -73684,8 +74044,8 @@ "openrouter/~openai/gpt-sol-latest": { "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_272k_tokens": 5e-06, - "cache_read_input_token_cost": 2e-07, - "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_above_272k_tokens": 2e-07, "input_cost_per_token": 2e-06, "input_cost_per_token_above_272k_tokens": 4e-06, "litellm_provider": "openrouter", @@ -79771,7 +80131,7 @@ "supports_vision": true, "supports_pdf_input": true, "supports_audio_input": false, - "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost": 1e-07, "supports_prompt_caching": true, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, @@ -80109,33 +80469,10 @@ "supports_tool_choice": true, "supports_vision": true }, - "xai.grok-4.7": { - "supports_regex_lookaround": false, - "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_flex": 2.5e-07, - "cache_read_input_token_cost_priority": 8.75e-07, - "input_cost_per_token": 2e-06, - "input_cost_per_token_flex": 1e-06, - "input_cost_per_token_priority": 3.5e-06, - "litellm_provider": "bedrock_converse", - "max_input_tokens": 500000, - "max_output_tokens": 500000, - "max_tokens": 500000, - "mode": "chat", - "output_cost_per_token": 6e-06, - "output_cost_per_token_flex": 3e-06, - "output_cost_per_token_priority": 1.05e-05, - "source": "https://b0.p.awsstatic.com/pricing/2.0/meteredUnitMaps/bedrock/USD/current/bedrock.json", - "supports_function_calling": true, - "supports_prompt_caching": false, - "supports_reasoning": true, - "supports_tool_choice": true, - "supports_vision": true - }, "openrouter/anthropic/claude-sonnet-5.5:batch": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, - "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost": 5e-08, "input_cost_per_token": 1e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, @@ -80155,12 +80492,12 @@ "supports_web_search": true }, "baseten/deepseek-ai/DeepSeek-V4.1-Flash-Fast": { - "cache_read_input_token_cost": 1.4e-07, + "cache_read_input_token_cost": 1.4e-08, "input_cost_per_token": 6e-07, "litellm_provider": "baseten", "max_input_tokens": 1048576, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 2.4e-06, "source": "https://inference.baseten.co/v1/models", @@ -80177,25 +80514,31 @@ "cache_creation_input_token_cost_above_272k_tokens_batches": 2.5e-06, "cache_creation_input_token_cost_above_272k_tokens_flex": 2.5e-06, "cache_creation_input_token_cost_above_272k_tokens_priority": 1e-05, + "cache_creation_input_token_cost_above_272k_tokens_ultrafast": 3e-05, "cache_creation_input_token_cost_batches": 1.25e-06, "cache_creation_input_token_cost_flex": 1.25e-06, "cache_creation_input_token_cost_priority": 5e-06, + "cache_creation_input_token_cost_ultrafast": 1.5e-05, "cache_read_input_token_cost": 1e-07, "cache_read_input_token_cost_above_272k_tokens": 2e-07, "cache_read_input_token_cost_above_272k_tokens_batches": 1e-07, "cache_read_input_token_cost_above_272k_tokens_flex": 1e-07, "cache_read_input_token_cost_above_272k_tokens_priority": 4e-07, + "cache_read_input_token_cost_above_272k_tokens_ultrafast": 1.2e-06, "cache_read_input_token_cost_batches": 5e-08, "cache_read_input_token_cost_flex": 5e-08, "cache_read_input_token_cost_priority": 2e-07, + "cache_read_input_token_cost_ultrafast": 6e-07, "input_cost_per_token": 2e-06, "input_cost_per_token_above_272k_tokens": 4e-06, "input_cost_per_token_above_272k_tokens_batches": 2e-06, "input_cost_per_token_above_272k_tokens_flex": 2e-06, "input_cost_per_token_above_272k_tokens_priority": 8e-06, + "input_cost_per_token_above_272k_tokens_ultrafast": 2.4e-05, "input_cost_per_token_batches": 1e-06, "input_cost_per_token_flex": 1e-06, "input_cost_per_token_priority": 4e-06, + "input_cost_per_token_ultrafast": 1.2e-05, "litellm_provider": "openai", "max_input_tokens": 922000, "max_output_tokens": 128000, @@ -80206,9 +80549,11 @@ "output_cost_per_token_above_272k_tokens_batches": 7.5e-06, "output_cost_per_token_above_272k_tokens_flex": 7.5e-06, "output_cost_per_token_above_272k_tokens_priority": 3e-05, + "output_cost_per_token_above_272k_tokens_ultrafast": 9e-05, "output_cost_per_token_batches": 5e-06, "output_cost_per_token_flex": 5e-06, "output_cost_per_token_priority": 2e-05, + "output_cost_per_token_ultrafast": 6e-05, "regional_processing_uplift_multiplier_eu": 1.1, "regional_processing_uplift_multiplier_us": 1.1, "search_context_cost_per_query": { @@ -80285,39 +80630,6 @@ "supports_sampling_params": false, "supports_xhigh_reasoning_effort": true }, - "openai.gpt-6.1-sol": { - "cache_creation_input_token_cost": 2.5e-06, - "cache_creation_input_token_cost_above_272k_tokens": 5e-06, - "cache_read_input_token_cost": 1e-07, - "cache_read_input_token_cost_above_272k_tokens": 2e-07, - "input_cost_per_token": 2e-06, - "input_cost_per_token_above_272k_tokens": 4e-06, - "litellm_provider": "bedrock_converse", - "max_input_tokens": 1050000, - "max_output_tokens": 131072, - "max_tokens": 131072, - "mode": "chat", - "output_cost_per_token": 1e-05, - "output_cost_per_token_above_272k_tokens": 1.5e-05, - "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-1-sol.html", - "supported_modalities": [ - "text", - "image" - ], - "supported_output_modalities": [ - "text" - ], - "supports_function_calling": true, - "supports_max_reasoning_effort": true, - "supports_minimal_reasoning_effort": false, - "supports_none_reasoning_effort": false, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_sampling_params": false, - "supports_xhigh_reasoning_effort": true - }, "bedrock_mantle/openai.gpt-6.1-sol": { "cache_creation_input_token_cost": 2.75e-06, "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, @@ -80349,6 +80661,7 @@ "supports_minimal_reasoning_effort": false, "supports_none_reasoning_effort": false, "supports_prompt_caching": true, + "supports_prompt_cache_breakpoint": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, @@ -80740,7 +81053,7 @@ "cache_read_input_token_cost": 7e-08, "input_cost_per_token": 6.8e-07, "litellm_provider": "openrouter", - "max_input_tokens": 524288, + "max_input_tokens": 1048576, "max_output_tokens": 262144, "max_tokens": 262144, "mode": "chat", @@ -81529,5 +81842,66 @@ "thinking_always_on": false, "prompt_cache_min_tokens": 512, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing" + }, + "openrouter/inclusionai/ling-3.0-flash-sante": { + "cache_read_input_token_cost": 8.4e-09, + "input_cost_per_token": 4.2e-08, + "litellm_provider": "openrouter", + "max_input_tokens": 262144, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 1.232e-07, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": false, + "supports_tool_choice": true, + "supports_vision": false, + "supports_web_search": false + }, + "openrouter/stepfun/step-5-preview": { + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 1e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 2.7e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_pdf_input": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": false, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": false + }, + "us.twelvelabs.pegasus-1-5-v1:0": { + "input_cost_per_video_per_second": 0.00049, + "litellm_provider": "bedrock", + "max_output_tokens": 98304, + "max_tokens": 98304, + "mode": "chat", + "output_cost_per_token": 7.5e-06, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", + "supports_video_input": true + }, + "global.twelvelabs.pegasus-1-5-v1:0": { + "input_cost_per_video_per_second": 0.00049, + "litellm_provider": "bedrock", + "max_output_tokens": 98304, + "max_tokens": 98304, + "mode": "chat", + "output_cost_per_token": 7.5e-06, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", + "supports_video_input": true } } diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index e51ae0453ea..d3843b78ff4 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -1039,6 +1039,9 @@ "supports_audio_output": { "type": "boolean" }, + "supports_bedrock_runtime_chat_completions_inline_reasoning": { + "type": "boolean" + }, "supports_bedrock_runtime_chat_completions_response_format": { "type": "boolean" }, diff --git a/packaging/litellm-core/pyproject.toml b/packaging/litellm-core/pyproject.toml new file mode 100644 index 00000000000..e3e2ca3d1db --- /dev/null +++ b/packaging/litellm-core/pyproject.toml @@ -0,0 +1,75 @@ +[project] +name = "litellm-core" +version = "1.105.0" +description = "Library to easily interface with LLM API providers" +readme = "README.md" +requires-python = ">=3.10, <3.15" +license = "MIT" +license-files = ["LICENSE"] +authors = [ + { name = "BerriAI" }, +] +dependencies = [ + "fastuuid>=0.14.0,<1.0", + "filelock>=3.16.1,<4.0", + "httpx[http2]>=0.28.0,<1.0", + "openai>=2.20.0,<3.0.0", + "python-dateutil>=2.8.2,<3.0", + "python-dotenv>=1.0.0,<2.0", + "pyyaml>=6.0.3,<7.0", + "packaging>=24.0", + "importlib-metadata>=8.0.0,<9.0", + "tiktoken>=0.8.0,<1.0; python_version < '3.14'", + "tiktoken>=0.12.0,<1.0; python_version >= '3.14'", + "click>=8.0.0,<9.0", + "jinja2>=3.1.6,<4.0", + "aiohttp>=3.14.2,<4.0", + "async-timeout>=4.0.3,<6.0; python_version < '3.11'", + "pydantic>=2.11.0,<3.0.0; python_version < '3.14'", + "pydantic>=2.12.0,<3.0.0; python_version >= '3.14'", + "pydantic-settings>=2.14.1,<3.0", + "jsonschema>=4.0.0,<5.0", + "typing-extensions>=4.13.0,<5.0", +] + +[project.urls] +Homepage = "https://litellm.ai" +Repository = "https://github.com/BerriAI/litellm" +Documentation = "https://docs.litellm.ai" + +[build-system] +requires = ["maturin==1.15.0"] +build-backend = "maturin" + +[tool.maturin] +manifest-path = "litellm-rust/crates/python-bridge/Cargo.toml" +module-name = "litellm.rust_bridge._native" +python-source = "." +bindings = "pyo3" +features = ["extension-module"] +profile = "release" +editable-profile = "dev" +include = [ + { path = "rust-toolchain.toml", format = "sdist" }, + { path = ".cargo/config.toml", format = "sdist" }, + "litellm/router_strategy/complexity_router/artifacts/*.json", + "litellm/router_strategy/complexity_router/fuse_presets.json", + "litellm/proxy/model_insights_tasks.json", + "litellm/proxy/common_utils/codex_base_instructions.md", + "litellm/proxy/common_utils/codex_bundled_models_0.159.3.json", + "litellm/proxy/lens/prompts/*.md", +] +exclude = [ + "litellm/proxy/_experimental/out", + "litellm/proxy/_experimental/out/**", + "litellm/proxy/enterprise", + "litellm/proxy/enterprise/**", + "litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/guardrail_benchmarks", + "litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/guardrail_benchmarks/**", + "**/__pycache__", + "**/__pycache__/**", + "**/.pytest_cache", + "**/.pytest_cache/**", + "**/.ruff_cache", + "**/.ruff_cache/**", +] diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 00904219b81..62e23388581 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2711,7 +2711,7 @@ } }, "voyage": { - "display_name": "Voyage AI (`voyage`)", + "display_name": "VoyageAI by MongoDB (`voyage`)", "url": "https://docs.litellm.ai/docs/providers/voyage", "endpoints": { "chat_completions": false, diff --git a/pyproject.toml b/pyproject.toml index d9ec313471e..e670b10547e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -77,8 +77,8 @@ proxy = [ "mcp>=2.2.0,<3", "httpx2>=2.5.0,<3", "pydantic>=2.12.0,<3", - "litellm-proxy-extras==0.4.106", - "litellm-enterprise==0.1.74", + "litellm-proxy-extras==0.4.107", + "litellm-enterprise==0.1.75", "RestrictedPython>=8.5,<9.0", "rich>=13.9.4,<14.0", "InquirerPy>=0.3.4,<1.0", @@ -378,6 +378,7 @@ version_files = [ ] [tool.pytest.ini_options] +pythonpath = ["scripts"] asyncio_mode = "auto" asyncio_default_fixture_loop_scope = "session" markers = [ diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json deleted file mode 100644 index d63a69de76f..00000000000 --- a/ruff-strict-budget.json +++ /dev/null @@ -1,260 +0,0 @@ -{ - "ANN001": { - "limit": 2918 - }, - "ANN002": { - "limit": 71 - }, - "ANN003": { - "limit": 806 - }, - "ANN201": { - "limit": 1965 - }, - "ANN202": { - "limit": 829 - }, - "ANN204": { - "limit": 683 - }, - "ANN205": { - "limit": 112 - }, - "ANN206": { - "limit": 133 - }, - "ANN401": { - "limit": 119 - }, - "ASYNC230": { - "limit": 11 - }, - "B004": { - "limit": 2 - }, - "B006": { - "limit": 176 - }, - "B008": { - "limit": 503 - }, - "B009": { - "limit": 52 - }, - "B010": { - "limit": 187 - }, - "B018": { - "limit": 2 - }, - "B019": { - "limit": 1 - }, - "B021": { - "limit": 1 - }, - "B026": { - "limit": 3 - }, - "BLE001": { - "limit": 2914 - }, - "C401": { - "limit": 8 - }, - "C404": { - "limit": 1 - }, - "C405": { - "limit": 19 - }, - "C408": { - "limit": 11 - }, - "C414": { - "limit": 4 - }, - "C419": { - "limit": 1 - }, - "C901": { - "limit": 306 - }, - "D419": { - "limit": 6 - }, - "DTZ001": { - "limit": 2 - }, - "DTZ003": { - "limit": 24 - }, - "DTZ005": { - "limit": 233 - }, - "DTZ006": { - "limit": 10 - }, - "DTZ007": { - "limit": 6 - }, - "DTZ011": { - "limit": 3 - }, - "EXE001": { - "limit": 4 - }, - "EXE002": { - "limit": 3 - }, - "F401": { - "limit": 12 - }, - "LOG015": { - "limit": 5 - }, - "N999": { - "limit": 1 - }, - "PERF102": { - "limit": 21 - }, - "PERF401": { - "limit": 12 - }, - "PERF403": { - "limit": 33 - }, - "PIE804": { - "limit": 18 - }, - "PIE810": { - "limit": 43 - }, - "PLC0206": { - "limit": 26 - }, - "PLC0414": { - "limit": 46 - }, - "PLR0124": { - "limit": 1 - }, - "PLR0206": { - "limit": 1 - }, - "PLR1704": { - "limit": 1 - }, - "PLR1714": { - "limit": 253 - }, - "PLW0127": { - "limit": 57 - }, - "PLW0602": { - "limit": 215 - }, - "PLW0603": { - "limit": 190 - }, - "PLW1508": { - "limit": 190 - }, - "PLW1510": { - "limit": 2 - }, - "PYI036": { - "limit": 3 - }, - "RET504": { - "limit": 173 - }, - "RUF012": { - "limit": 239 - }, - "RUF015": { - "limit": 8 - }, - "RUF019": { - "limit": 27 - }, - "RUF046": { - "limit": 4 - }, - "RUF059": { - "limit": 66 - }, - "RUF100": { - "limit": 0 - }, - "S110": { - "limit": 207 - }, - "S112": { - "limit": 22 - }, - "SIM101": { - "limit": 56 - }, - "SIM102": { - "limit": 310 - }, - "SIM103": { - "limit": 119 - }, - "SIM113": { - "limit": 3 - }, - "SIM115": { - "limit": 2 - }, - "SIM117": { - "limit": 6 - }, - "SIM201": { - "limit": 1 - }, - "SIM210": { - "limit": 8 - }, - "SIM211": { - "limit": 1 - }, - "SIM222": { - "limit": 1 - }, - "SIM401": { - "limit": 11 - }, - "TC004": { - "limit": 5 - }, - "TID251": { - "limit": 1035 - }, - "TRY002": { - "limit": 524 - }, - "TRY004": { - "limit": 96 - }, - "TRY201": { - "limit": 401 - }, - "TRY203": { - "limit": 109 - }, - "TRY300": { - "limit": 852 - }, - "UP028": { - "limit": 2 - }, - "UP031": { - "limit": 2 - }, - "UP036": { - "limit": 1 - } -} diff --git a/ruff.toml b/ruff.toml index fab3fe27aed..93396a723e7 100644 --- a/ruff.toml +++ b/ruff.toml @@ -1,6 +1,6 @@ lint.ignore = ["F405", "E402", "F403"] # The second group is the strict gate's graduates: rules the codebase already has zero -# violations of, so they hard-fail here instead of being ratcheted in ruff-strict-budget.json. +# violations of, so they hard-fail here instead of being counted by the strict gate. # That gives editors and `ruff check --fix` the diagnostic, which the gate script cannot. lint.extend-select = [ "T20", "PGH004", "RUF008", "RUF009", "RUF100", diff --git a/scripts/budget_ratchet_check.py b/scripts/budget_ratchet_check.py deleted file mode 100644 index 3ca5e9f3e9d..00000000000 --- a/scripts/budget_ratchet_check.py +++ /dev/null @@ -1,260 +0,0 @@ -#!/usr/bin/env python3 -"""Non-gating ratchet guard: budget limits may only fall, never rise. - -Every `*-budget.json` file (ruff-strict, type-discipline, basedpyright-code) is a -one-way ratchet: each rule's ceiling is its `limit`, and that limit is meant to be -driven DOWN over time. This check compares every budget file against its own -content at the merge-base with the target branch and fails (exits 1, red) if: - - * a rule's `limit` went up, - * a rule was dropped from a budget (its ceiling effectively became infinite) while - its checker still emits it, or - * an entire budget file was deleted. - -New rules and lowered/equal limits are fine. So is a rule that graduated: once a -paired config (ruff.toml for the ruff-strict budget) selects the rule outright it -hard-fails at the first violation, which is stricter than any ceiling the budget -could hold, so dropping its entry tightens the guard rather than removing it. -Likewise a retired rule: once the paired checker (check_test_quality.py for the -test-quality budget) no longer emits a code, its entry has no ceiling left to -loosen. - -This is deliberately NOT a gating check. It should turn the run red so that a -loosening is impossible to miss in review, but it must stay OUT of the -branch-protection required-checks list: a justified bump (e.g. banning a new API, -which mechanically raises a baseline) can then still be merged by a human who has -seen the red and accepted it. - -Usage: - python scripts/budget_ratchet_check.py [--base REF] [budget.json ...] - -""" - -from __future__ import annotations - -import argparse -import importlib.util -import json -import subprocess -import sys -from pathlib import Path -from types import MappingProxyType, ModuleType -from typing import Final, NamedTuple - -if sys.version_info >= (3, 11): - import tomllib -else: - import tomli as tomllib - -REPO_ROOT = Path(__file__).resolve().parent.parent -DEFAULT_BUDGETS: tuple[str, ...] = ( - "ruff-strict-budget.json", - "type-discipline-budget.json", - "basedpyright-code-budget.json", - "test-quality-budget.json", -) -GRADUATION_CONFIGS = MappingProxyType({"ruff-strict-budget.json": "ruff.toml"}) -RETIREMENT_SOURCES = MappingProxyType({"test-quality-budget.json": "check_test_quality"}) - - -class Regression(NamedTuple): - budget: str - rule: str - detail: str - - -def _run(cmd: list[str]) -> subprocess.CompletedProcess[str]: - return subprocess.run(cmd, cwd=REPO_ROOT, capture_output=True, text=True) - - -def _merge_base(base: str) -> str: - """The common ancestor of `base` and HEAD, so unrelated base drift is ignored.""" - proc = _run(["git", "merge-base", base, "HEAD"]) - return proc.stdout.strip() or base - - -def _load_head(rel: str) -> dict | None: - path = REPO_ROOT / rel - if not path.exists(): - return None - return json.loads(path.read_text()) - - -def _ref_is_commit(ref: str) -> bool: - return ( - _run( - ["git", "rev-parse", "--verify", "--quiet", f"{ref}^{{commit}}"] - ).returncode - == 0 - ) - - -def _load_base(rel: str, ref: str) -> dict | None: - """Budget content at `ref`, or None when the file did not exist there. - - `ref` is verified as a real commit by the caller, so a non-zero `git show` here means - the path was absent at that commit, not that the ref itself is unresolvable. - """ - proc = _run(["git", "show", f"{ref}:{rel}"]) - if proc.returncode != 0: - return None - return json.loads(proc.stdout) - - -def _ceiling(spec: dict) -> int: - """A rule's ceiling: its `limit`, or legacy `baseline + slack`. - - The base side of the diff can predate the `limit` migration, so a spec is read - under either schema and the two are compared on the same footing. - """ - if "limit" in spec: - return int(spec["limit"]) - return int(spec.get("baseline", 0)) + int(spec.get("slack", 0)) - - -def _limits(budget: dict) -> dict[str, int]: - """Map each rule to its ceiling; skip malformed specs.""" - return { - rule: _ceiling(spec) - for rule, spec in budget.items() - if isinstance(spec, dict) - } - - -def selectors_hard_failed_by(lint: dict) -> tuple[str, ...]: - """A ruff `[lint]` table's selected codes, minus anything `ignore` turns back off. - - `lint.ignore` wins over `lint.extend-select` in ruff, so an ignored code is not - actually enforced and must not count as a graduation. - """ - ignored = tuple(lint.get("ignore", ())) - return tuple( - selector - for selector in lint.get("extend-select", ()) - if not (ignored and selector.startswith(ignored)) - ) - - -def graduated_selectors(rel: str) -> tuple[str, ...]: - """Selectors the budget's paired ruff config hard-fails, so its ceiling is moot.""" - config = GRADUATION_CONFIGS.get(rel) - if config is None or not (REPO_ROOT / config).exists(): - return () - return selectors_hard_failed_by( - tomllib.loads((REPO_ROOT / config).read_text()).get("lint", {}) - ) - - -def _load_script(name: str) -> ModuleType: - if name in sys.modules: - return sys.modules[name] - spec: Final = importlib.util.spec_from_file_location(name, REPO_ROOT / "scripts" / f"{name}.py") - assert spec is not None and spec.loader is not None - module: Final = importlib.util.module_from_spec(spec) - sys.modules[name] = module - spec.loader.exec_module(module) - return module - - -def retired_rules(rel: str, base: dict[str, object]) -> frozenset[str]: - """Rules in the base budget that the paired checker can no longer emit, so there is no ceiling to loosen.""" - source: Final = RETIREMENT_SOURCES.get(rel) - if source is None: - return frozenset() - return frozenset(_limits(base)) - _load_script(source).RULE_CODES - - -def _regression_detail( - rule: str, - base_limits: dict[str, int], - head_limits: dict[str, int], - graduated: tuple[str, ...], - retired: frozenset[str] = frozenset(), -) -> str | None: - """Why `rule` regressed vs base, or None when it held flat, fell, or left the budget legitimately. - - A dropped rule is terminal unless it graduated or retired; otherwise the only - loosening left is a raised limit. - """ - base_limit = base_limits[rule] - if rule not in head_limits: - if rule in retired or (graduated and rule.startswith(graduated)): - return None - return f"rule dropped (limit {base_limit} -> removed)" - if head_limits[rule] > base_limit: - return f"limit raised {base_limit} -> {head_limits[rule]}" - return None - - -def regressions_for( - rel: str, - base: dict | None, - head: dict | None, - graduated: tuple[str, ...] = (), - retired: frozenset[str] = frozenset(), -) -> list[Regression]: - if base is None: - return [] # new budget file: nothing to ratchet against yet - if head is None: - return [Regression(rel, "*", "budget file was deleted (every limit removed)")] - - base_limits, head_limits = _limits(base), _limits(head) - return [ - Regression(rel, rule, detail) - for rule in sorted(base_limits) - if (detail := _regression_detail(rule, base_limits, head_limits, graduated, retired)) is not None - ] - - -def main() -> int: - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--base", help="Comparison ref (default: origin's current default branch)") - parser.add_argument("budgets", nargs="*", help="budget files to check") - args = parser.parse_args() - from default_branch import resolve_base_ref - - base_ref: Final = resolve_base_ref(args.base, REPO_ROOT) - budgets = args.budgets or list(DEFAULT_BUDGETS) - - ref = _merge_base(base_ref) - if not _ref_is_commit(ref): - print( - f"FAIL: base ref {ref!r} does not resolve to a commit, so the ratchet has nothing " - f"to compare against; refusing to pass vacuously (check the --base / BASE_SHA value)", - file=sys.stderr, - ) - return 1 - - regressions: list[Regression] = [] - checked: list[str] = [] - for rel in budgets: - base = _load_base(rel, ref) - head = _load_head(rel) - if base is None and head is None: - continue - if base is None: - print(f"skip {rel}: new file (no base at {base_ref} to ratchet against)") - continue - checked.append(rel) - regressions.extend(regressions_for(rel, base, head, graduated_selectors(rel), retired_rules(rel, base))) - - if regressions: - print( - f"FAIL: budget limit(s) loosened vs base {base_ref} (merge-base {ref[:12]}):" - ) - for reg in regressions: - print(f" {reg.budget} {reg.rule}: {reg.detail}") - print( - "Budgets are one-way ratchets and may only go down or stay flat. This " - "check is non-gating: if the increase is justified (e.g. a newly banned " - "API), a human can merge over the red after acknowledging it." - ) - return 1 - - suffix = f" ({', '.join(checked)})" if checked else "" - print(f"OK: no budget limit increased vs base {base_ref}{suffix}") - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/scripts/build_core_distribution.py b/scripts/build_core_distribution.py new file mode 100644 index 00000000000..36ec20a9031 --- /dev/null +++ b/scripts/build_core_distribution.py @@ -0,0 +1,86 @@ +"""Build the independent core distribution from shared repository sources.""" + +import argparse +import shutil +import subprocess +import sys +import tempfile +from pathlib import Path +from typing import Final + +ROOT: Final = Path(__file__).resolve().parents[1] +SOURCES: Final = ("litellm", "litellm-rust", ".cargo", "rust-toolchain.toml", "README.md", "LICENSE") + + +def stage_core_distribution(source: Path, destination: Path) -> None: + version: Final = subprocess.run( + ["uv", "version", "--short"], cwd=source, check=True, capture_output=True, text=True + ).stdout.strip() + ignored: Final = frozenset( + source / name + for name in subprocess.run( + [ + "git", + "ls-files", + "--ignored", + "--cached", + "--others", + "--exclude-standard", + "--directory", + "-z", + "--", + *SOURCES, + ], + cwd=source, + check=True, + capture_output=True, + text=True, + ).stdout.split("\0") + if name + ) + + def ignored_sources(directory: str, names: list[str]) -> set[str]: + return {name for name in names if Path(directory) / name in ignored} | shutil.ignore_patterns( + "__pycache__", ".pytest_cache", ".ruff_cache", "target", ".git", "*.so", "*.pyd" + )(directory, names) + + destination.mkdir(parents=True, exist_ok=True) + for name in SOURCES: + path: Final = source / name + if path.is_dir(): + shutil.copytree( + path, + destination / name, + ignore=ignored_sources, + ) + else: + shutil.copy2(path, destination / name) + shutil.copy2(source / "packaging/litellm-core/pyproject.toml", destination / "pyproject.toml") + subprocess.run(["uv", "version", version, "--frozen"], cwd=destination, check=True) + + +def build_core_distribution(output: Path, *, sdist_only: bool = False) -> None: + output_path: Final = output.resolve() + with tempfile.TemporaryDirectory(prefix="litellm-core-") as temporary: + stage: Final = Path(temporary) + stage_core_distribution(ROOT, stage) + subprocess.run( + ["uv", "build", "--python", sys.executable, "--out-dir", str(output_path)] + + (["--sdist"] if sdist_only else []), + cwd=stage, + check=True, + ) + + +def main() -> None: + parser: Final = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--out-dir", type=Path, default=ROOT / "dist/core") + parser.add_argument( + "--sdist-only", action="store_true", help="Build only the source archive, without compiling a wheel" + ) + args: Final = parser.parse_args() + build_core_distribution(Path(args.out_dir), sdist_only=args.sdist_only) + + +if __name__ == "__main__": + main() diff --git a/scripts/check_test_quality.py b/scripts/check_test_quality.py index 9f93023cd53..6f40dd6ea84 100644 --- a/scripts/check_test_quality.py +++ b/scripts/check_test_quality.py @@ -4,8 +4,8 @@ Sibling of scripts/check_type_discipline.py, same output contract (``path:line: CODE message``) and same stdlib-only constraint, aimed at the test tree instead of the package. Each rule is a shape the testing-strategy audit -measured and named; scripts/test_quality_gate.py caps the codebase total of each -one against test-quality-budget.json so the counts can only ratchet down. +measured and named; scripts/test_quality_gate.py fails any change that grows the +codebase total of one past its merge-base count. Rules ----- diff --git a/scripts/check_type_discipline.py b/scripts/check_type_discipline.py index 1adafd33490..3e645e02bc0 100644 --- a/scripts/check_type_discipline.py +++ b/scripts/check_type_discipline.py @@ -19,8 +19,8 @@ LIT003 noqa suppression without rule codes or without a reason. LIT004 pyright/mypy ignore without bracketed codes or without a reason. Required shape: `# pyright: ignore[reportArgumentType] # ` LIT005 A `# mutable-ok` / `# cast-ok` / `# guard-ok` / `# kwargs-ok` / - `# rebind-ok` / `# writable-ok` / `# comprehension-ok` suppression - without a reason. + `# rebind-ok` / `# writable-ok` / `# comprehension-ok` / `# frozen-ok` + suppression without a reason. LIT006 `cast(...)` call. typing.cast is an unchecked assertion (the moral equivalent of TypeScript's `as`); it lies to the type checker with zero runtime guarantee. Validate into a concrete frozen type at the boundary instead. @@ -92,6 +92,16 @@ LIT014 Comprehension with more than one `for` clause or more than one `if` clau `# comprehension-ok: ` on any line the comprehension spans. The marker belongs to the innermost violating comprehension spanning that line, and also to any single-line violating comprehension on that line. +LIT015 Pydantic model class that is not frozen. Set `frozen=True` in + `model_config = ConfigDict(...)`, `SettingsConfigDict(...)`, a dict-literal + `model_config`, an inner `class Config`, or the class keywords. Classes + inherit the frozen setting from in-module model bases, unless their own + configuration overrides it. Detection is name-based: `BaseModel`, + `pydantic.BaseModel`, `LiteLLMBaseModel`, `BaseLiteLLMOpenAIResponseObject`, + `LiteLLMPydanticObjectBase`, `OpenAIObject`, `RootModel`, and `BaseSettings` identify + models, while `TypedDict` classes are exempt. Replace in-place field writes + with `model_copy(update=...)`. Suppress with `# frozen-ok: ` on the + `class` line. LIT000 Setup failure: a target file could not be read, or contains a syntax error. Reported as a violation rather than crashing the run. @@ -111,12 +121,13 @@ import os import re import sys import tokenize +from collections.abc import Iterable, Iterator, Mapping, Sequence from dataclasses import dataclass +from itertools import groupby from multiprocessing import Pool from pathlib import Path -from collections.abc import Iterable, Iterator, Mapping, Sequence from types import MappingProxyType -from typing import NamedTuple +from typing import Final, NamedTuple # Mutable sequence and set types, banned in *every* annotation. Name-based, so `list`, # `typing.List`, `collections.deque`, and `collections.abc.MutableSequence` all match @@ -144,6 +155,19 @@ READONLY_QUALIFIER = "ReadOnly" # first argument is type syntax, the rest is metadata and never qualifies the field. FIELD_QUALIFIER_WRAPPERS = frozenset(("Required", "NotRequired", "Annotated")) TYPEDDICT_BASE = "TypedDict" +# Base names that mark a class as a pydantic model (LIT015). +PYDANTIC_BASES: Final = frozenset( + ( + "BaseModel", + "LiteLLMBaseModel", + "BaseLiteLLMOpenAIResponseObject", + "LiteLLMPydanticObjectBase", + "OpenAIObject", + "RootModel", + "BaseSettings", + ) +) +PYDANTIC_CONFIG_FACTORIES: Final = frozenset(("ConfigDict", "SettingsConfigDict")) MIN_REASON_LEN = 3 NOQA_RE = re.compile( @@ -161,6 +185,8 @@ KWARGS_OK_RE = re.compile(r"#\s*kwargs-ok(?::\s*(?P.*))?") REBIND_OK_RE = re.compile(r"#\s*rebind-ok(?::\s*(?P.*))?") WRITABLE_OK_RE = re.compile(r"#\s*writable-ok(?::\s*(?P.*))?") COMPREHENSION_OK_RE = re.compile(r"#\s*comprehension-ok(?::\s*(?P.*))?") +FROZEN_OK_RE: Final = re.compile(r"#\s*frozen-ok(?::\s*(?P.*))?") + @dataclass(frozen=True, slots=True) @@ -181,6 +207,7 @@ OK_SUPPRESSIONS: Final[tuple[_OkToken, ...]] = ( _OkToken("rebind-ok", REBIND_OK_RE, frozenset(("LIT010", "LIT011"))), _OkToken("writable-ok", WRITABLE_OK_RE, frozenset(("LIT012",))), _OkToken("comprehension-ok", COMPREHENSION_OK_RE, frozenset(("LIT014",))), + _OkToken("frozen-ok", FROZEN_OK_RE, frozenset(("LIT015",))), ) @@ -829,6 +856,147 @@ def iter_typeddict_violations(path: Path, tree: ast.AST) -> Iterator[Violation]: ) +# --------------------------------------------------------------------------- # +# Unfrozen pydantic models (LIT015) +# --------------------------------------------------------------------------- # + + +def _pydantic_classes(tree: ast.AST) -> tuple[ast.ClassDef, ...]: + classes: Final = tuple(node for node in ast.walk(tree) if isinstance(node, ast.ClassDef)) + + def expand(known: frozenset[str]) -> frozenset[str]: + grown: Final = known | frozenset(cls.name for cls in classes if _base_names(cls) & known) + return grown if grown == known else expand(grown) + + model_names: Final = expand(PYDANTIC_BASES) + typeddict_names: Final = expand(frozenset((TYPEDDICT_BASE,))) + return tuple(cls for cls in classes if _base_names(cls) & model_names and not _base_names(cls) & typeddict_names) + + +def _bool_constant(value: ast.expr) -> bool | None: + return value.value if isinstance(value, ast.Constant) and isinstance(value.value, bool) else None + + +def _module_assignment_items(tree: ast.AST) -> Iterator[tuple[str, tuple[int, ast.expr]]]: + if not isinstance(tree, ast.Module): + return + for stmt in tree.body: + if isinstance(stmt, ast.Assign): + yield from ( + (target.id, (stmt.lineno, stmt.value)) for target in stmt.targets if isinstance(target, ast.Name) + ) + elif isinstance(stmt, ast.AnnAssign) and isinstance(stmt.target, ast.Name) and stmt.value is not None: + yield stmt.target.id, (stmt.lineno, stmt.value) + + +def _model_config_frozen( + value: ast.expr, + module_assignments: Mapping[str, tuple[tuple[int, ast.expr], ...]], + class_lineno: int, +) -> bool | None: + config: Final = ( + next( + ( + expression + for lineno, expression in reversed(module_assignments.get(value.id, ())) + if lineno < class_lineno + ), + None, + ) + if isinstance(value, ast.Name) + else value + ) + if isinstance(config, ast.Call) and _head_name(config.func) in PYDANTIC_CONFIG_FACTORIES: + flags: Final = tuple(_bool_constant(kw.value) for kw in config.keywords if kw.arg == "frozen") + return flags[-1] if flags else None + if isinstance(config, ast.Dict): + flags: Final = tuple( + _bool_constant(item) + for key, item in zip(config.keys, config.values) + if isinstance(key, ast.Constant) and key.value == "frozen" + ) + return flags[-1] if flags else None + return None + + +def _assigns_name(stmt: ast.stmt, name: str) -> ast.expr | None: + value: Final = stmt.value if isinstance(stmt, (ast.Assign, ast.AnnAssign)) else None + targets: Final = ( + stmt.targets if isinstance(stmt, ast.Assign) else (stmt.target,) if isinstance(stmt, ast.AnnAssign) else () + ) + if value is not None and any(isinstance(target, ast.Name) and target.id == name for target in targets): + return value + return None + + +def _config_class_frozen(node: ast.ClassDef) -> bool | None: + values: Final = tuple(_assigns_name(stmt, "frozen") for stmt in node.body) + flags: Final = tuple(_bool_constant(value) for value in values if value is not None) + return flags[-1] if flags else None + + +def _stmt_frozen_flag( + stmt: ast.stmt, + module_assignments: Mapping[str, tuple[tuple[int, ast.expr], ...]], + class_lineno: int, +) -> bool | None: + config_value: Final = _assigns_name(stmt, "model_config") + if config_value is not None: + return _model_config_frozen(config_value, module_assignments, class_lineno) + if isinstance(stmt, ast.ClassDef) and stmt.name == "Config": + return _config_class_frozen(stmt) + return None + + +def _class_frozen_override( + cls: ast.ClassDef, + module_assignments: Mapping[str, tuple[tuple[int, ast.expr], ...]], + class_lineno: int, +) -> bool | None: + keyword_flags: Final = tuple(_bool_constant(keyword.value) for keyword in cls.keywords if keyword.arg == "frozen") + if keyword_flags: + return keyword_flags[-1] + body_flags: Final = tuple(_stmt_frozen_flag(stmt, module_assignments, class_lineno) for stmt in cls.body) + return next((flag for flag in reversed(body_flags) if flag is not None), None) + + +def iter_pydantic_violations(path: Path, tree: ast.AST) -> Iterator[Violation]: + module_assignments: Final = MappingProxyType( + { + name: tuple(binding for _, binding in assignments) + for name, assignments in groupby( + sorted(_module_assignment_items(tree), key=lambda item: item[0]), + key=lambda item: item[0], + ) + } + ) + models: Final = _pydantic_classes(tree) + bases_of: Final = {cls: _base_names(cls) for cls in models} + override_of: Final = {cls: _class_frozen_override(cls, module_assignments, cls.lineno) for cls in models} + + def frozen(known: frozenset[str]) -> frozenset[str]: + grown: Final = known | frozenset( + cls.name + for cls in models + if override_of[cls] is True or (override_of[cls] is None and bases_of[cls] & known) + ) + return grown if grown == known else frozen(grown) + + frozen_names: Final = frozen(frozenset()) + for cls in models: + if override_of[cls] is True or (override_of[cls] is None and bases_of[cls] & frozen_names): + continue + yield Violation( + path, + cls.lineno, + "LIT015", + f"pydantic model `{cls.name}` is not frozen: set `frozen=True` in " + f"`model_config`, an inner `class Config`, or the class keywords, and " + f"replace in-place field writes with `model_copy(update=...)` " + f"(suppress: `# frozen-ok: `)", + ) + + # --------------------------------------------------------------------------- # # Stacked comprehension clauses (LIT014) # --------------------------------------------------------------------------- # @@ -966,6 +1134,7 @@ def check_file(path: Path) -> tuple[Violation, ...]: *iter_final_violations(path, tree), *iter_param_violations(path, tree), *iter_typeddict_violations(path, tree), + *iter_pydantic_violations(path, tree), *(v for v, owned in comprehension_violations if owned), ), suppressions, @@ -997,7 +1166,7 @@ def _worker_count(path_count: int) -> int: def scan_paths(paths: Sequence[Path]) -> tuple[Violation, ...]: """check_file over every path. Pure per-file work, so it fans out across processes; callers sort, which is what keeps output order stable.""" - workers = _worker_count(len(paths)) + workers: Final = _worker_count(len(paths)) if workers == 1: return tuple(v for path in paths for v in check_file(path)) with Pool(workers) as pool: @@ -1005,13 +1174,13 @@ def scan_paths(paths: Sequence[Path]) -> tuple[Violation, ...]: def main(argv: Sequence[str]) -> int: - paths = tuple(a for a in argv if not a.startswith("-")) + paths: Final = tuple(a for a in argv if not a.startswith("-")) if not paths: print("usage: check_type_discipline.py ...", file=sys.stderr) return 2 - targets = tuple(collect_paths(paths)) - violations = sorted(scan_paths(targets)) + targets: Final = tuple(collect_paths(paths)) + violations: Final = sorted(scan_paths(targets)) for v in violations: print(v.render()) diff --git a/scripts/gate_slot_lock.py b/scripts/gate_slot_lock.py index 999b28ced1c..c1ab00461e0 100644 --- a/scripts/gate_slot_lock.py +++ b/scripts/gate_slot_lock.py @@ -1,7 +1,7 @@ #!/usr/bin/env python3 """Machine-wide slot lock for this repo's heavy entrypoints. -`make check`, `make lint`, and the standalone budget gates +`make check`, `make lint`, and the standalone lint gates (scripts/ruff_strict_gate.py, scripts/type_discipline_gate.py, scripts/type_check_gate.py) each hold one of N machine-wide slots while they run, so however many sessions and worktrees share one machine, at most N of diff --git a/scripts/lint_base_counts.py b/scripts/lint_base_counts.py new file mode 100644 index 00000000000..2b86d4ce55b --- /dev/null +++ b/scripts/lint_base_counts.py @@ -0,0 +1,327 @@ +#!/usr/bin/env python3 +"""Merge-base counts for the delta-vs-base lint gates. + +Each gate (scripts/ruff_strict_gate.py, scripts/type_discipline_gate.py, +scripts/type_check_gate.py, scripts/test_quality_gate.py) counts its rules +across the whole tree at HEAD and at the merge-base with the branch the change +merges into, and fails only when a rule grew past its ceiling: the merge-base +count, or the rule's fixed codebase-wide cap when the gate sets one and it is +higher. There is no committed budget, so a ceiling moves only when the base +branch does or when someone lowers a cap on it. + +The merge-base counts come from, in order, the disk cache under the git common +dir, the CI artifact publish-lint-base-counts.yml uploads for every push to the +default branch, and a scan of the base tree in a temporary worktree. Every +entry is keyed by the merge-base commit plus the checker's fingerprints (its +config, its rule logic, its tool version), so counts measured under a different +rule set are never matched, only recomputed. +""" + +from __future__ import annotations + +import hashlib +import io +import json +import os +import re +import subprocess +import sys +import zipfile +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from pathlib import Path +from types import MappingProxyType +from typing import Final, NamedTuple, TypeAlias + +REPO_ROOT: Final = Path(__file__).resolve().parents[1] +CACHE_DIR_NAME: Final = "litellm-lint-cache" +CACHE_KEEP_ENTRIES: Final = 8 +GH_TIMEOUT_SECONDS: Final = 10 + +_ORIGIN_SLUG: Final = re.compile(r"(?:git@github\.com:|https://github\.com/)([^/]+/[^/]+?)(?:\.git)?/?") + +Counts: TypeAlias = Mapping[str, int] +GhOutput: TypeAlias = Callable[[Sequence[str]], bytes | None] +NO_CAPS: Final[Counts] = MappingProxyType({}) + + +class Breach(NamedTuple): + rule: str + total: int + ceiling: int + added: int + + +@dataclass(frozen=True, slots=True) +class Checker: + name: str + fingerprints: tuple[str, ...] + + def key(self, base_point: str) -> str: + return cache_key(base_point, self.fingerprints) + + def artifact_name(self, base_point: str) -> str: + return f"{self.name}-counts-{self.key(base_point)}" + + def cache_file_name(self, base_point: str) -> str: + return f"{self.name}-base-{self.key(base_point)}.json" + + def cache_glob(self) -> str: + return f"{self.name}-base-*.json" + + +Fetch: TypeAlias = Callable[[Checker, str], Counts | None] + + +def sha256_of(path: Path) -> str: + return hashlib.sha256(path.read_bytes()).hexdigest() + + +def cache_key(base_point: str, fingerprints: Sequence[str]) -> str: + return hashlib.sha256("|".join((base_point, *fingerprints)).encode()).hexdigest()[:16] + + +Git: TypeAlias = Callable[[Sequence[str]], str] + + +def git_in(cwd: Path) -> Git: + def run(args: Sequence[str]) -> str: + proc: Final = subprocess.run(["git", *args], cwd=cwd, capture_output=True, text=True) + if proc.returncode not in (0, 1): + sys.stderr.write(proc.stderr) + raise SystemExit(f"git exited {proc.returncode}") + return proc.stdout + + return run + + +REPO_GIT: Final = git_in(REPO_ROOT) + + +def head_sha(git: Git = REPO_GIT) -> str: + return git(["rev-parse", "HEAD"]).strip() + + +def resolve_base_point(base_ref: str, git: Git = REPO_GIT) -> str: + """The snapshot commit base counts are measured at: merge-base(base_ref, HEAD), + made aware of an in-progress merge. Mid-merge, HEAD is still the pre-merge tip, + so its merge-base is the old branch point and every violation the base gained + since then would be blamed on this change. While MERGE_HEAD exists, prefer + merge-base(base_ref, MERGE_HEAD) whenever it is the newer of the two.""" + head_point: Final = git(["merge-base", base_ref, "HEAD"]).strip() + if not head_point: + return base_ref + merge_head: Final = git(["rev-parse", "--verify", "--quiet", "MERGE_HEAD"]).strip() + if not merge_head: + return head_point + merge_point: Final = git(["merge-base", base_ref, merge_head]).strip() + if not merge_point: + return head_point + older: Final = git(["merge-base", head_point, merge_point]).strip() + return merge_point if older == head_point else head_point + + +def default_cache_dir(git: Git = REPO_GIT) -> Path: + return Path(git(["rev-parse", "--path-format=absolute", "--git-common-dir"]).strip()) / CACHE_DIR_NAME + + +def validated_counts(data: object) -> Counts | None: + counts: Final = data.get("counts") if isinstance(data, dict) else None + if not isinstance(counts, dict): + return None + if not all( + isinstance(code, str) and isinstance(total, int) and not isinstance(total, bool) + for code, total in counts.items() + ): + return None + return counts + + +def load_cached_counts(path: Path) -> Counts | None: + try: + data: Final = json.loads(path.read_text()) + except (OSError, json.JSONDecodeError): + return None + return validated_counts(data) + + +def scratch_path(path: Path) -> Path: + return path.with_name(f".{path.name}.{os.getpid()}.tmp") + + +def counts_payload(base_point: str, counts: Counts) -> str: + return json.dumps({"base_point": base_point, "counts": dict(sorted(counts.items()))}, indent=2) + "\n" + + +def entry_recency(path: Path) -> float: + try: + return path.stat().st_mtime + except OSError: + return 0.0 + + +def evicted_beyond_cap(entries: Sequence[Path], keep: int) -> tuple[Path, ...]: + newest_first: Final = sorted(entries, key=entry_recency, reverse=True) + return tuple(newest_first[keep:]) + + +def store_counts(directory: Path, checker: Checker, base_point: str, counts: Counts) -> Path: + directory.mkdir(parents=True, exist_ok=True) + path: Final = directory / checker.cache_file_name(base_point) + scratch: Final = scratch_path(path) + scratch.write_text(counts_payload(base_point, counts)) + scratch.replace(path) + siblings: Final = tuple(entry for entry in directory.glob(checker.cache_glob()) if entry != path) + for stale in evicted_beyond_cap(siblings, CACHE_KEEP_ENTRIES - 1): + stale.unlink(missing_ok=True) + return path + + +def parse_origin_slug(url: str) -> str | None: + match: Final = _ORIGIN_SLUG.fullmatch(url.strip()) + return match.group(1) if match else None + + +def origin_slug(cwd: Path = REPO_ROOT) -> str | None: + proc: Final = subprocess.run(["git", "remote", "get-url", "origin"], cwd=cwd, capture_output=True, text=True) + return parse_origin_slug(proc.stdout) if proc.returncode == 0 else None + + +def gh_output(args: Sequence[str]) -> bytes | None: + try: + proc: Final = subprocess.run(["gh", *args], capture_output=True, timeout=GH_TIMEOUT_SECONDS) + except (OSError, subprocess.SubprocessError): + return None + return proc.stdout if proc.returncode == 0 else None + + +def _parsed_json(raw: bytes) -> object | None: + try: + return json.loads(raw) + except ValueError: + return None + + +def _artifact_download_url(listing: object) -> str | None: + artifacts: Final = listing.get("artifacts") if isinstance(listing, dict) else None + if not isinstance(artifacts, list) or not artifacts: + return None + newest: Final = artifacts[0] + if not isinstance(newest, dict) or newest.get("expired"): + return None + url: Final = newest.get("archive_download_url") + return url if isinstance(url, str) else None + + +def _counts_json_from_zip(zip_bytes: bytes) -> object | None: + try: + with zipfile.ZipFile(io.BytesIO(zip_bytes)) as archive: + members: Final = tuple(name for name in archive.namelist() if name.endswith(".json")) + if len(members) != 1: + return None + return json.loads(archive.read(members[0])) + except (zipfile.BadZipFile, ValueError, OSError): + return None + + +def counts_for_base(payload: object, base_point: str) -> Counts | None: + if not isinstance(payload, dict) or payload.get("base_point") != base_point: + return None + counts: Final = validated_counts(payload) + return counts if counts else None + + +def _fetch_fallback(reason: str) -> None: + sys.stderr.write(f"{reason}; computing base counts locally\n") + + +def fetch_ci_base_counts( + checker: Checker, + base_point: str, + gh: GhOutput = gh_output, + cwd: Path = REPO_ROOT, +) -> Counts | None: + """Base counts from the CI artifact published for `base_point`, or None. + + Every failure mode (no gh, no auth, offline, expired or missing artifact, + malformed payload, counts for a different commit) returns None so the + caller falls back to the local base scan; the fetch is an optimization and + must never make the gate less available than local compute alone.""" + slug: Final = origin_slug(cwd) + if slug is None: + return _fetch_fallback("origin remote is not a github.com URL") + name: Final = checker.artifact_name(base_point) + listing: Final = gh(["api", f"repos/{slug}/actions/artifacts?name={name}&per_page=1"]) + if listing is None: + return _fetch_fallback(f"could not list CI artifacts named {name}") + url: Final = _artifact_download_url(_parsed_json(listing)) + if url is None: + return _fetch_fallback(f"no usable CI artifact named {name}") + zip_bytes: Final = gh(["api", url]) + if zip_bytes is None: + return _fetch_fallback(f"download failed for CI artifact {name}") + counts: Final = counts_for_base(_counts_json_from_zip(zip_bytes), base_point) + if counts is None: + return _fetch_fallback(f"CI artifact {name} is not valid base counts for {base_point[:12]}") + sys.stderr.write(f"base counts fetched from CI artifact {name}\n") + return counts + + +def base_counts_cached( + checker: Checker, + base_point: str, + compute: Callable[[str], Counts], + cache_dir: Path | None = None, + fetch: Fetch = fetch_ci_base_counts, +) -> Counts: + """`compute` memoized on disk. The base tree at a given commit is immutable, + so its counts are a pure function of the merge-base plus the checker's + fingerprints in the cache key; an empty result is never stored because it is + the signature of a crashed pass, not a clean tree. On a disk miss the counts + CI already published for the merge-base are fetched before the expensive + local base scan; a fetch miss of any kind computes locally.""" + directory: Final = default_cache_dir() if cache_dir is None else cache_dir + cached: Final = load_cached_counts(directory / checker.cache_file_name(base_point)) + if cached is not None: + return cached + fetched: Final = fetch(checker, base_point) + if fetched: + store_counts(directory, checker, base_point, fetched) + return fetched + counts: Final = compute(base_point) + if counts: + store_counts(directory, checker, base_point, counts) + return counts + + +def emit_counts(checker: Checker, counts: Counts, directory: Path, head_point: str) -> Path: + """Write HEAD's per-rule counts as the file the publisher workflow uploads. + + The filename stem is exactly the artifact name `fetch_ci_base_counts` will + later look up for this commit, so emit and fetch cannot drift apart. Empty + counts are refused: a pass that produced nothing almost certainly crashed, + and publishing it would poison every branch that fetches it.""" + if not counts: + print( + f"FAIL: {checker.name} produced no violations; refusing to publish empty base " + "counts because the pass almost certainly crashed or emitted nothing." + ) + raise SystemExit(1) + name: Final = checker.artifact_name(head_point) + directory.mkdir(parents=True, exist_ok=True) + path: Final = directory / f"{name}.json" + path.write_text(counts_payload(head_point, counts)) + print(f"Emitted base counts for {head_point} as {name}.json ({sum(counts.values())} violations total)") + return path + + +def ceiling(rule: str, base: Counts, caps: Counts) -> int: + return max(base.get(rule, 0), caps.get(rule, 0)) + + +def evaluate(head: Counts, base: Counts, caps: Counts = NO_CAPS) -> tuple[Breach, ...]: + return tuple( + Breach(rule, total, ceiling(rule, base, caps), total - base.get(rule, 0)) + for rule, total in sorted(head.items()) + if total > ceiling(rule, base, caps) + ) diff --git a/scripts/pre_commit_lint.sh b/scripts/pre_commit_lint.sh index bc5341d265c..1c793705dc6 100755 --- a/scripts/pre_commit_lint.sh +++ b/scripts/pre_commit_lint.sh @@ -13,10 +13,10 @@ # - tests/e2e and tests/e2e_harness Python # -> `make lint-e2e-basedpyright` (test-linting.yml's e2e type-check step) # + raw HTTP client ban (test-code-quality.yml's check_e2e_no_raw_requests) -# - tests/ Python, ruff-tests.toml, test-quality-budget.json, scripts/check_test_quality.py, +# - tests/ Python, ruff-tests.toml, scripts/check_test_quality.py, # scripts/test_quality_gate.py # -> ruff over ruff-tests.toml + `make lint-test-quality` (test-linting.yml's -# test-tree ruff and test-quality budget steps) +# test-tree ruff and test-quality gate steps) # - dashboard -> prettier + eslint + lint budgets (test-litellm-ui-build.yml's frontend-lint) # - proxy/types -> regenerate the lazy OpenAPI snapshot and dashboard API types, fail on drift (check-ui-api-types.yml) # @@ -33,7 +33,7 @@ set -eu # before anything else, so N parallel `make check` runs across worktrees execute two # at a time instead of thrashing the machine. The wrapper exports # LITELLM_GATE_SLOT_HELD, so this re-exec happens exactly once and everything this -# script spawns (make lint, the budget gates) skips its own acquisition. +# script spawns (make lint, the lint gates) skips its own acquisition. script_dir=$(python3 -c 'import os, sys; print(os.path.dirname(os.path.realpath(sys.argv[1])))' "$0") if [ -z "${LITELLM_GATE_SLOT_HELD:-}" ]; then exec python3 "$script_dir/gate_slot_lock.py" "$0" "$@" @@ -97,7 +97,7 @@ existing_files() { litellm_py_pattern='^litellm/.*\.py$' e2e_py_pattern='^tests/e2e(_harness)?/.*\.py$' -test_tree_pattern='^(tests/.*\.py|ruff-tests\.toml|test-quality-budget\.json|scripts/(check_test_quality|test_quality_gate)\.py)$' +test_tree_pattern='^(tests/.*\.py|ruff-tests\.toml|scripts/(check_test_quality|test_quality_gate)\.py)$' spec_pattern='^(litellm/(proxy|types)/.*|ui/litellm-dashboard/(scripts/gen-api-types\.mjs|package\.json|package-lock\.json|src/lib/http/schema\.d\.ts))$' ui_prettier_pattern='^ui/litellm-dashboard/.*\.(js|jsx|ts|tsx|mjs|cjs|json|css|scss|md|mdx|yml|yaml|html)$' ui_eslint_pattern='^ui/litellm-dashboard/.*\.(js|jsx|ts|tsx|mjs|cjs)$' @@ -147,7 +147,7 @@ if [ -n "$staged" ]; then } warn_skipped "Python lint (make lint)" "$litellm_py_pattern" "$litellm_py_files" warn_skipped "tests/e2e checks (basedpyright + raw HTTP client ban)" "$e2e_py_pattern" "$e2e_py_files" - warn_skipped "test-tree lint (ruff-tests.toml + test-quality budget)" "$test_tree_pattern" "$test_tree_files" + warn_skipped "test-tree lint (ruff-tests.toml + test-quality gate)" "$test_tree_pattern" "$test_tree_files" warn_skipped "dashboard lint (prettier + eslint + lint budgets)" "$ui_prettier_pattern" "$ui_prettier_changed" warn_skipped "dashboard API-type sync (npm run gen:api)" "$spec_pattern" "$spec_files" fi @@ -304,9 +304,9 @@ if [ -n "$test_tree_files" ] && [ -z "$litellm_py_files" ]; then echo "check: linting the test tree (ruff check --config ruff-tests.toml tests)" uv run --no-sync ruff check --config ruff-tests.toml tests \ || { echo "✗ Test-tree ruff failed. Fix the errors above, then re-run make check." >&2; status=1; } - echo "check: checking the test-quality budget (make lint-test-quality)" + echo "check: checking the test-quality gate (make lint-test-quality)" make lint-test-quality \ - || { echo "✗ Test-quality budget failed. Fix the errors above, then re-run make check." >&2; status=1; } + || { echo "✗ Test-quality gate failed. Fix the errors above, then re-run make check." >&2; status=1; } fi if [ -n "${python_pid:-}" ]; then @@ -334,7 +334,7 @@ summary_item() { echo "check: summary" summary_item "Python lint (make lint)" "$litellm_py_files" "no litellm/ Python files in scope" summary_item "tests/e2e checks (basedpyright + raw HTTP client ban)" "$e2e_py_files" "no tests/e2e Python files in scope" -summary_item "test-tree lint (ruff-tests.toml + test-quality budget)" "$test_tree_files" \ +summary_item "test-tree lint (ruff-tests.toml + test-quality gate)" "$test_tree_files" \ "no tests/ Python files or test-tree lint inputs in scope" summary_item "dashboard lint (prettier + eslint + lint budgets)" "$ui_prettier_changed$ui_eslint_changed" "no dashboard files in scope" summary_item "dashboard API-type sync (npm run gen:api)" "$spec_files" "no litellm/proxy, litellm/types, or generator files in scope" diff --git a/scripts/ruff_strict_gate.py b/scripts/ruff_strict_gate.py index 8da10dd76f0..61ae0236b53 100644 --- a/scripts/ruff_strict_gate.py +++ b/scripts/ruff_strict_gate.py @@ -1,12 +1,17 @@ #!/usr/bin/env python3 -"""Total-count gate for the strict ruff rules in ruff-strict.toml. +"""Delta-vs-base gate for the strict ruff rules in ruff-strict.toml. -Each rule has a hard ``limit`` in ruff-strict-budget.json. The gate counts each -rule across the whole tree and fails when a rule is both over its limit and -higher than the base it merges into, so a change is blamed for the violations it -adds, never for drift that already exists in the base. ``--update`` ratchets each -rule's limit down by the number of violations this branch fixed relative to its -branch point (the merge-base). +Each rule is counted across the whole tree at HEAD and at the merge-base with +the branch this change merges into, and the gate fails only when a rule grew +past the merge-base count, so a change is blamed for the violations it adds, +never for drift that already sits in the base. There is no committed budget: the merge-base count is the +ceiling, so it moves only when the base branch does. + +The merge-base counts come from scripts/lint_base_counts.py: the disk cache, +then the CI artifact published for that commit, then a ruff pass over a +detached worktree at the merge-base under the current ruff configs. +``--emit-counts-dir`` writes HEAD's counts as the file that artifact is built +from. """ import argparse @@ -17,15 +22,26 @@ import subprocess import sys import tempfile from collections import Counter +from collections.abc import Callable, Mapping, Sequence from pathlib import Path from typing import Final, NamedTuple -REPO_ROOT = Path(__file__).resolve().parent.parent -STRICT_CONFIG = REPO_ROOT / "ruff-strict.toml" -BUDGET_PATH = REPO_ROOT / "ruff-strict-budget.json" -TARGET = "litellm" +from lint_base_counts import ( + Checker, + base_counts_cached, + emit_counts, + evaluate, + head_sha, + resolve_base_point, + sha256_of, +) -_HUNK = re.compile(r"^@@ -\d+(?:,\d+)? \+(\d+)(?:,(\d+))? @@") +REPO_ROOT: Final = Path(__file__).resolve().parent.parent +STRICT_CONFIG: Final = REPO_ROOT / "ruff-strict.toml" +BASE_CONFIG: Final = REPO_ROOT / "ruff.toml" +TARGET: Final = "litellm" + +_HUNK: Final = re.compile(r"^@@ -\d+(?:,\d+)? \+(\d+)(?:,(\d+))? @@") class Violation(NamedTuple): @@ -34,38 +50,22 @@ class Violation(NamedTuple): code: str -class Breach(NamedTuple): - rule: str - total: int - cap: int - added: int - - -def _run(cmd: list, cwd: Path = REPO_ROOT) -> str: - proc = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True) +def _run(cmd: Sequence[str], cwd: Path = REPO_ROOT) -> str: + proc: Final = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True) if proc.returncode not in (0, 1): sys.stderr.write(proc.stderr) raise SystemExit(f"{cmd[0]} exited {proc.returncode}") return proc.stdout -def resolve_base_point(base_ref: str, cwd: Path = REPO_ROOT) -> str: - """The snapshot commit base counts are measured at: merge-base(base_ref, HEAD), - made aware of an in-progress merge. Mid-merge, HEAD is still the pre-merge tip, - so its merge-base is the old branch point and every violation the base gained - since then would be blamed on this change. While MERGE_HEAD exists, prefer - merge-base(base_ref, MERGE_HEAD) whenever it is the newer of the two.""" - head_point: Final = _run(["git", "merge-base", base_ref, "HEAD"], cwd=cwd).strip() - if not head_point: - return base_ref - merge_head: Final = _run(["git", "rev-parse", "--verify", "--quiet", "MERGE_HEAD"], cwd=cwd).strip() - if not merge_head: - return head_point - merge_point: Final = _run(["git", "merge-base", base_ref, merge_head], cwd=cwd).strip() - if not merge_point: - return head_point - older: Final = _run(["git", "merge-base", head_point, merge_point], cwd=cwd).strip() - return merge_point if older == head_point else head_point +def ruff_version() -> str: + return _run(["ruff", "--version"]).strip() + + +def checker_identity( + strict_config: Path = STRICT_CONFIG, base_config: Path = BASE_CONFIG, version: Callable[[], str] = ruff_version +) -> Checker: + return Checker("ruff-strict", (sha256_of(strict_config), sha256_of(base_config), version())) def _ruff_json(cwd: Path, config: Path) -> list: @@ -76,7 +76,7 @@ def _ruff_json(cwd: Path, config: Path) -> list: return json.loads(raw or "[]") -def head_violations() -> list: +def head_violations() -> list[Violation]: out = [] for item in _ruff_json(REPO_ROOT, STRICT_CONFIG): name = Path(item["filename"]) @@ -90,47 +90,26 @@ def head_violations() -> list: return out -def count_by_rule(violations: list) -> dict: +def count_by_rule(violations: Sequence[Violation]) -> dict[str, int]: return dict(Counter(v.code for v in violations)) -def base_counts(ref: str) -> dict: - parent = Path(tempfile.mkdtemp(prefix="ruff_base_")) - worktree = parent / "wt" +def base_counts(ref: str) -> dict[str, int]: + parent: Final = Path(tempfile.mkdtemp(prefix="ruff_base_")) + worktree: Final = parent / "wt" try: _run(["git", "worktree", "add", "--detach", str(worktree), ref]) - shutil.copy(STRICT_CONFIG, worktree / "ruff-strict.toml") - items = _ruff_json(worktree, worktree / "ruff-strict.toml") + shutil.copy(BASE_CONFIG, worktree / BASE_CONFIG.name) + shutil.copy(STRICT_CONFIG, worktree / STRICT_CONFIG.name) + items: Final = _ruff_json(worktree, worktree / STRICT_CONFIG.name) return dict(Counter(item["code"] for item in items)) finally: _run(["git", "worktree", "remove", "--force", str(worktree)]) shutil.rmtree(parent, ignore_errors=True) -def over_ceiling(head: dict, budget: dict) -> frozenset: - """Rules whose head count already exceeds their limit. - - A rule can only breach when it is over its limit, so when none are the base - comparison cannot change the verdict and the base worktree scan can be skipped. - """ - return frozenset( - rule for rule, spec in budget.items() - if head.get(rule, 0) > spec["limit"] - ) - - -def evaluate(head: dict, base: dict, budget: dict) -> list: - breaches = [] - for rule, spec in budget.items(): - cap = spec["limit"] - total = head.get(rule, 0) - if total > cap and total > base.get(rule, 0): - breaches.append(Breach(rule, total, cap, total - base.get(rule, 0))) - return sorted(breaches) - - -def parse_changed_lines(diff_text: str) -> dict: - changed: dict = {} +def parse_changed_lines(diff_text: str) -> dict[str, set[int]]: + changed: dict[str, set[int]] = {} path = None for line in diff_text.splitlines(): if line.startswith("+++ b/"): @@ -142,84 +121,49 @@ def parse_changed_lines(diff_text: str) -> dict: return changed -def introduced(violations: list, changed: dict) -> list: +def introduced(violations: Sequence[Violation], changed: Mapping[str, set[int]]) -> list[Violation]: return [v for v in violations if v.line in changed.get(v.file, set())] def cmd_check(base: str) -> None: - budget = json.loads(BUDGET_PATH.read_text()) - head = head_violations() - head_counts = count_by_rule(head) - if not over_ceiling(head_counts, budget): - print(f"OK: every strict rule is within its codebase ceiling (base {base})") - return - base_point = resolve_base_point(base) - breaches = evaluate(head_counts, base_counts(base_point), budget) + head: Final = head_violations() + base_point: Final = resolve_base_point(base) + breaches: Final = evaluate(count_by_rule(head), base_counts_cached(checker_identity(), base_point, base_counts)) if not breaches: - print(f"OK: every strict rule is within its codebase ceiling (base {base})") + print(f"OK: no strict rule grew past its merge-base count (base {base})") return - new = introduced( - head, - parse_changed_lines( - _run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET]) - ), - ) - print(f"FAIL: strict-rule totals exceed their limit (base {base}):") + diff: Final = _run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET]) + new: Final = introduced(head, parse_changed_lines(diff)) + print(f"FAIL: strict-rule totals grew past their merge-base count (base {base}):") for breach in breaches: - print( - f" {breach.rule}: total {breach.total} over limit {breach.cap} (this change added {breach.added})" - ) + print(f" {breach.rule}: total {breach.total} over ceiling {breach.ceiling} (this change added {breach.added})") for violation in sorted(v for v in new if v.code == breach.rule): print(f" {violation.file}:{violation.line}") print( - "Reduce the new violations or remove an equal number elsewhere; the ceiling is the limit in ruff-strict-budget.json." + "Reduce the new violations or remove an equal number elsewhere; the ceiling is the merge-base count." ) raise SystemExit(1) -def ratcheted_budget(budget: dict, current: dict, base: dict) -> dict: - """Each rule's limit lowered by the violations `current` fixed vs `base`. - - `base` is the count at the branch point (the commit this branch diverged - from). The drop is clamped to what was actually cleared (a rule that grew - stays put), so the limit only ever falls. - """ - return { - rule: { - "limit": max(0, spec["limit"] - max(0, base.get(rule, 0) - current.get(rule, 0))) - } - for rule, spec in sorted(budget.items()) - } - - -def cmd_update(base_ref: str) -> None: - """Ratchet each rule's limit down by the violations this branch fixed. - - The working-tree count is compared against a ruff pass over a detached - worktree at the branch point (the merge-base with `base_ref`), so a branch's - fixes tighten its own ceilings by exactly what they cleared since it diverged. - """ - budget = json.loads(BUDGET_PATH.read_text()) - base_point = resolve_base_point(base_ref) - updated = ratcheted_budget( - budget, count_by_rule(head_violations()), base_counts(base_point) - ) - BUDGET_PATH.write_text(json.dumps(updated, indent=2, sort_keys=True) + "\n") - cleared = sum(budget[rule]["limit"] - updated[rule]["limit"] for rule in updated) - print(f"Ratcheted strict-rule limits down by {cleared} violations this branch fixed") - - def main() -> None: - parser = argparse.ArgumentParser(description=__doc__) + parser: Final = argparse.ArgumentParser(description=__doc__) parser.add_argument("--base", help="Comparison ref (default: origin's current default branch)") - parser.add_argument("--update", action="store_true") - args = parser.parse_args() + parser.add_argument( + "--emit-counts-dir", + type=Path, + help="Write HEAD's per-rule counts to this directory as a base-counts artifact instead of gating", + ) + args: Final = parser.parse_args() from default_branch import resolve_base_ref from gate_slot_lock import held_slot + if args.emit_counts_dir is not None: + with held_slot(): + emit_counts(checker_identity(), count_by_rule(head_violations()), args.emit_counts_dir, head_sha()) + return base_ref: Final = resolve_base_ref(args.base, REPO_ROOT) with held_slot(): - cmd_update(base_ref) if args.update else cmd_check(base_ref) + cmd_check(base_ref) if __name__ == "__main__": diff --git a/scripts/test_quality_gate.py b/scripts/test_quality_gate.py index e486324c741..a9ad60f088c 100644 --- a/scripts/test_quality_gate.py +++ b/scripts/test_quality_gate.py @@ -1,32 +1,24 @@ #!/usr/bin/env python3 -"""Total-count gate for the TQ* rules in scripts/check_test_quality.py. +"""Delta-vs-base gate for the TQ* rules in scripts/check_test_quality.py. Sibling of scripts/type_discipline_gate.py, pointed at the test tree instead of -the package. Each rule listed in test-quality-budget.json has a hard ``limit``. -The gate counts each rule across the whole `tests` tree and fails when a rule is -both over its limit and higher than the base it merges into, so a change is -blamed for the violations it adds, never for drift that already exists in the -base. +the package. Each rule is counted across the whole `tests` tree at HEAD and at +the merge-base with the branch this change merges into, and the gate fails only +when a rule grew past the merge-base count, so a change is blamed for the +violations it adds, never for drift that already exists in the base. There is +no committed budget: the merge-base count is the ceiling, so it moves only when +the base branch does. -Every rule is seeded at exactly its count on the day the gate landed, so the -suite's existing debt is grandfathered and any net-new violation trips the gate -immediately. ``--update`` ratchets a limit down by the violations fixed relative -to ``--base``, so the ceilings only ever fall. Base counts are measured with the -*current* checker, so a rule introduced on this branch is counted at the base too -and ratchets like every other one. The ratchet runs as a scheduled automation -against the repository's default branch, not on PR branches, so concurrent PRs never -race to edit the same limit. - -The deliberate difference from its sibling: this gate has no headroom anywhere. -Type discipline seeded LIT010/LIT011 at 1.5x to leave room for an in-flight -sweep; a test-quality violation has no such transition to absorb, so the line is -today's count and the only legal direction is down. +The merge-base counts come from scripts/lint_base_counts.py: the disk cache, +then the CI artifact published for that commit, then a pass of the current +checker over a detached worktree at the merge-base, so a rule introduced on +this branch is counted at the base too. ``--emit-counts-dir`` writes HEAD's +counts as the file that artifact is built from. """ from __future__ import annotations import argparse -import json import re import shutil import signal @@ -39,9 +31,18 @@ from pathlib import Path from types import FrameType, MappingProxyType from typing import Final, NamedTuple +from lint_base_counts import ( + Checker, + base_counts_cached, + emit_counts, + evaluate, + head_sha, + resolve_base_point, + sha256_of, +) + REPO_ROOT: Final = Path(__file__).resolve().parent.parent CHECKER: Final = REPO_ROOT / "scripts" / "check_test_quality.py" -BUDGET_PATH: Final = REPO_ROOT / "test-quality-budget.json" TARGET: Final = "tests" TERMINATION_SIGNALS: Final = (signal.SIGTERM, signal.SIGHUP) @@ -56,13 +57,6 @@ class Violation(NamedTuple): code: str -class Breach(NamedTuple): - rule: str - total: int - cap: int - added: int - - def _run(cmd: Sequence[str], cwd: Path = REPO_ROOT) -> str: proc: Final = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True) if proc.returncode not in (0, 1): @@ -71,22 +65,8 @@ def _run(cmd: Sequence[str], cwd: Path = REPO_ROOT) -> str: return proc.stdout -def resolve_base_point(base_ref: str, cwd: Path = REPO_ROOT) -> str: - """The snapshot commit base counts are measured at: merge-base(base_ref, HEAD), - made aware of an in-progress merge. Mid-merge, HEAD is still the pre-merge tip, - so its merge-base is the old branch point and every violation the base gained - since then would be blamed on this change.""" - head_point: Final = _run(["git", "merge-base", base_ref, "HEAD"], cwd=cwd).strip() - if not head_point: - return base_ref - merge_head: Final = _run(["git", "rev-parse", "--verify", "--quiet", "MERGE_HEAD"], cwd=cwd).strip() - if not merge_head: - return head_point - merge_point: Final = _run(["git", "merge-base", base_ref, merge_head], cwd=cwd).strip() - if not merge_point: - return head_point - older: Final = _run(["git", "merge-base", head_point, merge_point], cwd=cwd).strip() - return merge_point if older == head_point else head_point +def checker_identity(checker: Path = CHECKER) -> Checker: + return Checker("test-quality", (sha256_of(checker),)) def _check(root: Path, checker: Path) -> tuple[Violation, ...]: @@ -143,26 +123,6 @@ def base_counts(ref: str, repo_root: Path = REPO_ROOT, checker: Path = CHECKER) shutil.rmtree(parent, ignore_errors=True) -def over_ceiling(head: Mapping[str, int], budget: Mapping[str, Mapping[str, int]]) -> frozenset[str]: - """Rules whose head count already exceeds their limit. When none are, the base - comparison cannot change the verdict and the base worktree scan is skipped.""" - return frozenset( - rule for rule, spec in budget.items() if head.get(rule, 0) > spec["limit"] - ) - - -def evaluate( - head: Mapping[str, int], - base: Mapping[str, int], - budget: Mapping[str, Mapping[str, int]], -) -> tuple[Breach, ...]: - return tuple(sorted( - Breach(rule, head.get(rule, 0), spec["limit"], head.get(rule, 0) - base.get(rule, 0)) - for rule, spec in budget.items() - if head.get(rule, 0) > spec["limit"] and head.get(rule, 0) > base.get(rule, 0) - )) - - def _hunk_lines(body: str) -> frozenset[int]: return frozenset( line @@ -191,92 +151,46 @@ def introduced( def cmd_check(base: str) -> None: - budget: Final = json.loads(BUDGET_PATH.read_text()) head: Final = head_violations() - head_counts: Final = count_by_rule(head) - if not over_ceiling(head_counts, budget): - print(f"OK: every TQ rule is within its test-suite ceiling (base {base})") - return base_point: Final = resolve_base_point(base) - base_at_point: Final = base_counts(base_point) - breaches: Final = evaluate(head_counts, base_at_point, budget) + breaches: Final = evaluate(count_by_rule(head), base_counts_cached(checker_identity(), base_point, base_counts)) if not breaches: - print(f"OK: every TQ rule is within its test-suite ceiling (base {base})") + print(f"OK: no TQ rule grew past its merge-base count (base {base})") return - new: Final = introduced( - head, - parse_changed_lines( - _run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET]) - ), - ) - print(f"FAIL: TQ-rule totals exceed their limit (base {base}):") + diff: Final = _run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET]) + new: Final = introduced(head, parse_changed_lines(diff)) + print(f"FAIL: TQ-rule totals grew past their merge-base count (base {base}):") for breach in breaches: - print( - f" {breach.rule}: total {breach.total} over limit {breach.cap} " - f"(this change added {breach.added})" - ) + print(f" {breach.rule}: total {breach.total} over ceiling {breach.ceiling} (this change added {breach.added})") for violation in sorted(v for v in new if v.code == breach.rule): print(f" {violation.file}:{violation.line}") print( - "Fix the new violations, or give each one a reason " - "(`# test-quality-ok: `), or remove an equal number elsewhere; " - "the ceiling is the limit in test-quality-budget.json. " - "Run `python scripts/check_test_quality.py tests/` to see every finding." + "Fix the new violations, or give each one a reason (`# test-quality-ok: `), or remove an " + "equal number elsewhere; the ceiling is the merge-base count. Run " + "`python scripts/check_test_quality.py tests/` to see every finding." ) raise SystemExit(1) -def ratcheted_budget( - budget: Mapping[str, Mapping[str, int]], - current: Mapping[str, int], - base: Mapping[str, int], -) -> Mapping[str, Mapping[str, int]]: - """Each rule's limit lowered by the violations `current` fixed vs `base`. The drop - is clamped to what was actually cleared, so a limit only ever falls.""" - return MappingProxyType({ - rule: {"limit": max(0, spec["limit"] - max(0, base.get(rule, 0) - current.get(rule, 0)))} - for rule, spec in sorted(budget.items()) - }) - - -def cmd_update(base_ref: str) -> None: - """Ratchet each rule's limit down by the violations this branch fixed.""" - budget: Final = json.loads(BUDGET_PATH.read_text()) - base_point: Final = resolve_base_point(base_ref) - updated: Final = ratcheted_budget( - budget, count_by_rule(head_violations()), base_counts(base_point) - ) - BUDGET_PATH.write_text(json.dumps(dict(updated), indent=2, sort_keys=True) + "\n") - cleared: Final = sum(budget[rule]["limit"] - updated[rule]["limit"] for rule in updated) - print(f"Ratcheted TQ-rule limits down by {cleared} violations this branch fixed") - - -def cmd_seed() -> None: - """Write the budget from the working tree's current counts. Used once, to land - the gate; afterwards `--update` is the only thing that may move a limit.""" - counts: Final = count_by_rule(head_violations()) - BUDGET_PATH.write_text( - json.dumps({rule: {"limit": counts[rule]} for rule in sorted(counts)}, indent=2) + "\n" - ) - print(f"Seeded {BUDGET_PATH.name} at " + ", ".join(f"{r}={counts[r]}" for r in sorted(counts))) - - def main() -> None: parser: Final = argparse.ArgumentParser(description=__doc__) parser.add_argument("--base", help="Comparison ref (default: origin's current default branch)") - parser.add_argument("--update", action="store_true") - parser.add_argument("--seed", action="store_true") + parser.add_argument( + "--emit-counts-dir", + type=Path, + help="Write HEAD's per-rule counts to this directory as a base-counts artifact instead of gating", + ) args: Final = parser.parse_args() from default_branch import resolve_base_ref from gate_slot_lock import held_slot + if args.emit_counts_dir is not None: + with held_slot(): + emit_counts(checker_identity(), count_by_rule(head_violations()), args.emit_counts_dir, head_sha()) + return + base_ref: Final = resolve_base_ref(args.base, REPO_ROOT) with held_slot(): - if args.seed: - cmd_seed() - elif args.update: - cmd_update(resolve_base_ref(args.base, REPO_ROOT)) - else: - cmd_check(resolve_base_ref(args.base, REPO_ROOT)) + cmd_check(base_ref) if __name__ == "__main__": diff --git a/scripts/type_check_gate.py b/scripts/type_check_gate.py index 023b35af8fe..2a9f1c1229f 100644 --- a/scripts/type_check_gate.py +++ b/scripts/type_check_gate.py @@ -2,15 +2,17 @@ """Delta-vs-base per-rule gate for basedpyright. basedpyright's ``--outputjson`` is reduced to a count of errors per *rule* -(``reportAny``, ``reportArgumentType``, ...) and checked against a committed -budget of the form ``{rule: {limit}}``, the same shape as -``ruff-strict-budget.json``. A rule fails only when its codebase-wide total is -both over its ``limit`` *and* higher than the count on the base it merges into, -so a change is blamed for the errors it adds, never for drift that already sits -in the base. That ``> base`` guard is what stops an unrelated PR from inheriting -a red once two PRs each land near the limit and their sum crosses it: the -bystander's count equals its base, so it is spared, while any PR that actually -grows the rule past its limit still fails. +(``reportAny``, ``reportArgumentType``, ...) at HEAD and at the merge-base with +the branch this change merges into. A rule fails only when its codebase-wide +total grew past the merge-base count, so a change is blamed for the errors it +adds, never for drift that already sits in the base, and an unrelated PR never +inherits a red from what landed next to it: its count equals its base. + +``reportAny`` and ``reportExplicitAny`` are the exception, because Any spreads: +a correct change can surface new ones far from the lines it touched. Their +ceiling is the larger of the merge-base count and a fixed codebase-wide cap in +ANY_CAPS, so a change may add some while the total stays under the cap, and the +total can never pass it. The cap moves only when someone lowers it on main. Installed packages are part of the measurement: a typed dependency that is present changes what basedpyright can prove (and therefore which diagnostics @@ -28,21 +30,16 @@ The gate runs basedpyright itself, for both the head and the base pass, with ``NODE_OPTIONS`` raised to the heap this repo needs: basedpyright's node process OOMs at the ~4 GB default, and when callers had to remember the flag, every hand-copied pipeline (Makefile, CI, a dev running the recipe by hand) -was one forgotten env line away from an 80-second crash. The base count only -matters once some rule is over its limit, so when none is the base pass is -skipped outright. When it is needed, it is a second basedpyright pass over a -detached worktree at the merge-base, run under the same environment so import -resolution matches, and its per-rule counts are cached under the repo's git +was one forgotten env line away from an 80-second crash. The base pass is a +second basedpyright run over a detached worktree at the merge-base, under the +same environment so import resolution matches, and scripts/lint_base_counts.py +spares it whenever it can: the per-rule counts are cached under the repo's git common dir keyed by merge-base commit, ``pyrightconfig.json``, ``uv.lock``, -the Prisma schema, and the dependency-group set, so re-runs against the same -branch point pay for it once. A CI workflow publishes every main commit's counts as -an artifact (``--emit-counts-dir`` is its entry point), and on a disk-cache miss -the gate first tries to download the merge-base's artifact through the ``gh`` -CLI; any fetch failure falls back silently to the local base pass, so the gate -never gets worse than it was without CI. ``--update`` ratchets each rule's ``limit`` down by the -number of errors this branch fixed relative to its branch point (the merge-base), -so the headroom you were granted shrinks by exactly what you cleared and never -grows. +the Prisma schema, and the dependency-group set, and on a disk-cache miss the +artifact publish-lint-base-counts.yml uploaded for the merge-base is +downloaded through the ``gh`` CLI (``--emit-counts-dir`` is the publisher's +entry point); any fetch failure falls back silently to the local base pass, so +the gate never gets worse than it was without CI. ``--outputjson`` is used rather than text diagnostics because the latter wrap across lines, leaving the ``(reportRule)`` on a continuation line away from the @@ -53,33 +50,36 @@ carries an unambiguous ``rule`` field. import argparse import contextlib import hashlib -import io import json import os -import re import shutil import subprocess import sys import tempfile -import zipfile from collections import Counter -from collections.abc import Callable, Iterator, Mapping, Sequence +from collections.abc import Callable, Iterator, Mapping from pathlib import Path -from typing import Final, NamedTuple +from types import MappingProxyType +from typing import Final + +from lint_base_counts import ( + Checker, + base_counts_cached, + emit_counts, + evaluate, + head_sha, + resolve_base_point, +) REPO_ROOT = Path(__file__).resolve().parent.parent -BUDGET_PATH = REPO_ROOT / "basedpyright-code-budget.json" PYRIGHT_CONFIG = REPO_ROOT / "pyrightconfig.json" UV_LOCK = REPO_ROOT / "uv.lock" -CACHE_FILE_PREFIX = "basedpyright-base-" -CACHE_KEEP_ENTRIES = 8 -ARTIFACT_NAME_PREFIX = "basedpyright-counts-" -GH_TIMEOUT_SECONDS = 10 # The one environment every basedpyright pass measures in. The group set is # the slim one the CI publisher has always installed (not bootstrap's fatter -# --extra proxy env), so the committed budgets stay valid; changing it re-keys -# every cache and artifact fingerprint, so stale counts can never be matched. +# --extra proxy env), so published and cached counts stay comparable; changing +# it re-keys every cache and artifact fingerprint, so stale counts can never be +# matched. TYPECHECK_ENV_DIR = REPO_ROOT / ".venv-typecheck" TYPECHECK_DEP_GROUPS = ("proxy-dev", "e2e-dev") PRISMA_GENERATE_SCRIPT = REPO_ROOT / "scripts" / "prisma_generate_if_needed.py" @@ -93,17 +93,7 @@ NODE_HEAP_OPTION = "--max-old-space-size=8192" # Bucket for a basedpyright diagnostic with no `rule`. Counted so it's gated. UNCODED = "" -# Limit for a rule that shows up at HEAD but isn't in the budget at all -- a -# brand-new error category (new construct, or a tool/version change). The rule -# fails once it clears this many errors. -DEFAULT_LIMIT = 10 - - -class Breach(NamedTuple): - code: str - total: int - cap: int - added: int +ANY_CAPS: Final[Mapping[str, int]] = MappingProxyType({"reportAny": 6150, "reportExplicitAny": 1440}) def _to_relative(raw: str, root: Path) -> str | None: @@ -238,25 +228,6 @@ def run_basedpyright(cwd: Path = REPO_ROOT, env_dir: Path = TYPECHECK_ENV_DIR) - return proc.stdout -def resolve_base_point(base_ref: str, cwd: Path = REPO_ROOT) -> str: - """The snapshot commit base counts are measured at: merge-base(base_ref, HEAD), - made aware of an in-progress merge. Mid-merge, HEAD is still the pre-merge tip, - so its merge-base is the old branch point and every violation the base gained - since then would be blamed on this change. While MERGE_HEAD exists, prefer - merge-base(base_ref, MERGE_HEAD) whenever it is the newer of the two.""" - head_point: Final = _run(["git", "merge-base", base_ref, "HEAD"], cwd=cwd).strip() - if not head_point: - return base_ref - merge_head: Final = _run(["git", "rev-parse", "--verify", "--quiet", "MERGE_HEAD"], cwd=cwd).strip() - if not merge_head: - return head_point - merge_point: Final = _run(["git", "merge-base", base_ref, merge_head], cwd=cwd).strip() - if not merge_point: - return head_point - older: Final = _run(["git", "merge-base", head_point, merge_point], cwd=cwd).strip() - return merge_point if older == head_point else head_point - - @contextlib.contextmanager def _temp_worktree(ref: str) -> Iterator[Path]: parent = Path(tempfile.mkdtemp(prefix="bpr_base_")) @@ -283,21 +254,6 @@ def base_counts(ref: str) -> dict[str, int]: return count_basedpyright(run_basedpyright(worktree), root=worktree) -def over_ceiling( - head: Mapping[str, int], budget: Mapping[str, Mapping[str, int]] -) -> frozenset[str]: - """Rules whose head count already exceeds their limit. - - A rule can only breach when it is over its limit, so when none are the base - comparison cannot change the verdict and the base worktree pass can be skipped. - """ - return frozenset( - code - for code, total in head.items() - if total > (budget[code]["limit"] if code in budget else DEFAULT_LIMIT) - ) - - def environment_fingerprints( dep_groups: tuple[str, ...] = TYPECHECK_DEP_GROUPS, ) -> tuple[str, ...]: @@ -311,377 +267,70 @@ def environment_fingerprints( ) -def cache_key(base_point: str, fingerprints: tuple[str, ...]) -> str: - return hashlib.sha256("|".join((base_point, *fingerprints)).encode()).hexdigest()[ - :16 - ] - - -def cache_path( - directory: Path, base_point: str, fingerprints: tuple[str, ...] -) -> Path: - return directory / f"{CACHE_FILE_PREFIX}{cache_key(base_point, fingerprints)}.json" - - -def default_cache_dir() -> Path: - common = Path(_run(["git", "rev-parse", "--git-common-dir"]).strip()) - resolved = common if common.is_absolute() else REPO_ROOT / common - return resolved / "litellm-lint-cache" - - -def validated_counts(data: object) -> dict[str, int] | None: - counts: Final = data.get("counts") if isinstance(data, dict) else None - if not isinstance(counts, dict): - return None - if not all( - isinstance(code, str) and isinstance(total, int) and not isinstance(total, bool) - for code, total in counts.items() - ): - return None - return counts - - -def load_cached_counts(path: Path) -> dict[str, int] | None: - try: - data = json.loads(path.read_text()) - except (OSError, json.JSONDecodeError): - return None - return validated_counts(data) - - -def scratch_path(path: Path) -> Path: - """In-flight scratch for the tmp+rename write. Dot-prefixed so the prune - glob in `store_counts` can never match it (a concurrent run would otherwise - unlink it between write and rename), and pid-suffixed so two concurrent - writers of the same entry never share a scratch.""" - return path.with_name(f".{path.name}.{os.getpid()}.tmp") - - -def counts_payload(base_point: str, counts: Mapping[str, int]) -> str: - return ( - json.dumps( - {"base_point": base_point, "counts": dict(sorted(counts.items()))}, - indent=2, - ) - + "\n" - ) - - -def entry_recency(path: Path) -> float: - try: - return path.stat().st_mtime - except OSError: - return 0.0 - - -def evicted_beyond_cap(entries: Sequence[Path], keep: int) -> tuple[Path, ...]: - newest_first: Final = sorted(entries, key=entry_recency, reverse=True) - return tuple(newest_first[keep:]) - - -def store_counts( - directory: Path, path: Path, base_point: str, counts: Mapping[str, int] -) -> None: - directory.mkdir(parents=True, exist_ok=True) - scratch = scratch_path(path) - scratch.write_text(counts_payload(base_point, counts)) - scratch.replace(path) - siblings: Final = tuple( - entry for entry in directory.glob(f"{CACHE_FILE_PREFIX}*.json") if entry != path - ) - for stale in evicted_beyond_cap(siblings, CACHE_KEEP_ENTRIES - 1): - stale.unlink(missing_ok=True) - - -def parse_origin_slug(url: str) -> str | None: - match: Final = re.fullmatch( - r"(?:git@github\.com:|https://github\.com/)([^/]+/[^/]+?)(?:\.git)?/?", - url.strip(), - ) - return match.group(1) if match else None - - -def origin_slug() -> str | None: - proc: Final = subprocess.run( - ["git", "remote", "get-url", "origin"], - cwd=REPO_ROOT, - capture_output=True, - text=True, - ) - if proc.returncode != 0: - return None - return parse_origin_slug(proc.stdout) - - -def artifact_name(base_point: str) -> str: - return f"{ARTIFACT_NAME_PREFIX}{cache_key(base_point, environment_fingerprints())}" - - -def _gh_output(args: list[str]) -> bytes | None: - try: - proc = subprocess.run( - ["gh", *args], capture_output=True, timeout=GH_TIMEOUT_SECONDS - ) - except (OSError, subprocess.SubprocessError): - return None - return proc.stdout if proc.returncode == 0 else None - - -def _parsed_json(raw: bytes) -> object | None: - try: - return json.loads(raw) - except ValueError: - return None - - -def _artifact_download_url(listing: object) -> str | None: - artifacts: Final = listing.get("artifacts") if isinstance(listing, dict) else None - if not isinstance(artifacts, list) or not artifacts: - return None - newest: Final = artifacts[0] - if not isinstance(newest, dict) or newest.get("expired"): - return None - url: Final = newest.get("archive_download_url") - return url if isinstance(url, str) else None - - -def _counts_json_from_zip(zip_bytes: bytes) -> object | None: - try: - with zipfile.ZipFile(io.BytesIO(zip_bytes)) as archive: - members: Final = [ - name for name in archive.namelist() if name.endswith(".json") - ] - if len(members) != 1: - return None - return json.loads(archive.read(members[0])) - except (zipfile.BadZipFile, ValueError, OSError): - return None - - -def counts_for_base(payload: object, base_point: str) -> dict[str, int] | None: - if not isinstance(payload, dict) or payload.get("base_point") != base_point: - return None - counts: Final = validated_counts(payload) - return counts if counts else None - - -def _fetch_fallback(reason: str) -> None: - sys.stderr.write(f"{reason}; computing base counts locally\n") - - -def fetch_ci_base_counts( - base_point: str, - gh_output: Callable[[list[str]], bytes | None] = _gh_output, -) -> dict[str, int] | None: - """Base counts from the CI artifact published for `base_point`, or None. - - Every failure mode (no gh, no auth, offline, expired or missing artifact, - malformed payload, counts for a different commit) returns None so the - caller falls back to the local base pass; the fetch is an optimization and - must never make the gate less available than local compute alone.""" - slug: Final = origin_slug() - if slug is None: - return _fetch_fallback("origin remote is not a github.com URL") - name: Final = artifact_name(base_point) - listing: Final = gh_output( - ["api", f"repos/{slug}/actions/artifacts?name={name}&per_page=1"] - ) - if listing is None: - return _fetch_fallback(f"could not list CI artifacts named {name}") - url: Final = _artifact_download_url(_parsed_json(listing)) - if url is None: - return _fetch_fallback(f"no usable CI artifact named {name}") - zip_bytes: Final = gh_output(["api", url]) - if zip_bytes is None: - return _fetch_fallback(f"download failed for CI artifact {name}") - counts: Final = counts_for_base(_counts_json_from_zip(zip_bytes), base_point) - if counts is None: - return _fetch_fallback( - f"CI artifact {name} is not valid base counts for {base_point[:12]}" - ) - sys.stderr.write(f"base counts fetched from CI artifact {name}\n") - return counts - - -def base_counts_cached( - base_point: str, - cache_dir: Path | None = None, - compute: Callable[[str], dict[str, int]] = base_counts, - fetch: Callable[[str], dict[str, int] | None] = fetch_ci_base_counts, -) -> dict[str, int]: - """`base_counts` memoized on disk. The base tree at a given commit is - immutable, so its counts are a pure function of the merge-base plus the - environment fingerprints in the cache key; an empty result is never stored - because it is the signature of a crashed pass, not a clean tree. On a disk - miss the counts CI already published for the merge-base are fetched before - the expensive local base pass; a fetch miss of any kind computes locally.""" - directory = default_cache_dir() if cache_dir is None else cache_dir - path = cache_path(directory, base_point, environment_fingerprints()) - cached = load_cached_counts(path) - if cached is not None: - return cached - fetched: Final = fetch(base_point) - if fetched: - store_counts(directory, path, base_point, fetched) - return fetched - counts = compute(base_point) - if counts: - store_counts(directory, path, base_point, counts) - return counts - - -def evaluate( - head: Mapping[str, int], - base: Mapping[str, int], - budget: Mapping[str, Mapping[str, int]], -) -> list[Breach]: - breaches = [] - for code, total in head.items(): - spec = budget.get(code) - cap = spec["limit"] if spec else DEFAULT_LIMIT - prior = base.get(code, 0) - if total > cap and total > prior: - breaches.append(Breach(code, total, cap, total - prior)) - return sorted(breaches) - - -def is_vacuous_run( - counts: Mapping[str, int], budget: Mapping[str, Mapping[str, int]] -) -> bool: - """True when nothing was parsed but the budget expects errors -- the - signature of a type checker that produced no output. `run_basedpyright` - already fails crash exit codes, so this guards the remaining case: a run - that exits cleanly while emitting nothing, which would otherwise clear - every limit and pass silently.""" - return not counts and any(spec["limit"] for spec in budget.values()) - - -def ratcheted_budget( - budget: Mapping[str, Mapping[str, int]], - current: Mapping[str, int], - base: Mapping[str, int], -) -> dict[str, dict[str, int]]: - """Each rule's limit lowered by the errors `current` fixed vs `base`. - - `base` is the count at the branch point (the commit this branch diverged - from). The drop is clamped to what was actually cleared (a rule that grew - stays put), so the limit only ever falls. Rules absent from the budget are - dropped: a genuinely new error category is added to the JSON deliberately, - not on update. - """ - return { - code: { - "limit": max(0, spec["limit"] - max(0, base.get(code, 0) - current.get(code, 0))) - } - for code, spec in sorted(budget.items()) - } - - -def cmd_update(current: Mapping[str, int], base_ref: str) -> None: - """Ratchet each rule's limit down by the errors this branch fixed. - - `current` is the working-tree count; the reference count comes - from a second basedpyright pass over a detached worktree at the branch point - (the merge-base with `base_ref`), so a branch's fixes tighten its own ceilings - by exactly what they cleared since it diverged, and limits never rise. - """ - budget = json.loads(BUDGET_PATH.read_text()) if BUDGET_PATH.exists() else {} - base_point = resolve_base_point(base_ref) - updated = ratcheted_budget(budget, current, base_counts_cached(base_point)) - BUDGET_PATH.write_text(json.dumps(updated, indent=2, sort_keys=True) + "\n") - cleared = sum(budget[code]["limit"] - updated[code]["limit"] for code in updated) - print( - f"Ratcheted basedpyright limits down by {cleared} errors this branch fixed " - f"across {len(updated)} rules" - ) - - -def cmd_emit_counts(head: Mapping[str, int], directory: Path, head_sha: str) -> None: - """Write HEAD's per-rule counts as the file the publisher workflow uploads. - - The filename stem is exactly the artifact name `fetch_ci_base_counts` will - later look up for this commit, so emit and fetch cannot drift apart. Empty - counts are refused for the same reason `is_vacuous_run` exists: a pass that - produced nothing almost certainly crashed, and publishing it would poison - every branch that fetches it.""" - if not head: - print( - "FAIL: basedpyright produced no errors; refusing to publish empty base " - "counts because the pass almost certainly crashed or emitted nothing." - ) - raise SystemExit(1) - name: Final = artifact_name(head_sha) - directory.mkdir(parents=True, exist_ok=True) - (directory / f"{name}.json").write_text(counts_payload(head_sha, head)) - print( - f"Emitted base counts for {head_sha} as {name}.json " - f"({sum(head.values())} errors total)" - ) +def checker_identity(dep_groups: tuple[str, ...] = TYPECHECK_DEP_GROUPS) -> Checker: + return Checker("basedpyright", environment_fingerprints(dep_groups)) def cmd_check(head: Mapping[str, int], base_ref: str) -> None: - budget = json.loads(BUDGET_PATH.read_text()) - if is_vacuous_run(head, budget): - expected = sum(spec["limit"] for spec in budget.values()) + if not head: print( - f"FAIL: basedpyright produced no errors, but {BUDGET_PATH.name} allows " - f"up to ~{expected}. The type checker almost certainly crashed or emitted " - f"nothing; refusing to certify a vacuous run." + "FAIL: basedpyright produced no errors. The type checker almost certainly " + "crashed or emitted nothing; refusing to certify a vacuous run." ) raise SystemExit(1) - if not over_ceiling(head, budget): + base_point: Final = resolve_base_point(base_ref) + base: Final = base_counts_cached(checker_identity(), base_point, base_counts) + if not base: print( - f"OK: every rule is within its basedpyright limit ({sum(head.values())} errors total)" - ) - return - base_point = resolve_base_point(base_ref) - base = base_counts_cached(base_point) - if is_vacuous_run(base, budget): - print( - f"FAIL: basedpyright produced no errors for the base tree at " - f"{base_point[:12]}, so every rule would look freshly added. The base " - f"pass almost certainly crashed; refusing to blame this change for it." + f"FAIL: basedpyright produced no errors for the base tree at {base_point[:12]}, " + "so every rule would look freshly added. The base pass almost certainly " + "crashed; refusing to blame this change for it." ) raise SystemExit(1) - breaches = evaluate(head, base, budget) + judge(head, base, base_point) + + +def judge(head: Mapping[str, int], base: Mapping[str, int], base_point: str) -> None: + breaches: Final = evaluate(head, base, ANY_CAPS) if not breaches: print( - f"OK: every rule is within its basedpyright limit or no higher than base ({sum(head.values())} errors total)" + f"OK: every basedpyright rule is within its ceiling " + f"({sum(head.values())} errors total, base {base_point[:12]})" ) return - print("FAIL: basedpyright errors exceed the per-rule limit:") + print(f"FAIL: basedpyright errors grew past their ceiling (base {base_point[:12]}):") for breach in breaches: - print( - f" {breach.code}: total {breach.total} over limit {breach.cap} (this change added {breach.added})" - ) + print(f" {breach.rule}: total {breach.total} over ceiling {breach.ceiling} (this change added {breach.added})") print( - "Reduce the new errors or remove an equal number elsewhere; the ceiling is " - "the limit in basedpyright-code-budget.json." + "Reduce the new errors or remove an equal number elsewhere; the ceiling is the merge-base " + "count, or the cap in ANY_CAPS (scripts/type_check_gate.py) when that is higher." ) - summary = "; ".join(f"{b.code} {b.total}/{b.cap} (+{b.added})" for b in breaches) + summary: Final = "; ".join(f"{b.rule} {b.total}/{b.ceiling} (+{b.added})" for b in breaches) print(f"BREACHED RULES: {summary}") raise SystemExit(1) def main() -> None: - parser = argparse.ArgumentParser(description=__doc__) + parser: Final = argparse.ArgumentParser(description=__doc__) parser.add_argument("--base", help="Comparison ref (default: origin's current default branch)") - parser.add_argument("--update", action="store_true") - parser.add_argument("--emit-counts-dir", type=Path) - args = parser.parse_args() + parser.add_argument( + "--emit-counts-dir", + type=Path, + help="Write HEAD's per-rule counts to this directory as a base-counts artifact instead of gating", + ) + args: Final = parser.parse_args() from default_branch import resolve_base_ref from gate_slot_lock import held_slot - base_ref: Final = None if args.emit_counts_dir is not None else resolve_base_ref(args.base, REPO_ROOT) + if args.emit_counts_dir is not None: + with held_slot(): + ensure_typecheck_env() + emit_counts(checker_identity(), count_basedpyright(run_basedpyright()), args.emit_counts_dir, head_sha()) + return + base_ref: Final = resolve_base_ref(args.base, REPO_ROOT) with held_slot(): ensure_typecheck_env() - head = count_basedpyright(run_basedpyright()) - if args.emit_counts_dir is not None: - cmd_emit_counts( - head, args.emit_counts_dir, _run(["git", "rev-parse", "HEAD"]).strip() - ) - elif base_ref is not None: - cmd_update(head, base_ref) if args.update else cmd_check(head, base_ref) + cmd_check(count_basedpyright(run_basedpyright()), base_ref) if __name__ == "__main__": diff --git a/scripts/type_discipline_gate.py b/scripts/type_discipline_gate.py index 3293f32d565..5e6745c041d 100644 --- a/scripts/type_discipline_gate.py +++ b/scripts/type_discipline_gate.py @@ -1,56 +1,62 @@ #!/usr/bin/env python3 -"""Total-count gate for the LIT* rules in scripts/check_type_discipline.py. +"""Delta-vs-base gate for the LIT* rules in scripts/check_type_discipline.py. -Sibling of scripts/ruff_strict_gate.py. Each rule listed in -type-discipline-budget.json has a hard ``limit``. The gate counts each rule -across the whole `litellm` tree and fails when a rule is both over its limit and -higher than the base it merges into, so a change is blamed for the violations it -adds, never for drift that already exists in the base. +Sibling of scripts/ruff_strict_gate.py. Each rule is counted across the whole +`litellm` tree at HEAD and at the merge-base with the branch this change merges +into, and the gate fails only when a rule grew past the merge-base count, so a +change is blamed for the violations it adds, never for drift that already exists +in the base. There is no committed budget: the merge-base count is the ceiling, +so it moves only when the base branch does. -Rules not present in the budget are ignored, but today every rule the checker -emits is gated: LIT001 (mutable collection in any annotation), LIT003/LIT004 -(noqa / pyright-mypy ignore without codes or reason), LIT006 (cast), LIT008 (`**kwargs`), LIT009 (inert -`# type: ignore`, dead syntax while enableTypeIgnoreComments is false), LIT010 -(assignment without a Final declaration; suppress deliberate rebinding with -`# rebind-ok: `), LIT011 (parameter rebinding or in-place mutation), and -LIT012 (TypedDict field without a `ReadOnly[...]` qualifier; suppress with -`# writable-ok: `), and LIT014 (comprehension with more than one `for` -or `if` clause; suppress with `# comprehension-ok: ` on a spanned -line, which belongs to the innermost violating comprehension spanning it and -to any single-line violating comprehension on that line) carry limits at -or above their current count to ratchet down; LIT005 (`*-ok` suppression -without a reason) is frozen at limit 0 -so any net-new reasonless suppression trips the gate; LIT013 (`*-ok` suppression -that suppresses nothing) is frozen at 0 for the same reason; and LIT007 -(TypeGuard/TypeIs) is a hard zero. -LIT010 and LIT011 were seeded at 1.5x the count left after the sweep that -annotated every never-rebound name with Final, so that headroom is the hard -line new code cannot cross. -``--update`` ratchets a limit down by the violations this branch fixed relative -to its branch point (the merge-base). A rule absent from the budget at the -merge-base was seeded on this branch; ``--update`` leaves its limit untouched, -because the base tree predates the rule and its whole grandfathered count would -otherwise be misread as "fixed", collapsing the deliberate headroom to zero. +Every rule the checker emits is gated: LIT001 (mutable collection in any +annotation), LIT003/LIT004 (noqa / pyright-mypy ignore without codes or +reason), LIT005 (`*-ok` suppression without a reason), LIT006 (cast), LIT007 +(TypeGuard/TypeIs), LIT008 (`**kwargs`), LIT009 (inert `# type: ignore`, dead +syntax while enableTypeIgnoreComments is false), LIT010 (assignment without a +Final declaration; suppress deliberate rebinding with `# rebind-ok: `), +LIT011 (parameter rebinding or in-place mutation), LIT012 (TypedDict field +without a `ReadOnly[...]` qualifier; suppress with `# writable-ok: `), +LIT013 (`*-ok` suppression that suppresses nothing), LIT014 (comprehension +with more than one `for` or `if` clause; suppress with +`# comprehension-ok: ` on a spanned line, which belongs to the innermost +violating comprehension spanning it and to any single-line violating +comprehension on that line), and LIT015 (pydantic model not frozen; suppress +with `# frozen-ok: `). + +The merge-base counts come from scripts/lint_base_counts.py: the disk cache, +then the CI artifact published for that commit, then a pass of the current +checker over a detached worktree at the merge-base, so a rule change on this +branch is measured on both sides. ``--emit-counts-dir`` writes HEAD's counts +as the file that artifact is built from. """ import argparse -import json import re import shutil import subprocess import sys import tempfile from collections import Counter +from collections.abc import Mapping, Sequence from pathlib import Path from typing import Final, NamedTuple -REPO_ROOT = Path(__file__).resolve().parent.parent -CHECKER = REPO_ROOT / "scripts" / "check_type_discipline.py" -BUDGET_PATH = REPO_ROOT / "type-discipline-budget.json" -TARGET = "litellm" +from lint_base_counts import ( + Checker, + base_counts_cached, + emit_counts, + evaluate, + head_sha, + resolve_base_point, + sha256_of, +) -_HUNK = re.compile(r"^@@ -\d+(?:,\d+)? \+(\d+)(?:,(\d+))? @@") -_LINE = re.compile(r"^(?P.+?):(?P\d+): (?PLIT\d+) ") +REPO_ROOT: Final = Path(__file__).resolve().parent.parent +CHECKER: Final = REPO_ROOT / "scripts" / "check_type_discipline.py" +TARGET: Final = "litellm" + +_HUNK: Final = re.compile(r"^@@ -\d+(?:,\d+)? \+(\d+)(?:,(\d+))? @@") +_LINE: Final = re.compile(r"^(?P.+?):(?P\d+): (?PLIT\d+) ") class Violation(NamedTuple): @@ -59,41 +65,19 @@ class Violation(NamedTuple): code: str -class Breach(NamedTuple): - rule: str - total: int - cap: int - added: int - - -def _run(cmd: list, cwd: Path = REPO_ROOT) -> str: - proc = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True) +def _run(cmd: Sequence[str], cwd: Path = REPO_ROOT) -> str: + proc: Final = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True) if proc.returncode not in (0, 1): sys.stderr.write(proc.stderr) raise SystemExit(f"{cmd[0]} exited {proc.returncode}") return proc.stdout -def resolve_base_point(base_ref: str, cwd: Path = REPO_ROOT) -> str: - """The snapshot commit base counts are measured at: merge-base(base_ref, HEAD), - made aware of an in-progress merge. Mid-merge, HEAD is still the pre-merge tip, - so its merge-base is the old branch point and every violation the base gained - since then would be blamed on this change. While MERGE_HEAD exists, prefer - merge-base(base_ref, MERGE_HEAD) whenever it is the newer of the two.""" - head_point: Final = _run(["git", "merge-base", base_ref, "HEAD"], cwd=cwd).strip() - if not head_point: - return base_ref - merge_head: Final = _run(["git", "rev-parse", "--verify", "--quiet", "MERGE_HEAD"], cwd=cwd).strip() - if not merge_head: - return head_point - merge_point: Final = _run(["git", "merge-base", base_ref, merge_head], cwd=cwd).strip() - if not merge_point: - return head_point - older: Final = _run(["git", "merge-base", head_point, merge_point], cwd=cwd).strip() - return merge_point if older == head_point else head_point +def checker_identity(checker: Path = CHECKER) -> Checker: + return Checker("type-discipline", (sha256_of(checker),)) -def _check(root: Path, checker: Path) -> list: +def _check(root: Path, checker: Path) -> list[Violation]: # Resolve root first: on macOS tempfile dirs (/var/...) resolve to /private/var/..., # and the checker prints already-resolved absolute paths, so relative_to would fail. root = root.resolve() @@ -110,22 +94,22 @@ def _check(root: Path, checker: Path) -> list: return found -def head_violations() -> list: +def head_violations() -> list[Violation]: return _check(REPO_ROOT, CHECKER) -def count_by_rule(violations: list) -> dict: +def count_by_rule(violations: Sequence[Violation]) -> dict[str, int]: return dict(Counter(v.code for v in violations)) -def base_counts(ref: str) -> dict: - parent = Path(tempfile.mkdtemp(prefix="lit_base_")) - worktree = parent / "wt" +def base_counts(ref: str) -> dict[str, int]: + parent: Final = Path(tempfile.mkdtemp(prefix="lit_base_")) + worktree: Final = parent / "wt" try: _run(["git", "worktree", "add", "--detach", str(worktree), ref]) # Measure the base with the *current* rule logic, not whatever shipped at base. (worktree / "scripts").mkdir(parents=True, exist_ok=True) - checker = worktree / "scripts" / "check_type_discipline.py" + checker: Final = worktree / "scripts" / "check_type_discipline.py" shutil.copy(CHECKER, checker) return count_by_rule(_check(worktree, checker)) finally: @@ -140,27 +124,8 @@ def base_counts(ref: str) -> dict: shutil.rmtree(parent, ignore_errors=True) -def over_ceiling(head: dict, budget: dict) -> frozenset: - """Rules whose head count already exceeds their limit. - - A rule can only breach when it is over its limit, so when none are the base - comparison cannot change the verdict and the base worktree scan can be skipped. - """ - return frozenset(rule for rule, spec in budget.items() if head.get(rule, 0) > spec["limit"]) - - -def evaluate(head: dict, base: dict, budget: dict) -> list: - breaches = [] - for rule, spec in budget.items(): - cap = spec["limit"] - total = head.get(rule, 0) - if total > cap and total > base.get(rule, 0): - breaches.append(Breach(rule, total, cap, total - base.get(rule, 0))) - return sorted(breaches) - - -def parse_changed_lines(diff_text: str) -> dict: - changed: dict = {} +def parse_changed_lines(diff_text: str) -> dict[str, set[int]]: + changed: dict[str, set[int]] = {} path = None for line in diff_text.splitlines(): if line.startswith("+++ b/"): @@ -172,104 +137,54 @@ def parse_changed_lines(diff_text: str) -> dict: return changed -def introduced(violations: list, changed: dict) -> list: +def introduced(violations: Sequence[Violation], changed: Mapping[str, set[int]]) -> list[Violation]: return [v for v in violations if v.line in changed.get(v.file, set())] def cmd_check(base: str) -> None: - budget = json.loads(BUDGET_PATH.read_text()) - head = head_violations() - head_counts = count_by_rule(head) - if not over_ceiling(head_counts, budget): - print(f"OK: every LIT rule is within its codebase ceiling (base {base})") - return - base_point = resolve_base_point(base) - breaches = evaluate(head_counts, base_counts(base_point), budget) + head: Final = head_violations() + base_point: Final = resolve_base_point(base) + breaches: Final = evaluate(count_by_rule(head), base_counts_cached(checker_identity(), base_point, base_counts)) if not breaches: - print(f"OK: every LIT rule is within its codebase ceiling (base {base})") + print(f"OK: no LIT rule grew past its merge-base count (base {base})") return - new = introduced( - head, - parse_changed_lines(_run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET])), - ) - print(f"FAIL: LIT-rule totals exceed their limit (base {base}):") + diff: Final = _run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET]) + new: Final = introduced(head, parse_changed_lines(diff)) + print(f"FAIL: LIT-rule totals grew past their merge-base count (base {base}):") for breach in breaches: - print(f" {breach.rule}: total {breach.total} over limit {breach.cap} (this change added {breach.added})") + print(f" {breach.rule}: total {breach.total} over ceiling {breach.ceiling} (this change added {breach.added})") for violation in sorted(v for v in new if v.code == breach.rule): print(f" {violation.file}:{violation.line}") print( "Remove the new violations, give each a reason (`# noqa: XXX # `, " - "`# pyright: ignore[rule] # `, `# mutable-ok: `, " - "`# cast-ok: `, `# guard-ok: `, `# kwargs-ok: `, " - "`# rebind-ok: `, `# writable-ok: `, " - "`# comprehension-ok: `), or remove an equal " - "number elsewhere; the ceiling " - "is the limit in type-discipline-budget.json." + "`# pyright: ignore[rule] # `, `# mutable-ok: `, `# cast-ok: `, " + "`# guard-ok: `, `# kwargs-ok: `, `# rebind-ok: `, " + "`# writable-ok: `, `# comprehension-ok: `, `# frozen-ok: `), " + "or remove an equal number " + "elsewhere; the ceiling is the merge-base count." ) raise SystemExit(1) -def ratcheted_budget(budget: dict, current: dict, base: dict, seeded: frozenset = frozenset()) -> dict: - """Each rule's limit lowered by the violations `current` fixed vs `base`. - - `base` is the count at the branch point (the commit this branch diverged - from). The drop is clamped to what was actually cleared (a rule that grew - stays put), so the limit only ever falls. Rules in `seeded` were introduced - on this branch with deliberate grandfathered headroom; their limits pass - through untouched, since the base predates the rule and comparing against it - would misread the entire grandfathered count as fixed. - """ - return { - rule: { - "limit": spec["limit"] - if rule in seeded - else max(0, spec["limit"] - max(0, base.get(rule, 0) - current.get(rule, 0))) - } - for rule, spec in sorted(budget.items()) - } - - -def _base_budget_rules(base_point: str) -> frozenset: - proc = subprocess.run( - ["git", "show", f"{base_point}:{BUDGET_PATH.name}"], - cwd=REPO_ROOT, - capture_output=True, - text=True, - ) - if proc.returncode != 0: - return frozenset() - return frozenset(json.loads(proc.stdout)) - - -def cmd_update(base_ref: str) -> None: - """Ratchet each rule's limit down by the violations this branch fixed. - - The working-tree count is compared against a checker pass over a detached - worktree at the branch point (the merge-base with `base_ref`), so a branch's - fixes tighten its own ceilings by exactly what they cleared since it diverged. - """ - budget = json.loads(BUDGET_PATH.read_text()) - base_point = resolve_base_point(base_ref) - seeded = frozenset(budget) - _base_budget_rules(base_point) - updated = ratcheted_budget(budget, count_by_rule(head_violations()), base_counts(base_point), seeded) - BUDGET_PATH.write_text(json.dumps(updated, indent=2, sort_keys=True) + "\n") - cleared = sum(budget[rule]["limit"] - updated[rule]["limit"] for rule in updated) - print(f"Ratcheted LIT-rule limits down by {cleared} violations this branch fixed") - if seeded: - print("Left untouched (seeded on this branch, absent from the base budget): " + ", ".join(sorted(seeded))) - - def main() -> None: - parser = argparse.ArgumentParser(description=__doc__) + parser: Final = argparse.ArgumentParser(description=__doc__) parser.add_argument("--base", help="Comparison ref (default: origin's current default branch)") - parser.add_argument("--update", action="store_true") - args = parser.parse_args() + parser.add_argument( + "--emit-counts-dir", + type=Path, + help="Write HEAD's per-rule counts to this directory as a base-counts artifact instead of gating", + ) + args: Final = parser.parse_args() from default_branch import resolve_base_ref from gate_slot_lock import held_slot + if args.emit_counts_dir is not None: + with held_slot(): + emit_counts(checker_identity(), count_by_rule(head_violations()), args.emit_counts_dir, head_sha()) + return base_ref: Final = resolve_base_ref(args.base, REPO_ROOT) with held_slot(): - cmd_update(base_ref) if args.update else cmd_check(base_ref) + cmd_check(base_ref) if __name__ == "__main__": diff --git a/test-quality-budget.json b/test-quality-budget.json deleted file mode 100644 index 6f3dde8461b..00000000000 --- a/test-quality-budget.json +++ /dev/null @@ -1,26 +0,0 @@ -{ - "TQ001": { - "limit": 733 - }, - "TQ002": { - "limit": 737 - }, - "TQ003": { - "limit": 62 - }, - "TQ004": { - "limit": 469 - }, - "TQ005": { - "limit": 2399 - }, - "TQ006": { - "limit": 34 - }, - "TQ007": { - "limit": 117 - }, - "TQ009": { - "limit": 59 - } -} diff --git a/tests/audio_tests/test_audio_speech.py b/tests/audio_tests/test_audio_speech.py index 998de5ecc3b..781dff654f7 100644 --- a/tests/audio_tests/test_audio_speech.py +++ b/tests/audio_tests/test_audio_speech.py @@ -8,248 +8,12 @@ from dotenv import load_dotenv load_dotenv() from pathlib import Path -from unittest.mock import AsyncMock, MagicMock, patch import pytest import litellm -async def _run_audio_speech_litellm(sync_mode, model, api_base, api_key): - litellm.turn_on_debug() - speech_file_path = Path(__file__).parent / "speech.mp3" - - if sync_mode: - response = litellm.speech( - model=model, - voice="alloy", - input="the quick brown fox jumped over the lazy dogs", - api_base=api_base, - api_key=api_key, - organization=None, - project=None, - max_retries=1, - timeout=600, - client=None, - optional_params={}, - ) - - from litellm.types.llms.openai import HttpxBinaryResponseContent - - assert isinstance(response, HttpxBinaryResponseContent) - else: - response = await litellm.aspeech( - model=model, - voice="alloy", - input="the quick brown fox jumped over the lazy dogs", - api_base=api_base, - api_key=api_key, - organization=None, - project=None, - max_retries=1, - timeout=600, - client=None, - optional_params={}, - ) - - from litellm.llms.openai.openai import HttpxBinaryResponseContent - - assert isinstance(response, HttpxBinaryResponseContent) - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=1) -async def test_audio_speech_litellm_azure(sync_mode): - await _run_audio_speech_litellm( - sync_mode=sync_mode, - model="azure/tts", - api_base=os.getenv("AZURE_TTS_API_BASE"), - api_key=os.getenv("AZURE_TTS_API_KEY"), - ) - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=1) -async def test_audio_speech_litellm_openai(sync_mode): - await _run_audio_speech_litellm( - sync_mode=sync_mode, - model="openai/tts-1", - api_base=None, - api_key=os.getenv("OPENAI_API_KEY"), - ) - - - - -@pytest.mark.flaky(retries=6, delay=2) -@pytest.mark.asyncio -async def test_speech_litellm_vertex_async(): - # Mock the response - mock_response = AsyncMock() - - def return_val(): - return { - "audioContent": "dGVzdCByZXNwb25zZQ==", - } - - mock_response.json = return_val - mock_response.status_code = 200 - - # Set up the mock for asynchronous calls - with patch( - "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", - new_callable=AsyncMock, - ) as mock_async_post: - mock_async_post.return_value = mock_response - model = "vertex_ai/test" - - try: - response = await litellm.aspeech( - model=model, - input="async hello what llm guardrail do you have", - ) - except litellm.APIConnectionError as e: - if "Your default credentials were not found" in str(e): - pytest.skip("skipping test, credentials not found") - - # Assert asynchronous call - mock_async_post.assert_called_once() - _, kwargs = mock_async_post.call_args - print("call args", kwargs) - - assert kwargs["url"] == "https://texttospeech.googleapis.com/v1/text:synthesize" - - assert "x-goog-user-project" in kwargs["headers"] - assert kwargs["headers"]["Authorization"] is not None - - assert kwargs["json"] == { - "input": {"text": "async hello what llm guardrail do you have"}, - "voice": {"languageCode": "en-US", "name": "en-US-Studio-O"}, - "audioConfig": {"audioEncoding": "LINEAR16", "speakingRate": "1"}, - } - - -@pytest.mark.asyncio -async def test_speech_litellm_vertex_async_with_voice(): - # Mock the response - mock_response = AsyncMock() - - def return_val(): - return { - "audioContent": "dGVzdCByZXNwb25zZQ==", - } - - mock_response.json = return_val - mock_response.status_code = 200 - - # Set up the mock for asynchronous calls - with patch( - "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", - new_callable=AsyncMock, - ) as mock_async_post: - mock_async_post.return_value = mock_response - model = "vertex_ai/test" - - try: - response = await litellm.aspeech( - model=model, - input="async hello what llm guardrail do you have", - voice={ - "languageCode": "en-UK", - "name": "en-UK-Studio-O", - }, - audioConfig={ - "audioEncoding": "LINEAR22", - "speakingRate": "10", - }, - ) - except litellm.APIConnectionError as e: - if "Your default credentials were not found" in str(e): - pytest.skip("skipping test, credentials not found") - - # Assert asynchronous call - mock_async_post.assert_called_once() - _, kwargs = mock_async_post.call_args - print("call args", kwargs) - - assert kwargs["url"] == "https://texttospeech.googleapis.com/v1/text:synthesize" - - assert "x-goog-user-project" in kwargs["headers"] - assert kwargs["headers"]["Authorization"] is not None - - assert kwargs["json"] == { - "input": {"text": "async hello what llm guardrail do you have"}, - "voice": {"languageCode": "en-UK", "name": "en-UK-Studio-O"}, - "audioConfig": {"audioEncoding": "LINEAR22", "speakingRate": "10"}, - } - - -@pytest.mark.asyncio -async def test_speech_litellm_vertex_async_with_voice_ssml(): - # Mock the response - mock_response = AsyncMock() - - def return_val(): - return { - "audioContent": "dGVzdCByZXNwb25zZQ==", - } - - mock_response.json = return_val - mock_response.status_code = 200 - - ssml = """ - -

Hello, world!

-

This is a test of the text-to-speech API.

-
- """ - - # Set up the mock for asynchronous calls - with patch( - "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", - new_callable=AsyncMock, - ) as mock_async_post: - mock_async_post.return_value = mock_response - model = "vertex_ai/test" - - try: - response = await litellm.aspeech( - input=ssml, - model=model, - voice={ - "languageCode": "en-UK", - "name": "en-UK-Studio-O", - }, - audioConfig={ - "audioEncoding": "LINEAR22", - "speakingRate": "10", - }, - ) - except litellm.APIConnectionError as e: - if "Your default credentials were not found" in str(e): - pytest.skip("skipping test, credentials not found") - - # Assert asynchronous call - mock_async_post.assert_called_once() - _, kwargs = mock_async_post.call_args - print("call args", kwargs) - - assert kwargs["url"] == "https://texttospeech.googleapis.com/v1/text:synthesize" - - assert "x-goog-user-project" in kwargs["headers"] - assert kwargs["headers"]["Authorization"] is not None - - assert kwargs["json"] == { - "input": {"ssml": ssml}, - "voice": {"languageCode": "en-UK", "name": "en-UK-Studio-O"}, - "audioConfig": {"audioEncoding": "LINEAR22", "speakingRate": "10"}, - } - - - - @pytest.mark.asyncio @pytest.mark.flaky(retries=3, delay=1) async def test_azure_ava_tts_async(): diff --git a/tests/audio_tests/test_whisper.py b/tests/audio_tests/test_whisper.py index 5380a57c871..99b084be1d3 100644 --- a/tests/audio_tests/test_whisper.py +++ b/tests/audio_tests/test_whisper.py @@ -15,29 +15,21 @@ from dotenv import load_dotenv from openai import AsyncOpenAI import litellm -from litellm.integrations.custom_logger import CustomLogger # Get the current directory of the file being run pwd = os.path.dirname(os.path.realpath(__file__)) print(pwd) file_path = os.path.join(pwd, "gettysburg.wav") -file2_path = os.path.join(pwd, "eagle.wav") with open(file_path, "rb") as _f: _GETTYSBURG_BYTES = _f.read() -with open(file2_path, "rb") as _f: - _EAGLE_BYTES = _f.read() def _audio_file(): return ("gettysburg.wav", _GETTYSBURG_BYTES, "audio/wav") -def _audio_file2(): - return ("eagle.wav", _EAGLE_BYTES, "audio/wav") - - load_dotenv() from litellm import Router @@ -75,104 +67,3 @@ async def test_transcription_azure_whisper(response_format, timestamp_granularit response_format=response_format, timestamp_granularities=timestamp_granularities, ) - - -@pytest.mark.asyncio() -async def test_transcription_caching(): - import litellm - from litellm.caching.caching import Cache - - litellm.set_verbose = True - litellm.cache = Cache() - - # make raw llm api call - - response_1 = await litellm.atranscription( - model="whisper-1", - file=_audio_file(), - ) - - await asyncio.sleep(5) - - # cache hit - - response_2 = await litellm.atranscription( - model="whisper-1", - file=_audio_file(), - ) - - print("response_1", response_1) - print("response_2", response_2) - print("response2 hidden params", response_2._hidden_params) - assert response_2._hidden_params["cache_hit"] is True - - # cache miss - - response_3 = await litellm.atranscription( - model="whisper-1", - file=_audio_file2(), - ) - print("response_3", response_3) - print("response3 hidden params", response_3._hidden_params) - assert response_3._hidden_params.get("cache_hit") is not True - assert response_3.text != response_2.text - - litellm.cache = None - - -@pytest.mark.asyncio -async def test_whisper_log_pre_call(): - from litellm.litellm_core_utils.litellm_logging import Logging - from datetime import datetime - from unittest.mock import patch, MagicMock - - custom_logger = CustomLogger() - - litellm.callbacks = [custom_logger] - - with patch.object(custom_logger, "log_pre_api_call") as mock_log_pre_call: - await litellm.atranscription( - model="whisper-1", - file=_audio_file(), - ) - mock_log_pre_call.assert_called_once() - - -@pytest.mark.asyncio -async def test_gpt_4o_transcribe_model_mapping(): - """Test that GPT-4o transcription models are correctly mapped and not hardcoded to whisper-1""" - - # Test GPT-4o mini transcribe - response = await litellm.atranscription( - model="openai/gpt-4o-mini-transcribe", - file=_audio_file(), - response_format="json", - ) - - # Check that the response contains the correct model in hidden params - assert response._hidden_params is not None - assert response._hidden_params["model"] == "gpt-4o-mini-transcribe" - assert response._hidden_params["custom_llm_provider"] == "openai" - assert response.text is not None - - # Test GPT-4o transcribe - response2 = await litellm.atranscription( - model="openai/gpt-4o-transcribe", file=_audio_file(), response_format="json" - ) - - # Check that the response contains the correct model in hidden params - assert response2._hidden_params is not None - assert response2._hidden_params["model"] == "gpt-4o-transcribe" - assert response2._hidden_params["custom_llm_provider"] == "openai" - assert response2.text is not None - - # Test traditional whisper-1 still works - response3 = await litellm.atranscription( - model="openai/whisper-1", file=_audio_file(), response_format="json" - ) - - # Check that the response contains the correct model in hidden params - assert response3._hidden_params is not None - assert response3._hidden_params["model"] == "whisper-1" - assert response3._hidden_params["custom_llm_provider"] == "openai" - assert response3.text is not None diff --git a/tests/base_sdk_tests/check_base_sdk_install.py b/tests/base_sdk_tests/check_base_sdk_install.py index f680ba92645..e56d48f5e6e 100644 --- a/tests/base_sdk_tests/check_base_sdk_install.py +++ b/tests/base_sdk_tests/check_base_sdk_install.py @@ -6,6 +6,7 @@ pull ``packaging``, ``pluggy`` and ``iniconfig`` into the environment and could the very class of undeclared-dependency bug this guards against. """ +import argparse import importlib.util import sys import traceback @@ -29,12 +30,26 @@ def check_environment_is_base_only() -> str: def check_import() -> str: + from importlib.metadata import distributions as installed_distributions from importlib.metadata import version import litellm + from litellm.llms.brave.search.transformation import BraveSearchConfig + _require(callable(BraveSearchConfig), "Brave search configuration unavailable") _require(bool(litellm.__file__), "litellm has no __file__") - return f"imported litellm {version('litellm')}" + from litellm._version import version as sdk_version + + distributions = tuple( + distribution.metadata["Name"] + for distribution in installed_distributions() + if distribution.metadata["Name"] in ("litellm", "litellm-core") + ) + _require(len(distributions) == 1, f"expected one SDK distribution, found {distributions}") + distribution = distributions[0] + _require(sdk_version == version(distribution), "SDK version does not match installed metadata") + _require("litellm.proxy.proxy_cli" not in sys.modules, "SDK import loaded the proxy CLI") + return f"imported {distribution} {sdk_version}" def check_completion() -> str: @@ -97,12 +112,47 @@ def check_token_counter() -> str: return f"token_counter returned {count}" +def check_tokenizer_dependencies() -> str: + import litellm + from litellm.rust_bridge import tokenizer + from litellm.litellm_core_utils.tokenizer import HuggingFace, HuggingFaceTokenizer, Tokenizer + from litellm.utils import claude_json_str + + _require(isinstance(litellm.encoding, Tokenizer), "runtime alias rejects the default encoding") + native = tokenizer.native_anthropic() + if native is not None: + _require(bool(native.encode("hello")), "native tokenizer returned no tokens") + _require(isinstance(HuggingFaceTokenizer(native), HuggingFace), "runtime alias rejects native tokenizers") + if importlib.util.find_spec("tokenizers") is not None: + python_tokenizer = tokenizer.from_str(claude_json_str) + _require(bool(python_tokenizer.encode("hello").ids), "Python tokenizer returned no tokens") + _require(isinstance(python_tokenizer, HuggingFace), "runtime alias rejects Python tokenizers") + return "installed Python tokenizer and available native tokenizer work" + try: + tokenizer.from_str(claude_json_str) + except ImportError as error: + _require("pip install tokenizers" in str(error), f"missing tokenizer guidance: {error}") + else: + raise AssertionError("Python tokenizer loaded without tokenizers") + return "native tokenizer works; Python tokenizer reports its missing dependency" + + def check_bedrock_credential_resolution() -> str: import os from unittest import mock from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + if importlib.util.find_spec("boto3") is None: + try: + BaseAWSLLM()._sign_request( + service_name="bedrock", headers={}, optional_params={"aws_region_name": "us-east-1"}, + request_data={}, api_base="https://bedrock-runtime.us-east-1.amazonaws.com", api_key="", + ) + except ImportError as error: + _require("pip install boto3" in str(error), f"missing installation guidance: {error}") + return "AWS signing explains how to install boto3" + raise AssertionError("AWS signing unexpectedly worked without boto3") non_aws_environ = {k: v for k, v in os.environ.items() if not k.startswith("AWS_")} with mock.patch.dict(os.environ, non_aws_environ, clear=True): credentials = BaseAWSLLM().get_credentials( @@ -125,6 +175,7 @@ CHECKS: tuple[tuple[str, Callable[[], str]], ...] = ( ("embedding", check_embedding), ("bundled model metadata", check_bundled_model_metadata), ("token counter", check_token_counter), + ("tokenizer dependencies", check_tokenizer_dependencies), ("bedrock credential resolution", check_bedrock_credential_resolution), ) @@ -137,7 +188,14 @@ def _run(check: Callable[[], str]) -> tuple[bool, str]: def main() -> int: - print(f"base SDK smoke check on {sys.executable}") + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--profile", choices=("legacy", "core", "dependencies"), default="legacy") + profile = parser.parse_args().profile + for module in ("boto3", "botocore", "tokenizers", "huggingface_hub"): + present = importlib.util.find_spec(module) is not None + _require(present == (profile != "core"), f"{profile}: unexpected presence of {module}: {present}") + _require(importlib.util.find_spec("jsonschema") is not None, "jsonschema must remain mandatory") + print(f"{profile} SDK smoke check on {sys.executable}") for label, check in CHECKS: passed, detail = _run(check) if not passed: diff --git a/tests/base_sdk_tests/check_sdk_http.py b/tests/base_sdk_tests/check_sdk_http.py new file mode 100644 index 00000000000..d773e252127 --- /dev/null +++ b/tests/base_sdk_tests/check_sdk_http.py @@ -0,0 +1,202 @@ +"""Check installed SDK HTTP behavior against a recording loopback upstream.""" + +import asyncio +import importlib.util +import json +import os +import struct +import zlib +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from queue import Queue +from threading import Event, Thread +from typing import Final +from unittest.mock import patch + +ENVIRONMENT: Final = {"LITELLM_LOCAL_MODEL_COST_MAP": "True", "PYTHON_DOTENV_DISABLED": "1"} + +with patch.dict(os.environ, ENVIRONMENT): + import litellm + +RESPONSES: Final[Queue[tuple[int, bytes]]] = Queue() +REQUESTS: Final[Queue[tuple[str, dict[str, str], bytes]]] = Queue() +ARRIVED: Final = Event() +RELEASE: Final = Event() +MESSAGES: Final = [{"role": "user", "content": "ping"}] +CHAT: Final = { + "id": "chat-http-check", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "pong"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, +} + + +class RecordingUpstream(BaseHTTPRequestHandler): + def log_message(self, format: str, *args: object) -> None: + pass + + def do_POST(self) -> None: + REQUESTS.put((self.path, dict(self.headers), self.rfile.read(int(self.headers["Content-Length"])))) + status, body = RESPONSES.get(timeout=10) + ARRIVED.set() + if status == 0: + RELEASE.wait(timeout=10) + return + self.send_response(status) + self.send_header("Content-Type", "text/event-stream" if body.startswith(b"data:") else "application/vnd.amazon.eventstream" if body[:1] == b"\x00" else "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + +def enqueue(body: object, status: int = 200) -> None: + RESPONSES.put((status, json.dumps(body).encode())) + + +def enqueue_stream() -> None: + chunk: Final = {**CHAT, "object": "chat.completion.chunk", "choices": [{"index": 0, "delta": {"content": "pong"}}]} + RESPONSES.put((200, f"data: {json.dumps(chunk)}\n\ndata: [DONE]\n\n".encode())) + + +def bedrock_event(event: str, payload: object) -> bytes: + def encode_header(name: str, value: str) -> bytes: + return bytes([len(name)]) + name.encode() + b"\x07" + struct.pack(">H", len(value)) + value.encode() + + headers: Final = encode_header(":message-type", "event") + encode_header(":event-type", event) + content: Final = json.dumps(payload).encode() + prelude: Final = struct.pack(">II", len(headers) + len(content) + 16, len(headers)) + frame: Final = prelude + struct.pack(">I", zlib.crc32(prelude)) + headers + content + return frame + struct.pack(">I", zlib.crc32(frame)) + + +def check_http(base: str) -> None: + arguments: Final = dict(model="openai/gpt-4o", messages=MESSAGES, api_key="test-key", api_base=base + "/v1") + enqueue(CHAT) + response: Final = litellm.completion(**arguments) + assert response.choices[0].message.content == "pong" + assert response.usage.total_tokens == 5 + path, headers, body = REQUESTS.get(timeout=10) + assert path == "/v1/chat/completions" + assert headers["Authorization"] == "Bearer test-key" + assert json.loads(body)["messages"] == MESSAGES + enqueue_stream() + assert "".join(part.choices[0].delta.content or "" for part in litellm.completion(**arguments, stream=True)) == "pong" + REQUESTS.get(timeout=10) + for status, exception in ((401, litellm.AuthenticationError), (429, litellm.RateLimitError), (500, litellm.InternalServerError)): + enqueue({"error": {"message": "controlled upstream failure"}}, status) + try: + litellm.completion(**arguments, num_retries=0, max_retries=0) + except exception: + REQUESTS.get(timeout=10) + else: + raise AssertionError(f"HTTP {status} did not raise {exception.__name__}") + enqueue({"error": {"message": "retry once"}}, 429) + enqueue(CHAT) + assert litellm.completion(**arguments, num_retries=1).choices[0].message.content == "pong" + REQUESTS.get(timeout=10) + REQUESTS.get(timeout=10) + + async def check_async() -> None: + enqueue(CHAT) + result: Final = await litellm.acompletion(**arguments) + assert result.choices[0].message.content == "pong" + REQUESTS.get(timeout=10) + enqueue_stream() + stream: Final = await litellm.acompletion(**arguments, stream=True) + assert "".join([part.choices[0].delta.content or "" async for part in stream]) == "pong" + REQUESTS.get(timeout=10) + ARRIVED.clear() + RESPONSES.put((0, b"")) + task: Final = asyncio.create_task(litellm.acompletion(**arguments, num_retries=0, max_retries=0)) + assert await asyncio.to_thread(ARRIVED.wait, 10), "request did not reach upstream" + task.cancel() + try: + await task + except asyncio.CancelledError: + REQUESTS.get(timeout=10) + else: + raise AssertionError("cancellation did not propagate") + finally: + RELEASE.set() + + asyncio.run(check_async()) + enqueue({"object": "list", "model": "text-embedding-3-small", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}], + "usage": {"prompt_tokens": 3, "total_tokens": 3}}) + embedding: Final = litellm.embedding(model="text-embedding-3-small", input=["ping"], api_key="test-key", api_base=base + "/v1") + assert embedding.data[0]["embedding"] == [0.1, 0.2] + assert REQUESTS.get(timeout=10)[0] == "/v1/embeddings" + RELEASE.clear() + RESPONSES.put((0, b"")) + try: + litellm.completion(**arguments, timeout=0.1, num_retries=0, max_retries=0) + except litellm.Timeout: + REQUESTS.get(timeout=10) + else: + raise AssertionError("stalled upstream did not time out") + finally: + RELEASE.set() + for signed in ((False, True) if importlib.util.find_spec("boto3") else (False,)): + for asynchronous in (False, True): + enqueue({ + "output": {"message": {"role": "assistant", "content": [{"text": "pong"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 3, "outputTokens": 2, "totalTokens": 5}, + "metrics": {"latencyMs": 1}, + }) + bedrock: Final = dict( + model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0", messages=MESSAGES, + aws_region_name="us-east-1", aws_bedrock_runtime_endpoint=base, + **({"aws_access_key_id": "test-key", "aws_secret_access_key": "test-secret", "api_key": ""} + if signed else {"api_key": "bearer-key"}), + ) + result: Final = asyncio.run(litellm.acompletion(**bedrock)) if asynchronous else litellm.completion(**bedrock) + assert result.choices[0].message.content == "pong" + assert result.usage.total_tokens == 5 + path, headers, body = REQUESTS.get(timeout=10) + assert path.endswith("/converse") + assert {key.lower(): value for key, value in headers.items()}["authorization"].startswith( + "AWS4-HMAC-SHA256" if signed else "Bearer bearer-key" + ) + assert json.loads(body)["messages"][0]["content"][0]["text"] == "ping" + if importlib.util.find_spec("botocore") is not None: + RESPONSES.put((200, b"".join(( + bedrock_event("messageStart", {"role": "assistant"}), + bedrock_event("contentBlockDelta", {"contentBlockIndex": 0, "delta": {"text": "pong"}}), + bedrock_event("messageStop", {"stopReason": "end_turn"}), + bedrock_event("metadata", {"usage": {"inputTokens": 3, "outputTokens": 2, "totalTokens": 5}}), + )))) + streamed: Final = litellm.completion( + model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0", messages=MESSAGES, + api_key="bearer-key", aws_region_name="us-east-1", aws_bedrock_runtime_endpoint=base, stream=True, + ) + assert "".join(part.choices[0].delta.content or "" for part in streamed) == "pong" + assert REQUESTS.get(timeout=10)[0].endswith("/converse-stream") + enqueue({"id": "msg-http-check", "type": "message", "role": "assistant", "model": "claude-3-sonnet-20240229", + "content": [{"type": "text", "text": "pong"}], "stop_reason": "end_turn", "stop_sequence": None, + "usage": {"input_tokens": 3, "output_tokens": 2}}) + anthropic: Final = litellm.completion(model="anthropic/claude-3-sonnet-20240229", messages=MESSAGES, + api_key="test-key", api_base=base, max_tokens=16) + assert anthropic.choices[0].message.content == "pong" + assert anthropic.usage.total_tokens == 5 + assert REQUESTS.get(timeout=10)[0] == "/v1/messages" + print("PASS installed HTTP: sync/async, streaming, usage, error mapping, retries, cancellation, Bedrock auth") + + +def main() -> None: + with ThreadingHTTPServer(("127.0.0.1", 0), RecordingUpstream) as server: + worker: Final = Thread(target=server.serve_forever, daemon=True) + worker.start() + try: + check_http(f"http://127.0.0.1:{server.server_port}") + assert RESPONSES.empty() and REQUESTS.empty(), "unconsumed HTTP exchanges" + finally: + RELEASE.set() + server.shutdown() + worker.join(timeout=10) + + +if __name__ == "__main__": + with patch.dict(os.environ, ENVIRONMENT): + main() diff --git a/tests/base_sdk_tests/test_core_distribution.py b/tests/base_sdk_tests/test_core_distribution.py new file mode 100644 index 00000000000..e16226e1df8 --- /dev/null +++ b/tests/base_sdk_tests/test_core_distribution.py @@ -0,0 +1,253 @@ +import email +import os +import subprocess +import sys +import tarfile +import zipfile +from pathlib import Path +from typing import Final +from unittest.mock import patch + +import pytest +from packaging.requirements import Requirement + +if sys.version_info >= (3, 11): + import tomllib +else: + import tomli as tomllib + + +ROOT: Final = Path(__file__).resolve().parents[2] + + +@pytest.mark.parametrize("sdist_only", [False, True]) +def test_build_command_selects_requested_distributions(tmp_path: Path, sdist_only: bool) -> None: + from scripts.build_core_distribution import main + + output: Final = tmp_path / "dist" + arguments: Final = ["build_core_distribution.py", "--out-dir", str(output)] + ( + ["--sdist-only"] if sdist_only else [] + ) + with ( + patch.object(sys, "argv", arguments), + patch("scripts.build_core_distribution.stage_core_distribution") as stage, + patch("scripts.build_core_distribution.subprocess.run") as run, + ): + main() + run.assert_called_once() + command: Final = run.call_args.args[0] + assert command[:2] == ["uv", "build"] + assert ("--sdist" in command) is sdist_only + assert "--wheel" not in command + assert command[command.index("--out-dir") + 1] == str(output) + assert run.call_args.kwargs["check"] is True + assert run.call_args.kwargs["cwd"] == stage.call_args.args[1] + assert not stage.call_args.args[1].exists() + + +@pytest.fixture +def source_repository(tmp_path: Path) -> Path: + source: Final = tmp_path / "source" + source.mkdir() + subprocess.run(["git", "init", "--quiet", str(source)], check=True) + return source + + +def test_core_manifest_preserves_runtime_dependencies_without_extras() -> None: + manifest: Final = ROOT / "packaging/litellm-core/pyproject.toml" + assert manifest.is_file(), "The core distribution needs its own build manifest" + core: Final = tomllib.loads(manifest.read_text()) + legacy: Final = tomllib.loads((ROOT / "pyproject.toml").read_text()) + assert core["project"]["name"] == "litellm-core" + removed: Final = {"boto3", "tokenizers", "huggingface-hub"} + core_dependencies: Final = {Requirement(value).name for value in core["project"]["dependencies"]} + legacy_dependencies: Final = {Requirement(value).name for value in legacy["project"]["dependencies"]} + assert not core_dependencies & removed + assert removed <= legacy_dependencies + assert "jsonschema" in core_dependencies + assert legacy_dependencies - removed <= core_dependencies + core_requirements: Final = {Requirement(value) for value in core["project"]["dependencies"]} + retained_requirements: Final = { + Requirement(value) for value in legacy["project"]["dependencies"] if Requirement(value).name not in removed + } + assert retained_requirements <= core_requirements + assert not core["project"].get("optional-dependencies") + assert not core["project"].get("scripts") + + +def test_staging_stamps_release_version_without_modifying_sources(tmp_path: Path, source_repository: Path) -> None: + from scripts.build_core_distribution import SOURCES, stage_core_distribution + + source: Final = source_repository + for name in SOURCES: + path: Final = source / name + path.write_text(f"shared {name}") + (source / "pyproject.toml").write_text('[project]\nname = "litellm"\nversion = "9.8.7rc1"\n') + manifest: Final = source / "packaging/litellm-core/pyproject.toml" + manifest.parent.mkdir(parents=True) + manifest.write_bytes((ROOT / "packaging/litellm-core/pyproject.toml").read_bytes()) + original: Final = manifest.read_bytes() + stage: Final = tmp_path / "stage" + stage_core_distribution(source, stage) + assert tomllib.loads((stage / "pyproject.toml").read_text())["project"]["version"] == "9.8.7rc1" + assert manifest.read_bytes() == original + assert (source / "pyproject.toml").read_text() == '[project]\nname = "litellm"\nversion = "9.8.7rc1"\n' + assert all((stage / name).read_bytes() == (source / name).read_bytes() for name in SOURCES) + + +@pytest.fixture(scope="module") +def distribution_directory() -> Path: + directory: Final = os.environ.get("CORE_DISTRIBUTION_DIR") + if directory is None: + pytest.skip("CORE_DISTRIBUTION_DIR is supplied by the installed-distribution CI job") + return Path(directory) + + +@pytest.fixture(scope="module") +def distributions(distribution_directory: Path) -> tuple[Path, Path]: + wheels: Final = tuple(distribution_directory.glob("litellm_core-*.whl")) + sdists: Final = tuple(distribution_directory.glob("litellm_core-*.tar.gz")) + assert len(wheels) == len(sdists) == 1 + return wheels[0], sdists[0] + + +def test_core_wheel_metadata_and_resources(distributions: tuple[Path, Path]) -> None: + with zipfile.ZipFile(distributions[0]) as wheel: + names: Final = wheel.namelist() + metadata: Final = email.message_from_bytes( + wheel.read(next(n for n in names if n.endswith(".dist-info/METADATA"))) + ) + assert metadata["Name"] == "litellm-core" + assert metadata["Version"] == tomllib.loads((ROOT / "pyproject.toml").read_text())["project"]["version"] + assert not metadata.get_all("Provides-Extra") + assert not any(n.endswith(".dist-info/entry_points.txt") for n in names) + requirements: Final = tomllib.loads((ROOT / "packaging/litellm-core/pyproject.toml").read_text())["project"]["dependencies"] + for python_version in ("3.10", "3.11", "3.12", "3.13", "3.14"): + environment: Final = {"python_version": python_version, "python_full_version": python_version + ".0"} + assert { + (requirement.name, requirement.specifier, frozenset(requirement.extras)) + for item in metadata.get_all("Requires-Dist", ()) + for requirement in (Requirement(item),) + if requirement.marker is None or requirement.marker.evaluate(environment) + } == { + (requirement.name, requirement.specifier, frozenset(requirement.extras)) + for item in requirements + for requirement in (Requirement(item),) + if requirement.marker is None or requirement.marker.evaluate(environment) + } + assert "litellm/model_prices_and_context_window_backup.json" in names + assert "litellm/router_strategy/complexity_router/fuse_presets.json" in names + assert "litellm/proxy/proxy_cli.py" in names + assert any(n.startswith("litellm/rust_bridge/_native.") and n.endswith((".so", ".pyd")) for n in names) + assert not any(n.startswith("litellm/proxy/_experimental/out/") for n in names) + + +def test_core_artifacts_exclude_local_configuration(distributions: tuple[Path, Path]) -> None: + configurations: Final = ( + "litellm/proxy/_new_secret_config.yaml", + "litellm/proxy/_new_new_secret_config.yaml", + "litellm/proxy/_super_secret_config.yaml", + ) + with zipfile.ZipFile(distributions[0]) as wheel, tarfile.open(distributions[1]) as archive: + for name in configurations: + assert name not in wheel.namelist(), f"Core wheel contains ignored configuration {name}" + assert not any(member.name.endswith(f"/{name}") for member in archive.getmembers()) + + +@pytest.mark.parametrize("relative_path", ["rust-toolchain.toml", ".cargo/config.toml"]) +def test_core_sdist_preserves_native_build_configuration(distribution_directory: Path, relative_path: str) -> None: + sdists: Final = tuple(distribution_directory.glob("litellm_core-*.tar.gz")) + assert len(sdists) == 1 + with tarfile.open(sdists[0]) as archive: + name: Final = f"{sdists[0].name.removesuffix('.tar.gz')}/{relative_path}" + assert name in archive.getnames(), f"Core source distribution is missing {relative_path}" + content: Final = archive.extractfile(name) + assert content is not None + assert content.read() == (ROOT / relative_path).read_bytes() + + +def test_core_sdist_rebuilds_without_repository(distributions: tuple[Path, Path], tmp_path: Path) -> None: + with tarfile.open(distributions[1]) as archive: + archive.extractall(tmp_path, filter="data") + source: Final = next(tmp_path.glob("litellm_core-*")) + assert not (source / "scripts/build_core_distribution.py").exists() + assert tomllib.loads((source / "pyproject.toml").read_text())["project"]["name"] == "litellm-core" + result: Final = subprocess.run( + ["uv", "build", "--python", sys.executable, "--wheel", "--out-dir", str(tmp_path / "rebuilt")], + cwd=source, + capture_output=True, + text=True, + ) + assert result.returncode == 0, result.stderr + rebuilt: Final = next((tmp_path / "rebuilt").glob("*.whl")) + with zipfile.ZipFile(distributions[0]) as original, zipfile.ZipFile(rebuilt) as wheel: + assert set(original.namelist()) == set(wheel.namelist()) + for name in original.namelist(): + if (name.startswith("litellm/") or name.endswith("/METADATA")) and not name.endswith((".so", ".pyd")): + assert original.read(name) == wheel.read(name), name + + +def test_staging_copies_sources_without_build_artifacts(tmp_path: Path, source_repository: Path) -> None: + from scripts.build_core_distribution import SOURCES, stage_core_distribution + + source: Final = source_repository + for name in SOURCES: + (source / name).mkdir() + (source / name / "shared.txt").write_text("source payload") + (source / name / "stale.so").write_text("old native extension") + (source / name / "__pycache__").mkdir() + (source / name / "__pycache__/old.pyc").write_bytes(b"stale bytecode") + (source / "pyproject.toml").write_text('[project]\nname = "litellm"\nversion = "1.2.3"\n') + manifest: Final = source / "packaging/litellm-core/pyproject.toml" + manifest.parent.mkdir(parents=True) + manifest.write_bytes((ROOT / "packaging/litellm-core/pyproject.toml").read_bytes()) + stage: Final = tmp_path / "stage" + stage_core_distribution(source, stage) + assert all((stage / name / "shared.txt").read_text() == "source payload" for name in SOURCES) + assert not tuple(stage.rglob("*.so")) + assert not tuple(stage.rglob("*.pyc")) + assert all((source / name / "stale.so").is_file() for name in SOURCES) + + +def test_staging_preserves_gitignore_rules(tmp_path: Path, source_repository: Path) -> None: + from scripts.build_core_distribution import SOURCES, stage_core_distribution + + source: Final = source_repository + for name in SOURCES: + (source / name).mkdir() + (source / name / "shared.txt").write_text("shared payload") + (source / "pyproject.toml").write_text('[project]\nname = "litellm"\nversion = "1.2.3"\n') + manifest: Final = source / "packaging/litellm-core/pyproject.toml" + manifest.parent.mkdir(parents=True) + manifest.write_bytes((ROOT / "packaging/litellm-core/pyproject.toml").read_bytes()) + (source / ".gitignore").write_text("*.cfg\n!keep.cfg\nbuild-artifacts/\n") + (source / "litellm/.gitignore").write_text(".env\n") + (source / ".git/info/exclude").write_text("litellm/private-local.json\n") + (source / "litellm/private-local.json").write_text("synthetic packaging canary") + for name in ("tracked.cfg", "local secret.cfg", "keep.cfg", ".env"): + (source / "litellm" / name).write_text("synthetic packaging canary") + (source / "litellm/build-artifacts").mkdir() + (source / "litellm/build-artifacts/local.txt").write_text("generated artifact") + subprocess.run(["git", "add", "--force", "litellm/tracked.cfg"], cwd=source, check=True) + stage: Final = tmp_path / "stage" + stage_core_distribution(source, stage) + assert (stage / "litellm/keep.cfg").read_text() == "synthetic packaging canary" + assert (stage / "litellm/shared.txt").read_text() == "shared payload" + assert not (stage / "litellm/tracked.cfg").exists() + assert not (stage / "litellm/local secret.cfg").exists() + assert not (stage / "litellm/.env").exists() + assert not (stage / "litellm/private-local.json").exists() + assert not (stage / "litellm/build-artifacts").exists() + assert (source / "litellm/tracked.cfg").is_file() + + +def test_staging_rejects_missing_release_version(tmp_path: Path) -> None: + from scripts.build_core_distribution import stage_core_distribution + + manifest: Final = tmp_path / "packaging/litellm-core/pyproject.toml" + manifest.parent.mkdir(parents=True) + manifest.write_bytes((ROOT / "packaging/litellm-core/pyproject.toml").read_bytes()) + (tmp_path / "pyproject.toml").write_text('[project]\nname = "litellm"\n') + with pytest.raises(subprocess.CalledProcessError): + stage_core_distribution(tmp_path, tmp_path / "stage") + assert not (tmp_path / "stage").exists() diff --git a/tests/batches_tests/test_batch_rate_limits.py b/tests/batches_tests/test_batch_rate_limits.py deleted file mode 100644 index 87d35e11c74..00000000000 --- a/tests/batches_tests/test_batch_rate_limits.py +++ /dev/null @@ -1,896 +0,0 @@ -""" -Integration Tests for Batch Rate Limits -""" - -import asyncio -import json -import os - -import pytest -from fastapi import HTTPException - - -import litellm -from litellm import DualCache -from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.batch_rate_limiter import ( - BatchFileUsage, - PROXY_BatchRateLimiter, -) -from litellm.proxy.hooks.parallel_request_limiter_v3 import ( - PROXY_MaxParallelRequestsHandler_v3, -) -from litellm.proxy.utils import InternalUsageCache - - -def _build_batch_limiter() -> PROXY_BatchRateLimiter: - internal_usage_cache = InternalUsageCache(dual_cache=DualCache()) - return PROXY_BatchRateLimiter( - internal_usage_cache=internal_usage_cache, - parallel_request_limiter=PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=internal_usage_cache - ), - ) - - -def get_expected_batch_file_usage(file_path: str) -> tuple[int, int]: - """ - Helper function to calculate expected request count and token count from a batch JSONL file. - - Returns: - tuple[int, int]: (expected_request_count, expected_total_tokens) - """ - with open(file_path, "r") as f: - file_contents = [json.loads(line) for line in f if line.strip()] - - expected_request_count = len(file_contents) - expected_total_tokens = 0 - - for item in file_contents: - body = item.get("body", {}) - model = body.get("model", "") - messages = body.get("messages", []) - if messages: - item_tokens = litellm.token_counter(model=model, messages=messages) - expected_total_tokens += item_tokens - - return expected_request_count, expected_total_tokens - - -def _write_batch_file(tmp_path, file_name: str, content: str) -> str: - path = tmp_path / file_name - path.write_text(content) - return str(path) - - -@pytest.mark.asyncio() -@pytest.mark.skipif( - os.environ.get("OPENAI_API_KEY") is None, - reason="OPENAI_API_KEY not set - skipping integration test", -) -async def test_batch_rate_limits(): - """ - Integration test for batch rate limits with real OpenAI API calls. - Tests the full flow: file creation -> token counting -> cleanup - """ - litellm.turn_on_debug() - CUSTOM_LLM_PROVIDER = "openai" - BATCH_LIMITER = _build_batch_limiter() - - file_name = "openai_batch_completions.jsonl" - _current_dir = os.path.dirname(os.path.abspath(__file__)) - file_path = os.path.join(_current_dir, file_name) - - # Create file on OpenAI - print(f"Creating file from {file_path}") - file_obj = await litellm.acreate_file( - file=open(file_path, "rb"), - purpose="batch", - custom_llm_provider=CUSTOM_LLM_PROVIDER, - ) - print(f"Response from creating file: {file_obj}") - - assert file_obj.id is not None, "File ID should not be None" - - # Give API a moment to process the file - await asyncio.sleep(1) - - # Count requests and token usage in input file - tracked_batch_file_usage: BatchFileUsage = ( - await BATCH_LIMITER.count_input_file_usage( - file_id=file_obj.id, - custom_llm_provider=CUSTOM_LLM_PROVIDER, - ) - ) - print(f"Actual total tokens: {tracked_batch_file_usage.total_tokens}") - print(f"Actual request count: {tracked_batch_file_usage.request_count}") - - # Calculate expected values by reading the JSONL file - expected_request_count, expected_total_tokens = get_expected_batch_file_usage( - file_path=file_path - ) - - print(f"Expected request count: {expected_request_count}") - print(f"Expected total tokens: {expected_total_tokens}") - - # Verify token counting results - assert ( - tracked_batch_file_usage.request_count == expected_request_count - ), f"Expected {expected_request_count} requests, got {tracked_batch_file_usage.request_count}" - assert ( - tracked_batch_file_usage.total_tokens == expected_total_tokens - ), f"Expected {expected_total_tokens} total_tokens, got {tracked_batch_file_usage.total_tokens}" - - -@pytest.mark.asyncio() -async def test_batch_rate_limit_single_file(tmp_path): - """ - Test batch rate limiting with a single file. - - Key has TPM = 200 - - File with < 200 tokens: should go through - - File with > 200 tokens: should hit rate limit - """ - CUSTOM_LLM_PROVIDER = "openai" - - # Setup: Create internal usage cache and rate limiter - dual_cache = DualCache() - internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=internal_usage_cache - ) - - # Setup: Get batch rate limiter - batch_limiter = rate_limiter._get_batch_rate_limiter() - assert batch_limiter is not None, "Batch rate limiter should be available" - - # Setup: Create user API key with TPM = 200 - user_api_key_dict = UserAPIKeyAuth( - api_key="test-key-123", - tpm_limit=200, - rpm_limit=10, - ) - - # Test 1: File with < 200 tokens should go through - print("\n=== Test 1: File under 200 tokens ===") - - # Create a small batch file with ~150 tokens - small_batch_content = """{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]}} -{"custom_id": "request-2", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hi"}]}} -{"custom_id": "request-3", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hey"}]}}""" - - small_file_path = _write_batch_file( - tmp_path, "small-batch-rate-limit.jsonl", small_batch_content - ) - - try: - # Upload file to OpenAI - with open(small_file_path, "rb") as batch_file: - file_obj_small = await litellm.acreate_file( - file=batch_file, - purpose="batch", - custom_llm_provider=CUSTOM_LLM_PROVIDER, - ) - print(f"Created small file: {file_obj_small.id}") - await asyncio.sleep(1) # Give API time to process - - data_under_limit = { - "model": "gpt-3.5-turbo", - "input_file_id": file_obj_small.id, - "custom_llm_provider": CUSTOM_LLM_PROVIDER, - } - - # Should not raise an exception - result = await batch_limiter.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=dual_cache, - data=data_under_limit, - call_type="acreate_batch", - ) - print(f"✓ File with ~150 tokens passed (under limit of 200)") - print(f" Actual tokens: {result.get('_batch_token_count')}") - except HTTPException as e: - pytest.fail(f"Should not have hit rate limit with small file: {e.detail}") - - # Test 2: File with > 200 tokens should hit rate limit - print("\n=== Test 2: File over 200 tokens ===") - - # Reset cache for clean test - dual_cache = DualCache() - internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=internal_usage_cache - ) - batch_limiter = rate_limiter._get_batch_rate_limiter() - - # Create a larger batch file with ~10000+ tokens (100x larger to ensure it exceeds 200 token limit) - base_message = ( - "This is a longer message that will consume more tokens from the rate limit. " - * 100 - ) - - # Build JSONL content with json.dumps to avoid f-string nesting issues - import json as json_lib - - requests = [] - for i in range(1, 4): - request_obj = { - "custom_id": f"request-{i}", - "method": "POST", - "url": "/v1/chat/completions", - "body": { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": base_message}], - }, - } - requests.append(json_lib.dumps(request_obj)) - - large_batch_content = "\n".join(requests) - - large_file_path = _write_batch_file( - tmp_path, "large-batch-rate-limit.jsonl", large_batch_content - ) - - # Upload file to OpenAI - with open(large_file_path, "rb") as batch_file: - file_obj_large = await litellm.acreate_file( - file=batch_file, - purpose="batch", - custom_llm_provider=CUSTOM_LLM_PROVIDER, - ) - print(f"Created large file: {file_obj_large.id}") - await asyncio.sleep(1) # Give API time to process - - data_over_limit = { - "model": "gpt-3.5-turbo", - "input_file_id": file_obj_large.id, - "custom_llm_provider": CUSTOM_LLM_PROVIDER, - } - - # Should raise HTTPException with 429 status - with pytest.raises(HTTPException) as exc_info: - await batch_limiter.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=dual_cache, - data=data_over_limit, - call_type="acreate_batch", - ) - - assert exc_info.value.status_code == 429, "Should return 429 status code" - assert ( - "tokens" in exc_info.value.detail.lower() - ), "Error message should mention tokens" - print(f"✓ File with 250+ tokens correctly rejected (over limit of 200)") - print(f" Error: {exc_info.value.detail}") - - -@pytest.mark.asyncio() -async def test_batch_rate_limit_multiple_requests(tmp_path): - """ - Test batch rate limiting with multiple requests. - - Key has TPM = 200 - - Request 1: file with ~100 tokens (should go through, 100/200 used) - - Request 2: file with ~105 tokens (should hit limit, 100+105=205 > 200) - """ - CUSTOM_LLM_PROVIDER = "openai" - - # Setup: Create internal usage cache and rate limiter - dual_cache = DualCache() - internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=internal_usage_cache - ) - - # Setup: Get batch rate limiter - batch_limiter = rate_limiter._get_batch_rate_limiter() - assert batch_limiter is not None, "Batch rate limiter should be available" - - # Setup: Create user API key with TPM = 200 - user_api_key_dict = UserAPIKeyAuth( - api_key="test-key-456", - tpm_limit=200, - rpm_limit=10, - ) - - # Request 1: File with ~100 tokens - print("\n=== Request 1: File with ~100 tokens ===") - - # Create file with ~100 tokens - import json as json_lib - - message_1 = "This message has some content to reach about 100 tokens total. " * 4 - requests_1 = [] - for i in range(1, 3): - request_obj = { - "custom_id": f"request-{i}", - "method": "POST", - "url": "/v1/chat/completions", - "body": { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": message_1}], - }, - } - requests_1.append(json_lib.dumps(request_obj)) - - batch_content_1 = "\n".join(requests_1) - - file_path_1 = _write_batch_file( - tmp_path, "batch-rate-limit-request-1.jsonl", batch_content_1 - ) - - try: - # Upload file to OpenAI - with open(file_path_1, "rb") as batch_file: - file_obj_1 = await litellm.acreate_file( - file=batch_file, - purpose="batch", - custom_llm_provider=CUSTOM_LLM_PROVIDER, - ) - print(f"Created file 1: {file_obj_1.id}") - await asyncio.sleep(1) # Give API time to process - - data_request1 = { - "model": "gpt-3.5-turbo", - "input_file_id": file_obj_1.id, - "custom_llm_provider": CUSTOM_LLM_PROVIDER, - } - - # Should not raise an exception - result1 = await batch_limiter.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=dual_cache, - data=data_request1, - call_type="acreate_batch", - ) - tokens_used_1 = result1.get("_batch_token_count", 0) - print( - f"✓ Request 1 with {tokens_used_1} tokens passed ({tokens_used_1}/200 used)" - ) - except HTTPException as e: - pytest.fail(f"Request 1 should not have hit rate limit: {e.detail}") - - # Request 2: File with ~105+ tokens (total would exceed 200) - print("\n=== Request 2: File with ~105 tokens (should hit limit) ===") - - # Create file with ~105+ tokens - message_2 = ( - "This is another message with more content to exceed the remaining limit. " * 11 - ) - requests_2 = [] - for i in range(1, 3): - request_obj = { - "custom_id": f"request-{i}", - "method": "POST", - "url": "/v1/chat/completions", - "body": { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": message_2}], - }, - } - requests_2.append(json_lib.dumps(request_obj)) - - batch_content_2 = "\n".join(requests_2) - - file_path_2 = _write_batch_file( - tmp_path, "batch-rate-limit-request-2.jsonl", batch_content_2 - ) - - # Upload file to OpenAI - with open(file_path_2, "rb") as batch_file: - file_obj_2 = await litellm.acreate_file( - file=batch_file, - purpose="batch", - custom_llm_provider=CUSTOM_LLM_PROVIDER, - ) - print(f"Created file 2: {file_obj_2.id}") - await asyncio.sleep(1) # Give API time to process - - data_request2 = { - "model": "gpt-3.5-turbo", - "input_file_id": file_obj_2.id, - "custom_llm_provider": CUSTOM_LLM_PROVIDER, - } - - # Should raise HTTPException with 429 status - with pytest.raises(HTTPException) as exc_info: - await batch_limiter.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=dual_cache, - data=data_request2, - call_type="acreate_batch", - ) - - assert exc_info.value.status_code == 429, "Should return 429 status code" - assert ( - "tokens" in exc_info.value.detail.lower() - ), "Error message should mention tokens" - print(f"✓ Request 2 correctly rejected") - print(f" Error: {exc_info.value.detail}") - - -@pytest.mark.asyncio() -@pytest.mark.skipif( - os.environ.get("OPENAI_API_KEY") is None, - reason="OPENAI_API_KEY not set - skipping integration test", -) -async def test_batch_rate_limiter_with_managed_files(tmp_path): - """ - Test for GEN-2166: Verify batch rate limiter can read user files when managed files are enabled. - - This test ensures that: - 1. The batch rate limiter passes user_api_key_dict to afile_content() - 2. The managed files hook can verify file ownership correctly - 3. Rate limiting is enforced (not silently bypassed) - 4. No 403 Permission Denied errors occur for files owned by the user - """ - from unittest.mock import AsyncMock, MagicMock, patch - - CUSTOM_LLM_PROVIDER = "openai" - - # Setup: Create internal usage cache and rate limiter - dual_cache = DualCache() - internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=internal_usage_cache - ) - - # Setup: Get batch rate limiter - batch_limiter = rate_limiter._get_batch_rate_limiter() - assert batch_limiter is not None, "Batch rate limiter should be available" - - # Setup: Create user API key with TPM = 500, RPM = 10 - test_user_id = "test-user-abc123" - user_api_key_dict = UserAPIKeyAuth( - api_key="test-key-managed-files", - user_id=test_user_id, - tpm_limit=500, - rpm_limit=10, - ) - - print(f"\n=== Testing Batch Rate Limiter with Managed Files ===") - print(f"User ID: {test_user_id}") - - # Create a batch file with ~200 tokens - import json as json_lib - - message = "This is a test message for batch rate limiting with managed files. " * 5 - requests = [] - for i in range(1, 4): - request_obj = { - "custom_id": f"request-{i}", - "method": "POST", - "url": "/v1/chat/completions", - "body": { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": message}], - }, - } - requests.append(json_lib.dumps(request_obj)) - - batch_content = "\n".join(requests) - - file_path = _write_batch_file( - tmp_path, "managed-files-batch-rate-limit.jsonl", batch_content - ) - - try: - # Step 1: Upload file to OpenAI (simulating user upload) - print("\n1. Uploading batch input file...") - with open(file_path, "rb") as batch_file: - file_obj = await litellm.acreate_file( - file=batch_file, - purpose="batch", - custom_llm_provider=CUSTOM_LLM_PROVIDER, - ) - print(f" ✓ File uploaded: {file_obj.id}") - await asyncio.sleep(1) # Give API time to process - - # Step 2: Mock managed files hook to simulate file ownership check - # In a real scenario, the managed files hook would check if the user owns the file - # For this test, we'll verify that user_api_key_dict is passed correctly - print("\n2. Testing rate limiter file access with user context...") - - # Track if user_api_key_dict was passed to afile_content - original_afile_content = litellm.afile_content - user_context_passed = {"value": False} - - async def mock_afile_content(*args, **kwargs): - # Check if user_api_key_dict was passed - if ( - "user_api_key_dict" in kwargs - and kwargs["user_api_key_dict"] is not None - ): - user_context_passed["value"] = True - print(f" ✓ user_api_key_dict passed to afile_content") - print(f" User ID: {kwargs['user_api_key_dict'].user_id}") - else: - print(f" ✗ user_api_key_dict NOT passed to afile_content (BUG!)") - - # Call original function - return await original_afile_content(*args, **kwargs) - - # Patch afile_content to track the call - with patch("litellm.afile_content", side_effect=mock_afile_content): - data = { - "model": "gpt-3.5-turbo", - "input_file_id": file_obj.id, - "custom_llm_provider": CUSTOM_LLM_PROVIDER, - } - - # Step 3: Submit batch and verify rate limiting works - print("\n3. Submitting batch with rate limiting...") - result = await batch_limiter.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=dual_cache, - data=data, - call_type="acreate_batch", - ) - - tokens_used = result.get("_batch_token_count", 0) - requests_count = result.get("_batch_request_count", 0) - print(f" ✓ Batch submitted successfully") - print(f" Tokens counted: {tokens_used}") - print(f" Requests counted: {requests_count}") - print( - f" Rate limit usage: {tokens_used}/500 TPM, {requests_count}/10 RPM" - ) - - # Step 4: Verify user context was passed - print("\n4. Verifying fix for GEN-2166...") - assert user_context_passed["value"], ( - "FAILED: user_api_key_dict was not passed to afile_content(). " - "This means the bug GEN-2166 is not fixed!" - ) - print(" ✓ Fix verified: user_api_key_dict is correctly passed") - - # Step 5: Verify rate limiting is actually enforced (not bypassed) - print("\n5. Verifying rate limiting is enforced...") - assert tokens_used > 0, "Token count should be greater than 0" - assert requests_count > 0, "Request count should be greater than 0" - print(" ✓ Rate limiting is active (not silently bypassed)") - - print("\n=== Test Passed: GEN-2166 Fix Verified ===") - print("✓ Batch rate limiter can access user files") - print("✓ User context is correctly passed") - print("✓ Rate limiting is enforced") - print("✓ No silent failures") - - except HTTPException as e: - if e.status_code == 403: - pytest.fail( - f"FAILED: Got 403 Permission Denied error. " - f"This indicates the bug GEN-2166 is not fixed. " - f"Error: {e.detail}" - ) - else: - raise - except Exception as e: - pytest.fail(f"Unexpected error: {str(e)}") - - -@pytest.mark.asyncio() -async def test_batch_rate_limiter_without_user_context(tmp_path): - """ - Test that verifies the bug scenario from GEN-2166. - - When user_api_key_dict is NOT passed to count_input_file_usage(), - the function should still work for non-managed files, but would fail - for managed files (which is the bug we fixed). - - This test documents the expected behavior with and without user context. - """ - CUSTOM_LLM_PROVIDER = "openai" - - # Setup - BATCH_LIMITER = _build_batch_limiter() - - # Create a simple batch file - batch_content = """{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]}}""" - - file_path = _write_batch_file( - tmp_path, "without-user-context-batch-rate-limit.jsonl", batch_content - ) - - # Upload file - with open(file_path, "rb") as batch_file: - file_obj = await litellm.acreate_file( - file=batch_file, - purpose="batch", - custom_llm_provider=CUSTOM_LLM_PROVIDER, - ) - await asyncio.sleep(1) - - # Test 1: Without user context (old behavior - would fail with managed files) - print("\n=== Test 1: count_input_file_usage WITHOUT user context ===") - try: - usage_without_context = await BATCH_LIMITER.count_input_file_usage( - file_id=file_obj.id, - custom_llm_provider=CUSTOM_LLM_PROVIDER, - user_api_key_dict=None, # Explicitly passing None - ) - print( - f"✓ Works for non-managed files (tokens: {usage_without_context.total_tokens})" - ) - print(" Note: Would fail with 403 for managed files (GEN-2166 bug)") - except Exception as e: - print(f"✗ Failed: {str(e)}") - - # Test 2: With user context (new behavior - works with managed files) - print("\n=== Test 2: count_input_file_usage WITH user context ===") - user_api_key_dict = UserAPIKeyAuth( - api_key="test-key", - user_id="test-user-123", - ) - - usage_with_context = await BATCH_LIMITER.count_input_file_usage( - file_id=file_obj.id, - custom_llm_provider=CUSTOM_LLM_PROVIDER, - user_api_key_dict=user_api_key_dict, # Passing user context - ) - print(f"✓ Works with user context (tokens: {usage_with_context.total_tokens})") - print(" Note: This fixes GEN-2166 for managed files") - - # Verify both return the same results - assert usage_with_context.total_tokens == usage_without_context.total_tokens - assert usage_with_context.request_count == usage_without_context.request_count - print("\n✓ Both methods return identical results for non-managed files") - - -@pytest.mark.asyncio() -async def test_batch_rate_limiter_managed_files_regression(): - """ - Regression test for GEN-2166: Batch Rate Limiter Cannot Access User Files - - This test ensures that the batch rate limiter can properly access managed files - by verifying that: - 1. Managed files are detected correctly (base64 encoded unified file IDs) - 2. The _fetch_managed_file_content method uses the managed files hook - 3. User context (user_api_key_dict) is properly passed through - 4. No 403 errors occur when accessing files owned by the user - 5. The fix doesn't break non-managed file access - - This is a unit test that doesn't require external API calls. - """ - from unittest.mock import AsyncMock, MagicMock, patch - from litellm.llms.base_llm.files.transformation import BaseFileEndpoints - from litellm.types.llms.openai import HttpxBinaryResponseContent - import httpx - - print("\n=== Regression Test: GEN-2166 Batch Rate Limiter Managed Files ===") - - # Setup: Create batch rate limiter - dual_cache = DualCache() - internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) - rate_limiter = PROXY_MaxParallelRequestsHandler_v3( - internal_usage_cache=internal_usage_cache - ) - batch_limiter = rate_limiter._get_batch_rate_limiter() - assert batch_limiter is not None - - # Setup: Create user API key dict - user_api_key_dict = UserAPIKeyAuth( - api_key="test-key-regression", - user_id="test-user-regression", - tpm_limit=1000, - rpm_limit=10, - ) - - # Setup: Create mock file content (batch input file) - batch_content = b'{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Test message for regression"}]}}' - - # Mock managed file ID (base64 encoded unified file ID format) - managed_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9vY3RldC1zdHJlYW07dW5pZmllZF9pZCxyZWdyZXNzaW9uLXRlc3QtZmlsZQ==" - - # Test 1: Verify managed file detection - print("\n1. Verifying managed file detection...") - from litellm.proxy.openai_files_endpoints.common_utils import ( - is_base64_encoded_unified_file_id, - ) - - is_managed = is_base64_encoded_unified_file_id(managed_file_id) - assert is_managed, "Managed file should be detected correctly" - print(" ✓ Managed file detected") - - # Test 2: Verify _fetch_managed_file_content uses managed files hook - print("\n2. Verifying managed files hook integration...") - - # Create mock managed files hook - class MockManagedFiles(BaseFileEndpoints): - def __init__(self): - self._afile_content_called = False - self._last_call_args = None - - async def acreate_file(self, *args, **kwargs): - pass - - async def afile_content(self, *args, **kwargs): - self._afile_content_called = True - self._last_call_args = kwargs - # Return mock file content - mock_response = httpx.Response( - status_code=200, - content=batch_content, - headers={"content-type": "application/octet-stream"}, - ) - return HttpxBinaryResponseContent(response=mock_response) - - async def afile_delete(self, *args, **kwargs): - pass - - async def afile_list(self, *args, **kwargs): - pass - - async def afile_retrieve(self, *args, **kwargs): - pass - - mock_managed_files = MockManagedFiles() - mock_llm_router = MagicMock() - mock_proxy_logging_obj = MagicMock() - mock_proxy_logging_obj.get_proxy_hook.return_value = mock_managed_files - - # Patch proxy_server imports - with patch.dict( - "sys.modules", - { - "litellm.proxy.proxy_server": MagicMock( - llm_router=mock_llm_router, - proxy_logging_obj=mock_proxy_logging_obj, - ) - }, - ): - # Call _fetch_managed_file_content - result = await batch_limiter._fetch_managed_file_content( - file_id=managed_file_id, - user_api_key_dict=user_api_key_dict, - ) - - # Verify managed files hook was called - assert ( - mock_managed_files._afile_content_called - ), "REGRESSION: managed_files_obj.afile_content was not called! Bug GEN-2166 has returned." - - # Verify user context was passed - assert ( - mock_managed_files._last_call_args is not None - ), "REGRESSION: No arguments passed to afile_content" - assert ( - "file_id" in mock_managed_files._last_call_args - ), "REGRESSION: file_id not passed to managed files hook" - assert ( - mock_managed_files._last_call_args["file_id"] == managed_file_id - ), "REGRESSION: Incorrect file_id passed" - assert ( - "llm_router" in mock_managed_files._last_call_args - ), "REGRESSION: llm_router not passed to managed files hook" - - print(" ✓ Managed files hook called correctly") - print(" ✓ User context passed correctly") - - # Test 3: Verify count_input_file_usage uses managed files path - print("\n3. Verifying count_input_file_usage integration...") - - with patch.object(batch_limiter, "_fetch_managed_file_content") as mock_fetch: - mock_response = httpx.Response( - status_code=200, - content=batch_content, - headers={"content-type": "application/octet-stream"}, - ) - mock_fetch.return_value = HttpxBinaryResponseContent(response=mock_response) - - # Call count_input_file_usage with managed file - usage = await batch_limiter.count_input_file_usage( - file_id=managed_file_id, - custom_llm_provider="openai", - user_api_key_dict=user_api_key_dict, - ) - - # Verify _fetch_managed_file_content was called - assert ( - mock_fetch.called - ), "REGRESSION: _fetch_managed_file_content not called for managed files! Bug GEN-2166 has returned." - - # Verify correct parameters were passed - call_kwargs = mock_fetch.call_args.kwargs - assert ( - call_kwargs["file_id"] == managed_file_id - ), "REGRESSION: Incorrect file_id passed to _fetch_managed_file_content" - assert ( - call_kwargs["user_api_key_dict"] == user_api_key_dict - ), "REGRESSION: user_api_key_dict not passed! Bug GEN-2166 has returned." - - # Verify usage was calculated - assert usage.total_tokens > 0, "Token count should be greater than 0" - assert usage.request_count == 1, "Request count should be 1" - - print(" ✓ Managed file path used") - print(f" ✓ Token count: {usage.total_tokens}") - print(f" ✓ Request count: {usage.request_count}") - - # Test 4: Verify non-managed files still work - print("\n4. Verifying non-managed files still work...") - - non_managed_file_id = "file-abc123" # Standard OpenAI file ID - - with patch("litellm.afile_content") as mock_afile_content: - mock_response = httpx.Response( - status_code=200, - content=batch_content, - headers={"content-type": "application/octet-stream"}, - ) - mock_afile_content.return_value = HttpxBinaryResponseContent( - response=mock_response - ) - - # Call count_input_file_usage with non-managed file - usage = await batch_limiter.count_input_file_usage( - file_id=non_managed_file_id, - custom_llm_provider="openai", - user_api_key_dict=user_api_key_dict, - ) - - # Verify litellm.afile_content was called - assert ( - mock_afile_content.called - ), "REGRESSION: litellm.afile_content not called for non-managed files" - - print(" ✓ Standard file path used") - print(f" ✓ Token count: {usage.total_tokens}") - - # Test 5: Verify the fix prevents 403 errors - print("\n5. Verifying 403 error prevention...") - - # Simulate the bug scenario: managed files hook not being used - with patch.object(batch_limiter, "_fetch_managed_file_content") as mock_fetch: - # If this is NOT called for managed files, the bug has returned - mock_fetch.side_effect = Exception("Should not be called if bug exists") - - # This should call _fetch_managed_file_content - try: - with patch("litellm.afile_content") as mock_afile_content: - # If litellm.afile_content is called for managed files, bug exists - mock_afile_content.side_effect = Exception( - "Error code: 403 - User does not have access to the file" - ) - - # Reset mock_fetch to return valid content - mock_response = httpx.Response( - status_code=200, - content=batch_content, - headers={"content-type": "application/octet-stream"}, - ) - mock_fetch.side_effect = None - mock_fetch.return_value = HttpxBinaryResponseContent( - response=mock_response - ) - - # This should use _fetch_managed_file_content, not litellm.afile_content - usage = await batch_limiter.count_input_file_usage( - file_id=managed_file_id, - custom_llm_provider="openai", - user_api_key_dict=user_api_key_dict, - ) - - # Verify managed files path was used (not standard path that causes 403) - assert ( - mock_fetch.called - ), "REGRESSION: Managed files path not used! This would cause 403 errors." - assert ( - not mock_afile_content.called - ), "REGRESSION: Standard path used for managed files! This causes 403 errors." - - print(" ✓ 403 error prevention verified") - - except Exception as e: - if "403" in str(e): - pytest.fail( - f"REGRESSION: 403 error occurred! Bug GEN-2166 has returned. Error: {str(e)}" - ) - raise - - print("\n=== Regression Test Passed ===") - print("✓ Bug GEN-2166 is fixed and protected against regression") - print("✓ Managed files are properly accessed via managed files hook") - print("✓ User context is correctly passed through") - print("✓ No 403 errors occur") - print("✓ Non-managed files still work correctly\n") diff --git a/tests/batches_tests/test_openai_batches_and_files.py b/tests/batches_tests/test_openai_batches_and_files.py deleted file mode 100644 index 6341fe2f91c..00000000000 --- a/tests/batches_tests/test_openai_batches_and_files.py +++ /dev/null @@ -1,580 +0,0 @@ -# What is this? -## Unit Tests for OpenAI Batches API -import asyncio -import json -import os -import tempfile -from dotenv import load_dotenv - -load_dotenv() - -import logging -import time - -import pytest -from typing import Optional -import litellm -from litellm._logging import verbose_logger -import openai - -verbose_logger.setLevel(logging.DEBUG) - -from litellm.integrations.custom_logger import CustomLogger -from litellm.types.utils import StandardLoggingPayload -import socket -import httpx -from unittest.mock import patch, MagicMock, AsyncMock - - -def _can_resolve_openai(): - """Check if api.openai.com is reachable (DNS resolves).""" - try: - socket.getaddrinfo("api.openai.com", 443, socket.AF_UNSPEC, socket.SOCK_STREAM) - return True - except socket.gaierror: - return False - - -skip_if_no_openai_network = pytest.mark.skipif( - not _can_resolve_openai(), - reason="Cannot resolve api.openai.com - skipping integration test due to DNS issues", -) - - -async def _wait_for_standard_logging_object( - custom_logger: "TestCustomLogger", timeout: float = 15.0 -) -> StandardLoggingPayload: - from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER - - deadline = time.monotonic() + timeout - while time.monotonic() < deadline: - await GLOBAL_LOGGING_WORKER.flush() - if custom_logger.standard_logging_object is not None: - return custom_logger.standard_logging_object - await asyncio.sleep(0.25) - assert custom_logger.standard_logging_object is not None - return custom_logger.standard_logging_object - - -def load_vertex_ai_credentials(): - # Define the path to the vertex_key.json file - print("loading vertex ai credentials") - os.environ["GCS_FLUSH_INTERVAL"] = "1" - filepath = os.path.dirname(os.path.abspath(__file__)) - vertex_key_path = filepath + "/vertex_key.json" - - # Read the existing content of the file or create an empty dictionary - try: - with open(vertex_key_path, "r") as file: - # Read the file content - print("Read vertexai file path") - content = file.read() - - # If the file is empty or not valid JSON, create an empty dictionary - if not content or not content.strip(): - service_account_key_data = {} - else: - # Attempt to load the existing JSON content - file.seek(0) - service_account_key_data = json.load(file) - except FileNotFoundError: - # If the file doesn't exist, create an empty dictionary - service_account_key_data = {} - - # Update the service_account_key_data with environment variables - private_key_id = os.environ.get("VERTEX_AI_PRIVATE_KEY_ID", "") - private_key = os.environ.get("VERTEX_AI_PRIVATE_KEY", "") - private_key = private_key.replace("\\n", "\n") - service_account_key_data["private_key_id"] = private_key_id - service_account_key_data["private_key"] = private_key - - # Create a temporary file - with tempfile.NamedTemporaryFile(mode="w+", delete=False) as temp_file: - # Write the updated content to the temporary files - json.dump(service_account_key_data, temp_file, indent=2) - - # Export the temporary file as GOOGLE_APPLICATION_CREDENTIALS - os.environ["GCS_PATH_SERVICE_ACCOUNT"] = os.path.abspath(temp_file.name) - os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name) - print("created gcs path service account=", os.environ["GCS_PATH_SERVICE_ACCOUNT"]) - - -async def cancel_batch_unless_already_terminal(batch_id: str, provider: str) -> None: - try: - cancel_batch_response = await litellm.acancel_batch(batch_id=batch_id, custom_llm_provider=provider) - except openai.ConflictError as e: - if "Cannot cancel a batch with status 'completed'" in str(e): - print(f"Batch already completed, cannot cancel: {e}") - return - if "Cannot cancel a batch with status 'failed'" not in str(e): - raise - failed_batch = await litellm.aretrieve_batch(batch_id=batch_id, custom_llm_provider=provider) - print(f"Batch failed before cancel, errors={failed_batch.errors}") - failure_codes = {err.code for err in (failed_batch.errors.data if failed_batch.errors else None) or []} - assert failure_codes == {"token_limit_exceeded"}, ( - f"batch failed for a reason other than the org's enqueued token limit: {failed_batch.errors}" - ) - return - print("cancel_batch_response=", cancel_batch_response) - - -class TestCustomLogger(CustomLogger): - def __init__(self): - super().__init__() - self.standard_logging_object: Optional[StandardLoggingPayload] = None - - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - print( - "Success event logged with kwargs=", - kwargs, - "and response_obj=", - response_obj, - ) - self.standard_logging_object = kwargs["standard_logging_object"] - - -def cleanup_azure_files(): - """ - Delete all files for Azure - helper for when we run out of Azure Files Quota - """ - azure_files = litellm.file_list( - custom_llm_provider="azure", - api_key=os.getenv("AZURE_FT_API_KEY"), - api_base=os.getenv("AZURE_FT_API_BASE"), - ) - print("azure_files=", azure_files) - for _file in azure_files: - print("deleting file=", _file) - delete_file_response = litellm.file_delete( - file_id=_file.id, - custom_llm_provider="azure", - api_key=os.getenv("AZURE_FT_API_KEY"), - api_base=os.getenv("AZURE_FT_API_BASE"), - ) - print("delete_file_response=", delete_file_response) - assert delete_file_response.id == _file.id - - -def cleanup_azure_ft_models(): - """ - Test CLEANUP: Delete all existing fine tuning jobs for Azure - """ - try: - from openai import AzureOpenAI - import requests - - client = AzureOpenAI( - api_key=os.getenv("AZURE_AI_API_KEY"), - azure_endpoint=os.getenv("AZURE_AI_API_BASE"), - api_version=os.getenv("AZURE_AI_API_VERSION"), - ) - - _list_ft_jobs = client.fine_tuning.jobs.list() - print("_list_ft_jobs=", _list_ft_jobs) - - # delete all ft jobs make post request to this - # Delete all fine-tuning jobs - for job in _list_ft_jobs: - try: - endpoint = os.getenv("AZURE_FT_API_BASE").rstrip("/") - url = f"{endpoint}/openai/fine_tuning/jobs/{job.id}?api-version=2024-10-21" - print("url=", url) - - headers = { - "api-key": os.getenv("AZURE_FT_API_KEY"), - "Content-Type": "application/json", - } - - response = requests.delete(url, headers=headers) - print(f"Deleting job {job.id}: Status {response.status_code}") - if response.status_code != 204: - print(f"Error deleting job {job.id}: {response.text}") - - except Exception as e: - print(f"Error deleting job {job.id}: {str(e)}") - except Exception as e: - print(f"Error on cleanup_azure_ft_models: {str(e)}") - - -@pytest.mark.parametrize("provider", ["openai"]) -@pytest.mark.asyncio() -@skip_if_no_openai_network -async def test_async_create_batch(provider, tmp_path): - """ - 1. Create File for Batch completion - 2. Create Batch Request - 3. Retrieve the specific batch - """ - litellm.turn_on_debug() - print("Testing async create batch") - litellm.logging_callback_manager._reset_all_callbacks() - - file_name = "openai_batch_completions.jsonl" - _current_dir = os.path.dirname(os.path.abspath(__file__)) - file_path = os.path.join(_current_dir, file_name) - with open(file_path, "rb") as batch_file: - file_obj = await litellm.acreate_file( - file=batch_file, - purpose="batch", - custom_llm_provider=provider, - ) - print("Response from creating file=", file_obj) - - await asyncio.sleep(10) - batch_input_file_id = file_obj.id - assert ( - batch_input_file_id is not None - ), "Failed to create file, expected a non null file_id but got {batch_input_file_id}" - - extra_metadata_field = { - "user_api_key_alias": "special_api_key_alias", - "user_api_key_team_alias": "special_team_alias", - } - custom_logger = TestCustomLogger() - litellm.callbacks = [custom_logger, "datadog"] - create_batch_response = await litellm.acreate_batch( - completion_window="24h", - endpoint="/v1/chat/completions", - input_file_id=batch_input_file_id, - custom_llm_provider=provider, - metadata={"key1": "value1", "key2": "value2"}, - # litellm specific param - used for logging metadata on logging callback - litellm_metadata=extra_metadata_field, - ) - - print("response from litellm.create_batch=", create_batch_response) - - assert ( - create_batch_response.id is not None - ), f"Failed to create batch, expected a non null batch_id but got {create_batch_response.id}" - assert ( - create_batch_response.endpoint == "/v1/chat/completions" - or create_batch_response.endpoint == "/chat/completions" - ), f"Failed to create batch, expected endpoint to be /v1/chat/completions but got {create_batch_response.endpoint}" - assert ( - create_batch_response.input_file_id == batch_input_file_id - ), f"Failed to create batch, expected input_file_id to be {batch_input_file_id} but got {create_batch_response.input_file_id}" - - # Assert that the create batch event is logged on CustomLogger - standard_logging_object = await _wait_for_standard_logging_object(custom_logger) - print( - "standard_logging_object=", - json.dumps(standard_logging_object, indent=4, default=str), - ) - assert ( - standard_logging_object["metadata"]["user_api_key_alias"] - == extra_metadata_field["user_api_key_alias"] - ) - assert ( - standard_logging_object["metadata"]["user_api_key_team_alias"] - == extra_metadata_field["user_api_key_team_alias"] - ) - - retrieved_batch = await litellm.aretrieve_batch( - batch_id=create_batch_response.id, custom_llm_provider=provider - ) - print("retrieved batch=", retrieved_batch) - # just assert that we retrieved a non None batch - - assert retrieved_batch.id == create_batch_response.id - - # list all batches - list_batches = await litellm.alist_batches(custom_llm_provider=provider, limit=2) - print("list_batches=", list_batches) - - # try to get file content for our original file - - file_content = await litellm.afile_content( - file_id=batch_input_file_id, custom_llm_provider=provider - ) - - print("file content = ", file_content) - - # file obj - file_obj = await litellm.afile_retrieve( - file_id=batch_input_file_id, custom_llm_provider=provider - ) - print("file obj = ", file_obj) - assert file_obj.id == batch_input_file_id - - # delete file - delete_file_response = await litellm.afile_delete( - file_id=batch_input_file_id, custom_llm_provider=provider - ) - - print("delete file response = ", delete_file_response) - - assert delete_file_response.id == batch_input_file_id - - all_files_list = await litellm.afile_list( - custom_llm_provider=provider, - ) - - print("all_files_list = ", all_files_list) - - result_file_path = tmp_path / "batch_job_results_furniture.jsonl" - result_file_path.write_bytes(file_content.content) - - await cancel_batch_unless_already_terminal(batch_id=create_batch_response.id, provider=provider) - - -mock_file_response = { - "kind": "storage#object", - "id": "litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb/1739598666670574", - "selfLink": "https://www.googleapis.com/storage/v1/b/litellm-local/o/litellm-vertex-files%2Fpublishers%2Fgoogle%2Fmodels%2Fgemini-1.5-flash-001%2F5f7b99ad-9203-4430-98bf-3b45451af4cb", - "mediaLink": "https://storage.googleapis.com/download/storage/v1/b/litellm-local/o/litellm-vertex-files%2Fpublishers%2Fgoogle%2Fmodels%2Fgemini-1.5-flash-001%2F5f7b99ad-9203-4430-98bf-3b45451af4cb?generation=1739598666670574&alt=media", - "name": "litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb", - "bucket": "litellm-local", - "generation": "1739598666670574", - "metageneration": "1", - "contentType": "application/json", - "storageClass": "STANDARD", - "size": "416", - "md5Hash": "hbBNj7C8KJ7oVH+JmyRM6A==", - "crc32c": "oDmiUA==", - "etag": "CO7D0IT+xIsDEAE=", - "timeCreated": "2025-02-15T05:51:06.741Z", - "updated": "2025-02-15T05:51:06.741Z", - "timeStorageClassUpdated": "2025-02-15T05:51:06.741Z", - "timeFinalized": "2025-02-15T05:51:06.741Z", -} - -mock_vertex_batch_response = { - "name": "projects/123456789/locations/us-central1/batchPredictionJobs/test-batch-id-456", - "displayName": "litellm_batch_job", - "model": "projects/123456789/locations/us-central1/models/gemini-1.5-flash-001", - "modelVersionId": "v1", - "inputConfig": { - "gcsSource": { - "uris": [ - "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb" - ] - } - }, - "outputConfig": { - "gcsDestination": {"outputUriPrefix": "gs://litellm-local/batch-outputs/"} - }, - "dedicatedResources": { - "machineSpec": { - "machineType": "n1-standard-4", - "acceleratorType": "NVIDIA_TESLA_T4", - "acceleratorCount": 1, - }, - "startingReplicaCount": 1, - "maxReplicaCount": 1, - }, - "state": "JOB_STATE_RUNNING", - "createTime": "2025-02-15T05:51:06.741Z", - "startTime": "2025-02-15T05:51:07.741Z", - "updateTime": "2025-02-15T05:51:08.741Z", - "labels": {"key1": "value1", "key2": "value2"}, - "completionStats": {"successfulCount": 0, "failedCount": 0, "remainingCount": 100}, -} - -mock_vertex_list_response = { - "batchPredictionJobs": [ - mock_vertex_batch_response, - { - **mock_vertex_batch_response, - "name": "projects/123456789/locations/us-central1/batchPredictionJobs/test-batch-id-789", - "state": "JOB_STATE_SUCCEEDED", - }, - ], - "nextPageToken": "", -} - - -@pytest.mark.asyncio -async def test_avertex_batch_prediction(monkeypatch): - monkeypatch.setenv("GCS_BUCKET_NAME", "litellm-local") - monkeypatch.setenv("VERTEXAI_PROJECT", "mock-project") - monkeypatch.setenv("VERTEXAI_LOCATION", "us-central1") - - # Mock Google auth so the test doesn't need real credentials - mock_creds = MagicMock() - mock_creds.token = "mock-token" - mock_creds.valid = True - mock_creds.expiry = None - monkeypatch.setattr( - "google.auth.default", - lambda *args, **kwargs: (mock_creds, "mock-project"), - ) - - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - - # Configure mock response object - mock_response = MagicMock() - mock_response.raise_for_status.return_value = None - - async def mock_side_effect(*args, **kwargs): - print("args", args, "kwargs", kwargs) - url = kwargs.get("url", "") - if "files" in url: - mock_response.json.return_value = mock_file_response - elif "batch" in url: - mock_response.json.return_value = mock_vertex_batch_response - mock_response.status_code = 200 - return mock_response - - # Batch jsonl creation now stages the body to a temp file and issues a single - # uploadType=media POST against the raw httpx.AsyncClient (client.client) inside - # _astage_and_upload_media, not AsyncHTTPHandler.post. Patch that raw POST so the - # real staging/upload + response transform run while the GCS object response is - # mocked; AsyncHTTPHandler.post still handles the batch-prediction call. - with ( - patch( - "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", - side_effect=mock_side_effect, - ), - patch.object( - httpx.AsyncClient, - "post", - new_callable=AsyncMock, - return_value=httpx.Response( - 200, - json=mock_file_response, - request=httpx.Request("POST", "https://storage.googleapis.com/upload"), - ), - ) as mock_gcs_upload, - ): - litellm.set_verbose = True - litellm.turn_on_debug() - file_name = "vertex_batch_completions.jsonl" - _current_dir = os.path.dirname(os.path.abspath(__file__)) - file_path = os.path.join(_current_dir, file_name) - - # Create file - file_obj = await litellm.acreate_file( - file=open(file_path, "rb"), - purpose="batch", - custom_llm_provider="vertex_ai", - ) - print("Response from creating file=", file_obj) - - assert ( - file_obj.id - == "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb" - ) - - mock_gcs_upload.assert_awaited_once() - upload_url = str(mock_gcs_upload.call_args.args[0]) - assert "uploadType=media" in upload_url - assert "/b/litellm-local/o" in upload_url - assert ( - mock_gcs_upload.call_args.kwargs["headers"]["Content-Type"] - == "application/json" - ) - - # Create batch - create_batch_response = await litellm.acreate_batch( - completion_window="24h", - endpoint="/v1/chat/completions", - input_file_id=file_obj.id, - custom_llm_provider="vertex_ai", - metadata={"key1": "value1", "key2": "value2"}, - ) - print("create_batch_response=", create_batch_response) - - assert create_batch_response.id == "test-batch-id-456" - assert ( - create_batch_response.input_file_id - == "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb" - ) - - # Mock the retrieve batch response - with patch( - "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get" - ) as mock_get: - mock_get_response = MagicMock() - mock_get_response.json.return_value = mock_vertex_batch_response - mock_get_response.status_code = 200 - mock_get_response.is_redirect = False - mock_get_response.raise_for_status.return_value = None - mock_get_response.is_redirect = False - mock_get.return_value = mock_get_response - - retrieved_batch = await litellm.aretrieve_batch( - batch_id=create_batch_response.id, - custom_llm_provider="vertex_ai", - ) - print("retrieved_batch=", retrieved_batch) - - assert retrieved_batch.id == "test-batch-id-456" - - -@pytest.mark.asyncio - - -@pytest.mark.asyncio - - -@pytest.mark.asyncio -@skip_if_no_openai_network -async def test_delete_batch_output_file(): - """ - Test that deleting a batch output file works correctly. - - This test verifies the fix for: - - When a batch is retrieved and has an output_file_id, the file object is properly stored - - The output file can be deleted without validation errors - - The file_object is fetched and stored with proper metadata instead of None - """ - litellm.turn_on_debug() - print("Testing delete batch output file") - - file_name = "openai_batch_completions.jsonl" - _current_dir = os.path.dirname(os.path.abspath(__file__)) - file_path = os.path.join(_current_dir, file_name) - - # Create file for batch - file_obj = await litellm.acreate_file( - file=open(file_path, "rb"), - purpose="batch", - custom_llm_provider="openai", - ) - print("Response from creating file=", file_obj) - batch_input_file_id = file_obj.id - - # Create batch - create_batch_response = await litellm.acreate_batch( - completion_window="24h", - endpoint="/v1/chat/completions", - input_file_id=batch_input_file_id, - custom_llm_provider="openai", - ) - print("Batch created with ID=", create_batch_response.id) - - # Retrieve batch to get output_file_id - retrieved_batch = await litellm.aretrieve_batch( - batch_id=create_batch_response.id, custom_llm_provider="openai" - ) - print("Retrieved batch=", retrieved_batch) - - # If batch has completed and has output file, test deleting it - if retrieved_batch.output_file_id: - print(f"Testing deletion of output file: {retrieved_batch.output_file_id}") - - # This is the key test - deleting the output file should work - # without validation errors (file_object should not be None) - delete_output_file_response = await litellm.afile_delete( - file_id=retrieved_batch.output_file_id, custom_llm_provider="openai" - ) - - print("Delete output file response=", delete_output_file_response) - assert delete_output_file_response.id == retrieved_batch.output_file_id - assert delete_output_file_response.deleted is True or hasattr( - delete_output_file_response, "id" - ) - print("✓ Successfully deleted batch output file") - else: - print( - "⚠ Batch has not completed yet or no output file available, skipping output file deletion test" - ) - - # Clean up - delete the input file - delete_input_file_response = await litellm.afile_delete( - file_id=batch_input_file_id, custom_llm_provider="openai" - ) - print("Delete input file response=", delete_input_file_response) - assert delete_input_file_response.id == batch_input_file_id - print("✓ Successfully deleted batch input file") diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index cbc09edc357..9b9206194b4 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -110,6 +110,8 @@ ignored_function_names = [ "_aanthropic_messages_yield_recovered", # Tested through every mid-stream retry and fallback test in test_router.py "_anthropic_messages_policy_retries", # Tested through the retry budget precedence test in test_router.py "_get_wildcard_deployments", # Tested through the get_model_list_of_routed_group wildcard test in test_router.py + "_is_fallback_hop", # Tested through the order fallback hop tests in test_router_order_fallback.py + "_deployment_that_just_failed", # Tested through the same-boundary hop test in test_router_order_fallback.py ] diff --git a/tests/code_coverage_tests/test_unit_passed_gate.py b/tests/code_coverage_tests/test_unit_passed_gate.py new file mode 100644 index 00000000000..7b511538d40 --- /dev/null +++ b/tests/code_coverage_tests/test_unit_passed_gate.py @@ -0,0 +1,68 @@ +import json +import os +import subprocess +from pathlib import Path +from typing import Final + +import pytest +import yaml + +_UNIT_WORKFLOW: Final = Path(__file__).resolve().parents[2] / ".github" / "workflows" / "test-unit.yml" +_GATE_JOB: Final = "unit-passed" + + +def _jobs() -> dict[str, dict[str, object]]: + return yaml.safe_load(_UNIT_WORKFLOW.read_text())["jobs"] + + +def _gate_script() -> str: + steps: Final = _jobs()[_GATE_JOB]["steps"] + assert isinstance(steps, list) and len(steps) == 1 + return steps[0]["run"] + + +def _run_gate(results: dict[str, str]) -> subprocess.CompletedProcess[str]: + needs: Final = {job: {"result": result, "outputs": {}} for job, result in results.items()} + return subprocess.run( + ("bash", "--noprofile", "--norc", "-eo", "pipefail", "-c", _gate_script()), + env={**os.environ, "NEEDS": json.dumps(needs)}, + capture_output=True, + text=True, + timeout=30, + check=False, + ) + + +def _needed_jobs() -> tuple[str, ...]: + needs: Final = _jobs()[_GATE_JOB]["needs"] + assert isinstance(needs, list) + return tuple(needs) + + +def test_the_gate_passes_when_every_needed_job_succeeded() -> None: + jobs: Final = _needed_jobs() + assert jobs + + result: Final = _run_gate(dict.fromkeys(jobs, "success")) + + assert result.returncode == 0, result.stdout + result.stderr + assert all(f"{job}: success" in result.stdout for job in jobs), result.stdout + + +@pytest.mark.parametrize("outcome", ("failure", "cancelled", "skipped")) +def test_the_gate_fails_when_any_needed_job_did_not_succeed(outcome: str) -> None: + jobs: Final = _needed_jobs() + assert len(jobs) > 1 + + result: Final = _run_gate({**dict.fromkeys(jobs, "success"), jobs[-1]: outcome}) + + assert result.returncode != 0 + assert f"{jobs[-1]}: {outcome}" in result.stdout, result.stdout + + +def test_the_gate_waits_for_every_other_job_and_reports_even_when_they_fail() -> None: + jobs: Final = _jobs() + gate: Final = jobs[_GATE_JOB] + + assert set(_needed_jobs()) == set(jobs) - {_GATE_JOB} + assert gate["if"] == "always()" diff --git a/tests/code_coverage_tests/unbounded_in_baseline.txt b/tests/code_coverage_tests/unbounded_in_baseline.txt index 77c7c362f39..da56a25542d 100644 --- a/tests/code_coverage_tests/unbounded_in_baseline.txt +++ b/tests/code_coverage_tests/unbounded_in_baseline.txt @@ -1,6 +1,5 @@ # Grandfathered findings of check_unbounded_in_lists.py: path::scope::kind::subject::occurrence. # Fix a site and delete its line; regenerate with `check_unbounded_in_lists.py --update-baseline`. -enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py CheckResponsesCost.check_responses_cost prisma id.in `[job.id for job in completed_jobs]` 0 litellm/integrations/shadow_eval_logger.py ShadowEvalLogger._active_jobs prisma job_id.in `[str(record.id) for record in records]` 0 litellm/llms/litellm_proxy/skills/handler.py LiteLLMSkillsHandler.list_skills prisma created_by.in `owner_scopes` 0 litellm/proxy/_experimental/mcp_server/db.py get_mcp_servers prisma server_id.in `server_ids` 0 diff --git a/tests/e2e/llm_translation/responses_helpers.py b/tests/e2e/llm_translation/responses_helpers.py new file mode 100644 index 00000000000..be0fac74bb6 --- /dev/null +++ b/tests/e2e/llm_translation/responses_helpers.py @@ -0,0 +1,17 @@ +from __future__ import annotations + +from typing import Final + +from models import LiteLLMParamsBody + +AZURE_OPENAI_BACKEND: Final = "azure/gpt-5.4-nano" +AZURE_OPENAI_API_VERSION: Final = "v1" + + +def azure_openai_params(api_version: str = AZURE_OPENAI_API_VERSION) -> LiteLLMParamsBody: + return LiteLLMParamsBody( + model=AZURE_OPENAI_BACKEND, + api_base="os.environ/AZURE_API_BASE", + api_key="os.environ/AZURE_API_KEY", + api_version=api_version, + ) diff --git a/tests/e2e/llm_translation/test_messages_count_tokens_bedrock_e2e.py b/tests/e2e/llm_translation/test_messages_count_tokens_bedrock_e2e.py new file mode 100644 index 00000000000..ee5cd54e4e0 --- /dev/null +++ b/tests/e2e/llm_translation/test_messages_count_tokens_bedrock_e2e.py @@ -0,0 +1,70 @@ +"""Live e2e: `/v1/messages/count_tokens` on a Bedrock Claude model that bedrock-runtime +cannot count. + +Claude Opus 4.8 is offered only through cross-region inference, and bedrock-runtime's +CountTokens answers 400 for it. The proxy then has to count through bedrock-mantle's +Anthropic count_tokens, and the answer must sit within a few percent of what `/v1/messages` +bills as `usage.input_tokens`. The local tokenizer fallback undercounts these models by +about 40%, so this is the line that proves the real count is served. Both calls go through +the real Anthropic SDK, the client customers count with +""" + +from __future__ import annotations + +from typing import Final + +import pytest +from anthropic.types import MessageParam +from e2e_config import unique_marker +from e2e_metadata import Domain, Provider, Route, Subject, meta +from lifecycle import ResourceManager +from models import LiteLLMParamsBody +from proxy_client import ProxyClient +from sdk_clients import NO_PROXY_CACHE, SdkClients + +pytestmark = pytest.mark.e2e + +CROSS_REGION_ONLY_CLAUDE_BACKEND: Final = "bedrock/global.anthropic.claude-opus-4-8" +COUNT_TOLERANCE: Final = 0.05 + + +def _register(proxy: ProxyClient, resources: ResourceManager, backend: str) -> tuple[str, str]: + model: Final = f"e2e-count-tokens-bedrock-{unique_marker()}" + model_id: Final = proxy.create_model( + model, + LiteLLMParamsBody( + model=backend, + 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", + ), + ) + resources.defer(lambda: proxy.delete_model(model_id)) + return model, resources.key() + + +class TestBedrockMessagesCountTokens: + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.COUNT_TOKENS, + providers=(Provider.BEDROCK,), + models=(CROSS_REGION_ONLY_CLAUDE_BACKEND,), + ) + ) + def test_count_matches_billed_input_tokens_for_a_model_bedrock_runtime_cannot_count( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model, key = _register(proxy, resources, CROSS_REGION_ONLY_CLAUDE_BACKEND) + client: Final = sdk.anthropic(key) + prompt: Final = f"{unique_marker()} " + "The quick brown fox jumps over the lazy dog. " * 40 + message: Final[MessageParam] = {"role": "user", "content": prompt} + + counted: Final = client.messages.count_tokens(model=model, messages=[message]) + answered: Final = client.messages.create( + model=model, max_tokens=1, messages=[message], extra_body=NO_PROXY_CACHE + ) + + billed: Final = answered.usage.input_tokens + assert billed, answered.usage + assert abs(counted.input_tokens - billed) <= billed * COUNT_TOLERANCE, (counted, answered.usage) diff --git a/tests/e2e/llm_translation/test_responses_e2e.py b/tests/e2e/llm_translation/test_responses_e2e.py index 00a1d3611f9..5a78f77ebca 100644 --- a/tests/e2e/llm_translation/test_responses_e2e.py +++ b/tests/e2e/llm_translation/test_responses_e2e.py @@ -42,6 +42,7 @@ from provider_edge import LiveEdge, start_provider_edge from provider_edge_bedrock import bedrock_signer from proxy_client import ProxyClient from pydantic import BaseModel, TypeAdapter +from responses_helpers import AZURE_OPENAI_BACKEND, azure_openai_params from sdk_clients import NO_PROXY_CACHE, SdkClients pytestmark = pytest.mark.e2e @@ -58,8 +59,8 @@ OPENAI_VISION_BACKEND: Final = "openai/gpt-4o" ANTHROPIC_BACKEND: Final = "anthropic/claude-haiku-4-5" BEDROCK_CONVERSE_BACKEND: Final = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" VERTEX_BACKEND: Final = "vertex_ai/gemini-2.5-flash" -AZURE_OPENAI_BACKEND: Final = "azure/gpt-5.4-nano" -AZURE_OPENAI_API_VERSION: Final = "v1" +GEMINI_BACKEND: Final = "gemini/gemini-2.5-flash" +OPENAI_RESPONSES_BACKEND: Final = "openai/gpt-5.5" INSTRUCTIONS = "You are a helpful assistant" CAT_IMAGE_URL = "https://upload.wikimedia.org/wikipedia/commons/3/3a/Cat03.jpg" BEDROCK_EDGE_REGION: Final = "us-east-1" @@ -102,11 +103,36 @@ WEATHER_TOOL: FunctionToolParam = { "strict": False, } +LOCATIONS_TOOL: Final[FunctionToolParam] = { + "type": "function", + "name": "get_locations", + "description": "Return locations that need weather information", + "parameters": { + "type": "object", + "properties": {"locations": {"type": "array", "items": {"type": "string"}}}, + "required": ["locations"], + "additionalProperties": False, + }, + "strict": True, +} + + +class LocationsArguments(BaseModel): + locations: list[str] + + +class ResponseUsageCost(BaseModel): + cost: float | None = None + def _openai_params() -> LiteLLMParamsBody: return LiteLLMParamsBody(model=OPENAI_MINI_BACKEND, api_key="os.environ/OPENAI_API_KEY") +def _openai_responses_params() -> LiteLLMParamsBody: + return LiteLLMParamsBody(model=OPENAI_RESPONSES_BACKEND, api_key="os.environ/OPENAI_API_KEY") + + def _anthropic_params() -> LiteLLMParamsBody: return LiteLLMParamsBody(model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY") @@ -128,13 +154,8 @@ def _vertex_params() -> LiteLLMParamsBody: ) -def _azure_openai_params() -> LiteLLMParamsBody: - return LiteLLMParamsBody( - model=AZURE_OPENAI_BACKEND, - api_base="os.environ/AZURE_API_BASE", - api_key="os.environ/AZURE_API_KEY", - api_version=AZURE_OPENAI_API_VERSION, - ) +def _gemini_params() -> LiteLLMParamsBody: + return LiteLLMParamsBody(model=GEMINI_BACKEND, api_key="os.environ/GEMINI_API_KEY") def _register( @@ -197,22 +218,147 @@ class TestResponses: def test_responses_streaming_returns_completion( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: - model = _register(proxy, resources, _openai_params()) - client = sdk.openai(resources.key()) + model: Final = _register(proxy, resources, _openai_params()) + client: Final = sdk.openai(resources.key()) - stream = client.responses.create( + stream: Final = client.responses.create( model=model, input="reply with one word", instructions=INSTRUCTIONS, stream=True, extra_body=NO_PROXY_CACHE, ) - events = tuple(stream) + events: Final = tuple(stream) assert events, "responses stream returned no events" - deltas = tuple(event.delta for event in events if event.type == "response.output_text.delta") + deltas: Final = tuple(event.delta for event in events if event.type == "response.output_text.delta") assert any(delta for delta in deltas), "responses stream returned no text deltas" - assert events[-1].type == "response.completed", ( - f"responses stream did not terminate with response.completed: {events[-1].type}" + completed: Final = events[-1] + assert isinstance(completed, ResponseCompletedEvent), ( + f"responses stream did not terminate with response.completed: {completed.type}" + ) + usage: Final = completed.response.usage + assert usage is not None, f"response.completed had no usage: {completed.response!r}" + assert usage.input_tokens > 0, f"response.completed had no input tokens: {usage!r}" + assert usage.output_tokens > 0, f"response.completed had no output tokens: {usage!r}" + assert usage.total_tokens == usage.input_tokens + usage.output_tokens, ( + f"response.completed token totals were inconsistent: {usage!r}" + ) + usage_cost: Final = TypeAdapter(ResponseUsageCost).validate_python( + cast(object, usage.model_extra if usage.model_extra is not None else {}) + ) + assert usage_cost.cost is not None, f"response.completed usage had no cost: {usage.model_extra!r}" + assert usage_cost.cost > 0, f"response.completed cost was not positive: {usage_cost.cost}" + + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.GEMINI,), + models=(GEMINI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) + def test_responses_gemini_returns_completion( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model: Final = _register(proxy, resources, _gemini_params(), prefix="e2e-responses-gemini") + client: Final = sdk.openai(resources.key()) + + response: Final = client.responses.create( + model=model, input="reply with one word", instructions=INSTRUCTIONS, extra_body=NO_PROXY_CACHE + ) + assert response.output_text.strip(), f"/responses over gemini returned no output text: {response.output!r}" + + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.GEMINI,), + models=(GEMINI_BACKEND,), + mode=Mode.STREAM, + ) + ) + def test_responses_gemini_streaming_returns_completion( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model: Final = _register(proxy, resources, _gemini_params(), prefix="e2e-responses-gemini") + client: Final = sdk.openai(resources.key()) + + stream: Final = client.responses.create( + model=model, + input="reply with one word", + instructions=INSTRUCTIONS, + stream=True, + extra_body=NO_PROXY_CACHE, + ) + events: Final = tuple(stream) + deltas: Final = tuple(event.delta for event in events if event.type == "response.output_text.delta") + assert any(deltas), "responses stream over gemini returned no text deltas" + assert isinstance(events[-1], ResponseCompletedEvent), ( + f"responses stream over gemini did not end with response.completed: {events[-1].type}" + ) + + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.GEMINI,), + models=(GEMINI_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) + def test_responses_gemini_replays_legacy_function_call_output( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model: Final = _register(proxy, resources, _gemini_params(), prefix="e2e-responses-gemini-tool") + client: Final = sdk.openai(resources.key()) + function_call_id: Final = f"fc_{unique_marker()}" + input_items: Final[ResponseInputParam] = [ + { + "type": "message", + "role": "user", + "content": "What is the temperature in Paris today?", + }, + { + "type": "function_call", + "arguments": '{"location": "Paris, France"}', + "call_id": function_call_id, + "name": "get_temperature", + "id": function_call_id, + "status": "completed", + }, + { + "type": "function_call_output", + "call_id": function_call_id, + "output": "Temperature is exactly 31 Celsius.", + }, + ] + tools: Final[tuple[FunctionToolParam, ...]] = ( + { + "type": "function", + "name": "get_temperature", + "description": "Get the current temperature for a location", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string"}}, + "required": ["location"], + "additionalProperties": False, + }, + "strict": False, + }, + ) + + response: Final = client.responses.create( + model=model, + input=input_items, + tools=tools, + store=False, + extra_body=NO_PROXY_CACHE, + ) + assert response.status == "completed", f"legacy tool replay was not completed: {response.status}" + assert "31" in response.output_text, ( + f"legacy tool result was missing from output text: {response.output_text!r}" ) @pytest.mark.covers("llm.responses.openai.basic.nonstream.cost_logged") @@ -361,6 +507,91 @@ class TestResponses: ) _assert_weather_call(response) + @pytest.mark.covers("llm.responses.anthropic.tool_use.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) + def test_responses_anthropic_strict_array_schema_tool_call( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model: Final = _register(proxy, resources, _anthropic_params()) + client: Final = sdk.openai(resources.key()) + + response: Final = client.responses.create( + model=model, + input="Find the weather locations for Tokyo and Paris using get_locations.", + instructions=INSTRUCTIONS, + tools=[LOCATIONS_TOOL], + tool_choice="required", + extra_body=NO_PROXY_CACHE, + ) + function_call: Final = next( + (call for call in _function_calls(response) if call.name == "get_locations"), + None, + ) + assert function_call is not None, f"response had no get_locations call: {response.output!r}" + arguments: Final = LocationsArguments.model_validate_json(function_call.arguments) + assert arguments.locations, f"get_locations returned no locations: {function_call.arguments}" + + @pytest.mark.covers("llm.responses.anthropic.multi_turn.nonstream.works") + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.ANTHROPIC,), + models=(ANTHROPIC_BACKEND,), + capabilities=(Capability.FUNCTION_CALLING,), + mode=Mode.NONSTREAM, + ) + ) + def test_responses_anthropic_tool_output_continues_with_previous_response_id( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model: Final = _register(proxy, resources, _anthropic_params()) + client: Final = sdk.openai(resources.key()) + + first: Final = client.responses.create( + model=model, + input="Find the weather locations for Tokyo and Paris using get_locations.", + instructions=INSTRUCTIONS, + tools=[LOCATIONS_TOOL], + tool_choice="required", + extra_body=NO_PROXY_CACHE, + ) + function_call: Final = next( + (call for call in _function_calls(first) if call.name == "get_locations"), + None, + ) + assert function_call is not None, f"response had no get_locations call: {first.output!r}" + assert function_call.call_id, f"get_locations call had no call_id: {function_call!r}" + arguments: Final = LocationsArguments.model_validate_json(function_call.arguments) + assert arguments.locations, f"get_locations call had no locations: {function_call.arguments}" + + tool_result: Final = "Distinctive forecast: 47 degrees Celsius" + follow_up_input: Final[ResponseInputParam] = [ + { + "type": "function_call_output", + "call_id": function_call.call_id, + "output": tool_result, + } + ] + second: Final = client.responses.create( + model=model, + previous_response_id=first.id, + input=follow_up_input, + instructions=INSTRUCTIONS, + tools=[LOCATIONS_TOOL], + extra_body=NO_PROXY_CACHE, + ) + assert "47" in second.output_text, f"follow-up omitted tool result: {second.output_text!r}" + @pytest.mark.covers("llm.responses.bedrock_converse.basic.nonstream.works") @meta( Subject( @@ -469,7 +700,7 @@ class TestResponses: def test_responses_azure_openai_returns_completion( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: - model = _register(proxy, resources, _azure_openai_params(), prefix="e2e-responses-azure-openai") + model = _register(proxy, resources, azure_openai_params(), prefix="e2e-responses-azure-openai") client = sdk.openai(resources.key()) response = client.responses.create( @@ -479,6 +710,150 @@ class TestResponses: f"/responses over azure openai returned no output text: {response.output!r}" ) + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.AZURE,), + models=(AZURE_OPENAI_BACKEND,), + mode=Mode.STREAM, + ) + ) + def test_responses_azure_openai_streaming_returns_completion( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model: Final = _register(proxy, resources, azure_openai_params(), prefix="e2e-responses-azure-stream") + client: Final = sdk.openai(resources.key()) + + stream: Final = client.responses.create( + model=model, + input="reply with one word", + instructions=INSTRUCTIONS, + stream=True, + extra_body=NO_PROXY_CACHE, + ) + events: Final = tuple(stream) + deltas: Final = tuple(event.delta for event in events if event.type == "response.output_text.delta") + assert any(deltas), "responses stream over azure openai returned no text deltas" + assert isinstance(events[-1], ResponseCompletedEvent), ( + f"responses stream over azure openai did not end with response.completed: {events[-1].type}" + ) + + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.AZURE,), + models=(AZURE_OPENAI_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) + def test_responses_azure_openai_preview_api_version_accepts_truncation( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model: Final = _register( + proxy, + resources, + azure_openai_params(api_version="preview"), + prefix="e2e-responses-azure-preview", + ) + client: Final = sdk.openai(resources.key()) + + response: Final = client.responses.create( + model=model, + input="reply with one word", + instructions=INSTRUCTIONS, + truncation="auto", + extra_body=NO_PROXY_CACHE, + ) + assert response.output_text.strip(), ( + f"/responses over azure openai preview returned no output text: {response.output!r}" + ) + + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_RESPONSES_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) + def test_responses_compact_returns_compacted_conversation( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model: Final = _register( + proxy, + resources, + _openai_responses_params(), + prefix="e2e-responses-compact", + ) + client: Final = sdk.openai(resources.key()) + conversation: Final[ResponseInputParam] = [ + {"role": "user", "content": "Remember that my favorite color is blue."}, + {"role": "assistant", "content": "I will remember that your favorite color is blue."}, + ] + + compacted: Final = client.responses.compact( + model=model, + input=conversation, + extra_body=NO_PROXY_CACHE, + ) + assert compacted.id, f"/responses/compact returned no id: {compacted!r}" + assert any(item.type == "compaction" for item in compacted.output), ( + f"/responses/compact returned no compaction item: {compacted.output!r}" + ) + compacted_input: Final[ResponseInputParam] = TypeAdapter(ResponseInputParam).validate_python( + [item.model_dump(exclude_none=True) for item in compacted.output] + + [{"role": "user", "content": "What is my favorite color?"}] + ) + + response: Final = client.responses.create( + model=model, + input=compacted_input, + extra_body=NO_PROXY_CACHE, + ) + assert "blue" in response.output_text.lower(), ( + f"compacted conversation did not retain the favorite color: {response.output_text!r}" + ) + + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(Provider.OPENAI,), + models=(OPENAI_RESPONSES_BACKEND,), + mode=Mode.NONSTREAM, + ) + ) + def test_responses_context_management_compacts_server_side( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model: Final = _register( + proxy, + resources, + _openai_responses_params(), + prefix="e2e-responses-context-compaction", + ) + client: Final = sdk.openai(resources.key()) + filler: Final = "The archive record has a blue marker beside every stored entry. " * 350 + conversation: Final[ResponseInputParam] = [ + {"role": "user", "content": filler}, + {"role": "assistant", "content": "I have read the archive and retained its details."}, + {"role": "user", "content": "Reply with one word to verify server-side compaction."}, + ] + + response: Final = client.responses.create( + model=model, + input=conversation, + context_management=[{"type": "compaction", "compact_threshold": 1000}], + extra_body=NO_PROXY_CACHE, + ) + assert response.status == "completed", f"context management did not complete: {response.status}" + assert any(item.type == "compaction" for item in response.output), ( + f"context management returned no compaction item: {response.output!r}" + ) + @pytest.mark.covers("llm.responses.azure_openai.tool_use.nonstream.works") @meta( Subject( @@ -493,7 +868,7 @@ class TestResponses: def test_responses_azure_openai_returns_function_call( self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients ) -> None: - model = _register(proxy, resources, _azure_openai_params(), prefix="e2e-responses-azure-openai-tool") + model = _register(proxy, resources, azure_openai_params(), prefix="e2e-responses-azure-openai-tool") client = sdk.openai(resources.key()) response = client.responses.create( diff --git a/tests/e2e/llm_translation/test_responses_retrieve_e2e.py b/tests/e2e/llm_translation/test_responses_retrieve_e2e.py index b4a9634fbef..dd94d5cab72 100644 --- a/tests/e2e/llm_translation/test_responses_retrieve_e2e.py +++ b/tests/e2e/llm_translation/test_responses_retrieve_e2e.py @@ -6,7 +6,7 @@ Creates a stored response, retrieves it by id, and pins invalid-id error handlin from __future__ import annotations import time -from typing import Final +from typing import Final, Literal import openai import pytest @@ -15,7 +15,9 @@ from e2e_http import NoBody, Success, UnknownApiError, unwrap from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta from lifecycle import ResourceManager from models import LiteLLMParamsBody +from responses_helpers import AZURE_OPENAI_BACKEND, azure_openai_params from openai.types.responses import ( + ResponseCompletedEvent, ResponseCreatedEvent, ResponseInputMessageItem, ResponseInputText, @@ -138,12 +140,42 @@ class TestResponsesRetrieve: def _register_openai(proxy: ProxyClient, resources: ResourceManager, prefix: str) -> str: - model = f"{prefix}-{unique_marker()}" - model_id = proxy.create_model(model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY")) + return _register_response_deployment(proxy, resources, "openai", prefix) + + +def _register_response_deployment( + proxy: ProxyClient, + resources: ResourceManager, + deployment: Literal["openai", "azure"], + prefix: str, +) -> str: + model: Final = f"{prefix}-{unique_marker()}" + params: Final = ( + azure_openai_params() + if deployment == "azure" + else LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") + ) + model_id: Final = proxy.create_model(model, params) resources.defer(lambda: proxy.delete_model(model_id)) return model +def _deployment_param(deployment: Literal["openai", "azure"], provider: Provider, backend: str, mode: Mode) -> object: + return pytest.param( + deployment, + id=deployment, + marks=meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.RESPONSES, + providers=(provider,), + models=(backend,), + mode=mode, + ) + ), + ) + + def _input_texts(item: object) -> tuple[str, ...]: if not isinstance(item, ResponseInputMessageItem): return () @@ -176,57 +208,96 @@ class TestStoredResponseLifecycle: texts = tuple(text for item in items for text in _input_texts(item)) assert any(marker in text for text in texts), f"input_items did not list the stored prompt: {items!r}" - @meta( - Subject( - domain=Domain.LLM_TRANSLATION, - route=Route.RESPONSES, - providers=(Provider.OPENAI,), - models=(OPENAI_BACKEND,), - mode=Mode.NONSTREAM, - ) + @pytest.mark.parametrize( + "deployment", + [ + _deployment_param("openai", Provider.OPENAI, OPENAI_BACKEND, Mode.NONSTREAM), + _deployment_param("azure", Provider.AZURE, AZURE_OPENAI_BACKEND, Mode.NONSTREAM), + ], ) def test_deleted_response_is_no_longer_retrievable( - self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + self, + proxy: ProxyClient, + resources: ResourceManager, + sdk: SdkClients, + deployment: Literal["openai", "azure"], ) -> None: - model = _register_openai(proxy, resources, "e2e-resp-delete") - client = sdk.openai(resources.key()) + model: Final = _register_response_deployment(proxy, resources, deployment, "e2e-resp-delete") + client: Final = sdk.openai(resources.key()) - created = client.responses.create( + created: Final = client.responses.create( model=model, input=f"Reply with one word. {unique_marker()}", store=True, extra_body=NO_PROXY_CACHE ) - retrieved = client.responses.retrieve(created.id) + retrieved: Final = client.responses.retrieve(created.id) assert retrieved.status == "completed", f"stored response not retrievable as completed: {retrieved!r}" + assert retrieved.output_text == created.output_text, ( + f"retrieved output changed: created={created.output_text!r}, retrieved={retrieved.output_text!r}" + ) client.responses.delete(created.id) - with pytest.raises(openai.APIStatusError) as gone: - client.responses.retrieve(created.id) + gone: Final = pytest.raises(openai.APIStatusError, client.responses.retrieve, created.id) assert 400 <= gone.value.status_code < 500, f"retrieve after delete expected a 4xx: {gone.value!r}" + @pytest.mark.parametrize( + "deployment", + [ + _deployment_param("openai", Provider.OPENAI, OPENAI_BACKEND, Mode.STREAM), + _deployment_param("azure", Provider.AZURE, AZURE_OPENAI_BACKEND, Mode.STREAM), + ], + ) + def test_streamed_response_can_be_deleted( + self, + proxy: ProxyClient, + resources: ResourceManager, + sdk: SdkClients, + deployment: Literal["openai", "azure"], + ) -> None: + model: Final = _register_response_deployment(proxy, resources, deployment, "e2e-resp-delete-stream") + client: Final = sdk.openai(resources.key()) + events: Final = tuple( + client.responses.create( + model=model, + input=f"Reply with one word. {unique_marker()}", + store=True, + stream=True, + extra_body=NO_PROXY_CACHE, + ) + ) + completed: Final = next((event for event in events if isinstance(event, ResponseCompletedEvent)), None) + assert completed is not None, f"stream did not complete: {events!r}" + response_id: Final = completed.response.id + + client.responses.delete(response_id) + gone: Final = pytest.raises(openai.APIStatusError, client.responses.retrieve, response_id) + assert 400 <= gone.value.status_code < 500, f"retrieve after streamed delete expected a 4xx: {gone.value!r}" + @pytest.mark.provider_live class TestBackgroundResponseCancel: - @meta( - Subject( - domain=Domain.LLM_TRANSLATION, - route=Route.RESPONSES, - providers=(Provider.OPENAI,), - models=(OPENAI_BACKEND,), - mode=Mode.NONSTREAM, - ) + @pytest.mark.parametrize( + "deployment", + [ + _deployment_param("openai", Provider.OPENAI, OPENAI_BACKEND, Mode.NONSTREAM), + _deployment_param("azure", Provider.AZURE, AZURE_OPENAI_BACKEND, Mode.NONSTREAM), + ], ) def test_cancel_background_response( - self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + self, + proxy: ProxyClient, + resources: ResourceManager, + sdk: SdkClients, + deployment: Literal["openai", "azure"], ) -> None: - model = _register_openai(proxy, resources, "e2e-resp-cancel") - client = sdk.openai(resources.key()) + model: Final = _register_response_deployment(proxy, resources, deployment, "e2e-resp-cancel") + client: Final = sdk.openai(resources.key()) - created = client.responses.create( + created: Final = client.responses.create( model=model, input=f"{LONG_TASK} {unique_marker()}", background=True, extra_body=NO_PROXY_CACHE ) assert created.status in CANCELLABLE_STATUSES, f"background response was not queued: {created.status}" - cancelled = client.responses.cancel(created.id) + cancelled: Final = client.responses.cancel(created.id) assert cancelled.status == "cancelled", f"cancel did not stop the response: {cancelled.status}" @meta( diff --git a/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md b/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md index d9380098891..e3f3bf009dc 100644 --- a/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md +++ b/tests/e2e/quota_management/spend_tracking/SPEND_TRACKING_COVERAGE_MATRIX.md @@ -43,7 +43,7 @@ proxy + SpendLogs rows. Status: `covered` / `partial` / `gap`. | Tag | `test_update_daily_tag_spend.py` | partial | yes (`test_tag_spend_matches_sum_of_tagged_logs`) | | End-user | `test_proxy_update_spend.py` | covered | yes | | Spend == sum(logs) consistency | none | gap | yes (key + tag aggregate == sum of rows) | -| Concurrent increments (one key, parallel writers) | `tests/spend_tracking_tests/test_spend_accuracy_tests.py` (burst) | partial | yes (`test_burst_of_concurrent_calls_loses_no_spend`) | +| Concurrent increments (one key, parallel writers) | `tests/integration/spend/test_spend_rollup_accuracy.py`, `tests/integration/spend/test_chaos_burst_spend_once.py` (burst) | partial | yes (`test_burst_of_concurrent_calls_loses_no_spend`) | ## Spend read endpoints (verification surface) diff --git a/tests/e2e/ui/fixtures/migratedPages.ts b/tests/e2e/ui/fixtures/migratedPages.ts index d44dfc625aa..5fa1f9d2444 100644 --- a/tests/e2e/ui/fixtures/migratedPages.ts +++ b/tests/e2e/ui/fixtures/migratedPages.ts @@ -117,7 +117,7 @@ export const MIGRATED_E2E_PAGES: Readonly> = { linkName: "AI Hub", content: { role: "heading", name: "AI Hub" }, }, - new_usage: { segment: "usage", linkName: "Usage", content: { role: "heading", name: "Usage View" } }, + new_usage: { segment: "usage", linkName: "Usage", content: { role: "heading", name: "Usage" } }, usage: { segment: "old-usage", linkName: "Old Usage", diff --git a/tests/e2e/ui/tests/usage/usagePage.spec.ts b/tests/e2e/ui/tests/usage/usagePage.spec.ts index 3d057cfa2c9..d532042c7af 100644 --- a/tests/e2e/ui/tests/usage/usagePage.spec.ts +++ b/tests/e2e/ui/tests/usage/usagePage.spec.ts @@ -13,9 +13,8 @@ import { /** Covers /ui/usage. The legacy /ui/old-usage view is deprecated and deliberately not covered. */ -/** Stepping up from the title is exact; the page renders several other tables. */ const topKeysCard = (page: PlaywrightPage): Locator => - page.getByText("Top Virtual Keys", { exact: true }).locator("xpath=.."); + page.locator("section").filter({ has: page.getByRole("heading", { name: "Top Virtual Keys", exact: true }) }); async function openUsage(page: PlaywrightPage): Promise { await navigateToPage(page, Page.NewUsage); diff --git a/tests/guardrails_tests/test_bedrock_guardrails.py b/tests/guardrails_tests/test_bedrock_guardrails.py deleted file mode 100644 index aeeb34f4e8d..00000000000 --- a/tests/guardrails_tests/test_bedrock_guardrails.py +++ /dev/null @@ -1,100 +0,0 @@ -import io, asyncio -import pytest - -import litellm -from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( - BedrockGuardrail, - _redact_pii_matches, -) -from litellm.proxy._types import UserAPIKeyAuth -from unittest.mock import MagicMock, AsyncMock, patch - - -@pytest.mark.asyncio -async def test_bedrock_guardrails_pii_masking(): - # Create proper mock objects - mock_user_api_key_dict = UserAPIKeyAuth() - - guardrail = BedrockGuardrail( - guardrailIdentifier="wf0hkdb5x07f", - guardrailVersion="DRAFT", - ) - - request_data = { - "model": "gpt-5.5", - "messages": [ - {"role": "user", "content": "Hello, my phone number is +1 412 555 1212"}, - {"role": "assistant", "content": "Hello, how can I help you today?"}, - {"role": "user", "content": "I need to cancel my order"}, - { - "role": "user", - "content": "ok, my credit card number is 1234-5678-9012-3456", - }, - ], - } - - response = await guardrail.async_moderation_hook( - data=request_data, - user_api_key_dict=mock_user_api_key_dict, - call_type="completion", - ) - print("response after moderation hook", response) - - if response: # Only assert if response is not None - assert response["messages"][0]["content"] == "Hello, my phone number is {PHONE}" - assert response["messages"][1]["content"] == "Hello, how can I help you today?" - assert response["messages"][2]["content"] == "I need to cancel my order" - assert ( - response["messages"][3]["content"] - == "ok, my credit card number is {CREDIT_DEBIT_CARD_NUMBER}" - ) - - -@pytest.mark.asyncio -async def test_bedrock_guardrails_pii_masking_content_list(): - # Create proper mock objects - mock_user_api_key_dict = UserAPIKeyAuth() - - guardrail = BedrockGuardrail( - guardrailIdentifier="wf0hkdb5x07f", - guardrailVersion="DRAFT", - ) - - request_data = { - "model": "gpt-5.5", - "messages": [ - { - "role": "user", - "content": [ - { - "type": "text", - "text": "Hello, my phone number is +1 412 555 1212", - }, - {"type": "text", "text": "what time is it?"}, - ], - }, - {"role": "assistant", "content": "Hello, how can I help you today?"}, - {"role": "user", "content": "who is the president of the united states?"}, - ], - } - - response = await guardrail.async_moderation_hook( - data=request_data, - user_api_key_dict=mock_user_api_key_dict, - call_type="completion", - ) - print(response) - - if response: # Only assert if response is not None - # Verify that the list content is properly masked - assert isinstance(response["messages"][0]["content"], list) - assert ( - response["messages"][0]["content"][0]["text"] - == "Hello, my phone number is {PHONE}" - ) - assert response["messages"][0]["content"][1]["text"] == "what time is it?" - assert response["messages"][1]["content"] == "Hello, how can I help you today?" - assert ( - response["messages"][2]["content"] - == "who is the president of the united states?" - ) diff --git a/tests/guardrails_tests/test_presidio_pii.py b/tests/guardrails_tests/test_presidio_pii.py deleted file mode 100644 index 595cc95fa9f..00000000000 --- a/tests/guardrails_tests/test_presidio_pii.py +++ /dev/null @@ -1,211 +0,0 @@ -import os -import pytest -from litellm import mock_completion -from unittest.mock import patch - -import litellm -from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - OPTIONAL_PresidioPIIMasking, - PresidioPerRequestConfig, -) -from litellm.types.guardrails import PiiEntityType, PiiAction -from litellm.proxy._types import UserAPIKeyAuth -from litellm.caching.caching import DualCache -from litellm.exceptions import BlockedPiiEntityError - - -@pytest.mark.asyncio -async def test_presidio_with_blocked_entities(): - """Test for Presidio guardrail with blocked entities - requires actual Presidio API""" - # Setup the guardrail with specific entities config - BLOCK for credit card - litellm.turn_on_debug() - pii_entities_config = { - PiiEntityType.CREDIT_CARD: PiiAction.BLOCK, # This entity should cause a block - PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK, # This entity should be masked - } - - presidio_guardrail = OPTIONAL_PresidioPIIMasking( - pii_entities_config=pii_entities_config, - presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"), - presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"), - ) - - # Test text with blocked PII type - test_text = ( - "My credit card number is 4111-1111-1111-1111 and my email is test@example.com" - ) - - # Verify the analyze request configuration - analyze_request = presidio_guardrail._get_presidio_analyze_request_payload( - text=test_text, presidio_config=None, request_data={} - ) - - # Verify entities were passed correctly - assert "entities" in analyze_request - assert set(analyze_request["entities"]) == set(pii_entities_config.keys()) - - # Test that BlockedPiiEntityError is raised when check_pii is called - with pytest.raises(BlockedPiiEntityError) as excinfo: - await presidio_guardrail.check_pii( - text=test_text, output_parse_pii=True, presidio_config=None, request_data={} - ) - - # Verify the error contains the correct entity type - assert excinfo.value.entity_type == PiiEntityType.CREDIT_CARD - assert excinfo.value.guardrail_name == presidio_guardrail.guardrail_name - - -@pytest.mark.asyncio -async def test_presidio_pre_call_hook_with_blocked_entities(): - """Test for Presidio guardrail pre-call hook with blocked entities on a chat completion request""" - # Setup the guardrail with specific entities config - pii_entities_config = { - PiiEntityType.CREDIT_CARD: PiiAction.BLOCK, # This entity should cause a block - PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK, # This entity should be masked - } - - presidio_guardrail = OPTIONAL_PresidioPIIMasking( - pii_entities_config=pii_entities_config, - presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"), - presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"), - ) - - # Create a sample chat completion request with PII data - data = { - "messages": [ - {"role": "system", "content": "You are a helpful assistant."}, - { - "role": "user", - "content": "My credit card is 4111-1111-1111-1111 and my email is test@example.com.", - }, - ], - "model": "gpt-5-mini", - } - - # Mock objects needed for the pre-call hook - user_api_key_dict = UserAPIKeyAuth(api_key="test_key") - cache = DualCache() - - # Call the pre-call hook and expect BlockedPiiEntityError - with pytest.raises(BlockedPiiEntityError) as excinfo: - await presidio_guardrail.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, - cache=cache, - data=data, - call_type="completion", - ) - - print(f"got error: {excinfo}") - - # Verify the error contains the correct entity type - assert excinfo.value.entity_type == PiiEntityType.CREDIT_CARD - assert excinfo.value.guardrail_name == presidio_guardrail.guardrail_name - - - - - - -# asyncio.run(test_output_parsing()) - - -### UNIT TESTS FOR PRESIDIO PII MASKING ### - -input_a_anonymizer_results = { - "text": "hello world, my name is . My number is: ", - "items": [ - { - "start": 48, - "end": 62, - "entity_type": "PHONE_NUMBER", - "text": "", - "operator": "replace", - }, - { - "start": 24, - "end": 32, - "entity_type": "PERSON", - "text": "", - "operator": "replace", - }, - ], -} - -input_b_anonymizer_results = { - "text": "My name is , who are you? Say my name in your response", - "items": [ - { - "start": 11, - "end": 19, - "entity_type": "PERSON", - "text": "", - "operator": "replace", - } - ], -} - - -# Test if PII masking works with input A - - -# Test if PII masking works with input B (also test if the response != A's response) - - - - -@pytest.mark.asyncio -@patch.dict( - os.environ, - { - "PRESIDIO_ANALYZER_API_BASE": "http://localhost:5002", - "PRESIDIO_ANONYMIZER_API_BASE": "http://localhost:5001", - }, -) -async def test_presidio_pii_masking_logging_output_only_logged_response_guardrails_config(): - from typing import Dict, List, Optional - - import litellm - from litellm.proxy.guardrails.init_guardrails import initialize_guardrails - from litellm.types.guardrails import ( - GuardrailItemSpec, - GuardrailEventHooks, - ) - - litellm.set_verbose = True - # Environment variables are now patched via the decorator instead of setting them directly - - guardrails_config: List[Dict[str, GuardrailItemSpec]] = [ - { - "pii_masking": { - "callbacks": ["presidio"], - "default_on": True, - "logging_only": True, - } - } - ] - litellm_settings = {"guardrails": guardrails_config} - - assert len(litellm.guardrail_name_config_map) == 0 - initialize_guardrails( - guardrails_config=guardrails_config, - premium_user=True, - config_file_path="", - litellm_settings=litellm_settings, - ) - - assert len(litellm.guardrail_name_config_map) == 1 - - pii_masking_obj: Optional[OPTIONAL_PresidioPIIMasking] = None - for callback in litellm.callbacks: - print(f"CALLBACK: {callback}") - if isinstance(callback, OPTIONAL_PresidioPIIMasking): - pii_masking_obj = callback - - assert pii_masking_obj is not None - - assert hasattr(pii_masking_obj, "logging_only") - assert pii_masking_obj.event_hook == GuardrailEventHooks.logging_only - - assert pii_masking_obj.should_run_guardrail( - data={}, event_type=GuardrailEventHooks.logging_only - ) diff --git a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py deleted file mode 100644 index e7d526beeb8..00000000000 --- a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py +++ /dev/null @@ -1,112 +0,0 @@ -import logging -import traceback - -from dotenv import load_dotenv -from openai.types.image import Image - - -from litellm.llms.bedrock.image_generation.amazon_nova_canvas_transformation import ( - AmazonNovaCanvasConfig, -) - -logging.basicConfig(level=logging.DEBUG) -load_dotenv() -import asyncio - -import pytest -from litellm.llms.bedrock.image_generation.cost_calculator import cost_calculator -from litellm.types.utils import ImageResponse, ImageObject - -import litellm -from litellm.llms.bedrock.image_generation.amazon_stability3_transformation import ( - AmazonStability3Config, -) -from litellm.llms.bedrock.image_generation.amazon_stability1_transformation import ( - AmazonStabilityConfig, -) -from litellm.types.llms.bedrock import ( - AmazonStability3TextToImageRequest, - AmazonStability3TextToImageResponse, -) -from unittest.mock import MagicMock, patch -from litellm.llms.bedrock.image_generation.image_handler import ( - BedrockImageGeneration, - BedrockImagePreparedRequest, -) -from litellm.llms.bedrock.common_utils import BedrockError - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -# Test cases for issue #14373 - Bedrock Application Inference Profiles with Nova Canvas - - - - - - - - -def test_amazon_nova_canvas_image_gen(): - """Test Amazon Nova Canvas image generation with cost tracking.""" - from litellm import image_generation - - model_id = "bedrock/amazon.nova-canvas-v1:0" - - response = litellm.image_generation( - model=model_id, - prompt="A serene mountain landscape at sunset with a lake reflection", - aws_region_name="us-east-1", - ) - - print(f"response cost: {response._hidden_params['response_cost']}") - - assert response._hidden_params["response_cost"] > 0 diff --git a/tests/image_gen_tests/test_image_edits.py b/tests/image_gen_tests/test_image_edits.py index fdcbf6fcd8b..eca22196287 100644 --- a/tests/image_gen_tests/test_image_edits.py +++ b/tests/image_gen_tests/test_image_edits.py @@ -193,182 +193,3 @@ async def test_openai_image_edit_litellm_router(): f.write(image_bytes) except litellm.ContentPolicyViolationError as e: pass - - -@pytest.mark.flaky(retries=3, delay=2) -@pytest.mark.asyncio -async def test_openai_image_edit_with_bytesio(): - """Test image editing using BytesIO objects instead of file readers""" - from litellm import aimage_edit, image_edit - - litellm.turn_on_debug() - try: - prompt = """ - Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. - """ - - # Get images as BytesIO objects - bytesio_images = get_test_images_as_bytesio() - - result = await aimage_edit( - prompt=prompt, - model="gpt-image-1", - image=bytesio_images, - ) - print("result from image edit with BytesIO", result) - - # Validate the response meets expected schema - ImageResponse.model_validate(result) - - if isinstance(result, ImageResponse) and result.data: - image_base64 = result.data[0].b64_json - if image_base64: - image_bytes = base64.b64decode(image_base64) - - # Save the image to a file - with open("test_image_edit_bytesio.png", "wb") as f: - f.write(image_bytes) - except litellm.ContentPolicyViolationError as e: - pass - - - - - - -@pytest.mark.asyncio -async def test_azure_image_edit_cost_tracking(): - """Test Azure image edit cost tracking with custom logger""" - from litellm import aimage_edit, image_edit - - test_custom_logger = TestCustomLogger() - litellm.logging_callback_manager._reset_all_callbacks() - litellm.callbacks = [test_custom_logger] - - # Mock response for Azure image edit with usage data for cost tracking - mock_response = { - "created": 1589478378, - "data": [ - { - "b64_json": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg==" - } - ], - "usage": { - "total_tokens": 1100, - "input_tokens": 100, - "input_tokens_details": {"image_tokens": 50, "text_tokens": 50}, - "output_tokens": 1000, - }, - } - - class MockResponse: - def __init__(self, json_data, status_code): - self._json_data = json_data - self.status_code = status_code - self.text = json.dumps(json_data) - self.headers = {} - - def json(self): - return self._json_data - - with patch( - "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", - new_callable=AsyncMock, - ) as mock_post: - # Configure the mock to return our response - mock_post.return_value = MockResponse(mock_response, 200) - - litellm.turn_on_debug() - - prompt = """ - Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. - """ - - # Set up test environment variables - - result = await aimage_edit( - prompt=prompt, - model="azure/CUSTOM_AZURE_DEPLOYMENT_NAME", - base_model="azure/gpt-image-1", - image=_make_test_images(), - ) - - # Verify the request was made correctly - mock_post.assert_called_once() - - # Validate the response meets expected schema - ImageResponse.model_validate(result) - - if isinstance(result, ImageResponse) and result.data: - image_base64 = result.data[0].b64_json - if image_base64: - image_bytes = base64.b64decode(image_base64) - - # Save the image to a file - with open("test_image_edit.png", "wb") as f: - f.write(image_bytes) - - await asyncio.sleep(5) - print( - "standard logging payload", - json.dumps( - test_custom_logger.standard_logging_payload, indent=4, default=str - ), - ) - - # check model - assert ( - test_custom_logger.standard_logging_payload["model"] - == "CUSTOM_AZURE_DEPLOYMENT_NAME" - ) - assert ( - test_custom_logger.standard_logging_payload["custom_llm_provider"] - == "azure" - ) - - # check response_cost - assert test_custom_logger.standard_logging_payload["response_cost"] is not None - assert test_custom_logger.standard_logging_payload["response_cost"] > 0 - - - - - - -@pytest.mark.flaky(retries=3, delay=2) -@pytest.mark.asyncio -async def test_multiple_image_edit_with_different_formats(): - """Test multiple images editing with different file formats and types""" - from litellm import aimage_edit - - litellm.turn_on_debug() - - try: - prompt = "Create a cohesive artistic style across all images" - - mixed_images = [ - _make_single_test_image(), - get_test_images_as_bytesio()[1], - ] - - result = await aimage_edit( - prompt=prompt, - model="gpt-image-1", - image=mixed_images, - ) - - print("Mixed format images result:", result) - ImageResponse.model_validate(result) - - assert result is not None - assert result.data is not None - assert len(result.data) > 0 - - # Save result if available - if result.data and result.data[0].b64_json: - image_bytes = base64.b64decode(result.data[0].b64_json) - with open("test_multiple_image_edit_mixed.png", "wb") as f: - f.write(image_bytes) - - except litellm.ContentPolicyViolationError as e: - pytest.skip(f"Content policy violation: {e}") diff --git a/tests/integration/_support/forward_proxy.py b/tests/integration/_support/forward_proxy.py new file mode 100644 index 00000000000..9f5cff94277 --- /dev/null +++ b/tests/integration/_support/forward_proxy.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +import threading +from collections.abc import Generator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from queue import SimpleQueue +from typing import Final + + +@dataclass(frozen=True, slots=True) +class Tunnel: + method: str + target: str + headers: Mapping[str, str] + + +@dataclass(frozen=True, slots=True) +class ForwardProxy: + url: str + port: int + received: SimpleQueue[Tunnel] + + def drain(self) -> tuple[Tunnel, ...]: + return tuple(self.received.get_nowait() for _ in range(self.received.qsize())) + + def targets(self) -> tuple[str, ...]: + return tuple(tunnel.target for tunnel in self.drain()) + + +@contextmanager +def refusing_forward_proxy(port: int = 0) -> Generator[ForwardProxy, None, None]: + """Owned HTTP forward proxy: records the host each CONNECT asks for and refuses the tunnel with 403. + + A litellm proxy booted with ``HTTPS_PROXY`` pointed here makes the host it dials for a provider + observable without a byte leaving the box; the caller sees litellm's own connection error. + """ + received: Final[SimpleQueue[Tunnel]] = SimpleQueue() + + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + timeout = 5 + + def refuse(self) -> None: + received.put(Tunnel(self.command, self.path, {name.lower(): value for name, value in self.headers.items()})) + self.send_response(403) + self.send_header("content-length", "0") + self.send_header("connection", "close") + self.end_headers() + self.close_connection = True + + do_CONNECT = refuse + do_GET = refuse + do_POST = refuse + do_PUT = refuse + do_DELETE = refuse + + def log_message(self, format: str, *args: object) -> None: + pass + + with ThreadingHTTPServer(("127.0.0.1", port), Handler) as server: + thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05}) + thread.start() + try: + yield ForwardProxy(f"http://127.0.0.1:{server.server_port}", server.server_port, received) + finally: + server.shutdown() + thread.join(timeout=6) + assert not thread.is_alive(), "Owned forward proxy survived cleanup" + server.server_close() diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 12780540380..23297b28584 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -260,6 +260,8 @@ def owned_proxy_process( "127.0.0.1", "--num_workers", str(workers), + "--timeout_worker_healthcheck", + str(int(graceful_stop_seconds())), *database_setup, *extra_arguments, ) @@ -294,6 +296,8 @@ def owned_gateway_image( str(workers), "--host", "127.0.0.1", + "--timeout-worker-healthcheck", + str(int(graceful_stop_seconds())), ) launch: Final = _launch_until_bound(command, root, environment, output, _PORT_ATTEMPTS) try: @@ -397,9 +401,7 @@ class UpstreamSlot: __slots__ = ("certificate", "directory", "port", "process", "root") - def __init__( - self, directory: Path, port: int, root: Path, certificate: UpstreamCertificate | None = None - ) -> None: + def __init__(self, directory: Path, port: int, root: Path, certificate: UpstreamCertificate | None = None) -> None: self.directory = directory self.port = port self.root = root diff --git a/tests/integration/_support/prompt_cache_breakpoint.py b/tests/integration/_support/prompt_cache_breakpoint.py index ce3b97b86be..bfc7f1efbce 100644 --- a/tests/integration/_support/prompt_cache_breakpoint.py +++ b/tests/integration/_support/prompt_cache_breakpoint.py @@ -32,13 +32,13 @@ _SCRIPTED_FAILURE: Final = re.compile(r"fail-(\d{3})") _MINTED_RESPONSE: Final = re.compile(r"^resp_([0-9a-f]{32})-[0-9a-f]{32}$") _STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") -Kind: TypeAlias = Literal["text", "image_url", "file", "input_audio"] -KINDS: Final[tuple[Kind, ...]] = ("text", "image_url", "file", "input_audio") +Kind: TypeAlias = Literal["text", "image_url", "file", "video_url"] +KINDS: Final[tuple[Kind, ...]] = ("text", "image_url", "file", "video_url") WIRE_TYPE: Final[Mapping[Kind, str]] = { "text": "input_text", "image_url": "input_image", "file": "input_file", - "input_audio": "input_text", + "video_url": "input_text", } @@ -91,12 +91,19 @@ def block(kind: Kind, value: str) -> dict[str, JsonValue]: return {"type": "image_url", "image_url": {"url": "https://example.com/breakpoint.png"}} case "file": return {"type": "file", "file": {"file_id": "file-breakpoint"}} - case "input_audio": - return {"type": "input_audio", "input_audio": {"data": "Zm9v", "format": "wav"}} + case "video_url": + return {"type": "video_url", "video_url": {"url": "https://example.com/clip.mp4"}} case _: assert_never(kind) +AUDIO_PAYLOAD: Final[Mapping[str, JsonValue]] = {"data": "Zm9v", "format": "wav"} + + +def audio() -> dict[str, JsonValue]: + return {"type": "input_audio", "input_audio": dict(AUDIO_PAYLOAD)} + + def drained_posts(wire: Wire) -> tuple[Request, ...]: return tuple(request for request in wire.drain() if request.method == "POST") diff --git a/tests/integration/_support/wire.py b/tests/integration/_support/wire.py index d118da03daf..e6a166e8457 100644 --- a/tests/integration/_support/wire.py +++ b/tests/integration/_support/wire.py @@ -3,13 +3,13 @@ from __future__ import annotations import ssl import threading import time -from collections.abc import Callable, Generator, Mapping +from collections.abc import Callable, Generator, Iterator, Mapping from contextlib import contextmanager from dataclasses import dataclass from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from queue import SimpleQueue from types import MappingProxyType -from typing import Final +from typing import BinaryIO, Final @dataclass(frozen=True, slots=True) @@ -20,6 +20,37 @@ class Request: body: bytes +class AbortedBody(Exception): + """The client closed the connection before the body it announced was complete.""" + + +def exactly(stream: BinaryIO, size: int) -> bytes: + data: Final = stream.read(size) + if len(data) < size: + raise AbortedBody + return data + + +def chunked_body(stream: BinaryIO) -> Iterator[bytes]: + while True: + size_line: Final = stream.readline() + if not size_line: + raise AbortedBody + size: Final = int(size_line.split(b";")[0].strip(), 16) + if size == 0: + while stream.readline().strip(): + pass + return + yield exactly(stream, size) + stream.readline() + + +def read_body(headers: Mapping[str, str], stream: BinaryIO) -> bytes: + if headers.get("transfer-encoding", "").lower() == "chunked": + return b"".join(chunked_body(stream)) + return exactly(stream, int(headers.get("content-length", "0"))) + + @dataclass(frozen=True, slots=True) class Reply: status: int = 200 @@ -78,12 +109,14 @@ def wire_server( connected.put(f"{self.client_address[0]}:{self.client_address[1]}") def respond(self) -> None: - request: Final = Request( - self.command, - self.path, - {name.lower(): value for name, value in self.headers.items()}, - self.rfile.read(int(self.headers.get("content-length", "0"))), - ) + headers: Final = {name.lower(): value for name, value in self.headers.items()} + try: + body: Final = read_body(headers, self.rfile) + except AbortedBody: + self.close_connection = True + disconnected.put(self.path) + return + request: Final = Request(self.command, self.path, headers, body) received.put(request) try: reply = respond(request) diff --git a/tests/integration/authorization/test_cli_sso_login_ui_disabled.py b/tests/integration/authorization/test_cli_sso_login_ui_disabled.py index 126fe27bef2..a5e9b2752ea 100644 --- a/tests/integration/authorization/test_cli_sso_login_ui_disabled.py +++ b/tests/integration/authorization/test_cli_sso_login_ui_disabled.py @@ -23,7 +23,14 @@ import pytest import yaml from pydantic import JsonValue -from tests.integration._support.client import JSON_OBJECT, Gateway, eventually, gateway_from_environment, string_value +from tests.integration._support.client import ( + JSON_OBJECT, + Gateway, + eventually, + gateway_from_environment, + object_value, + string_value, +) from tests.integration._support.database import read_rows from tests.integration._support.process import OwnedProxy, group_members, owned_proxy_process from tests.integration._support.provider import SharedProvider @@ -648,3 +655,52 @@ def test_flag_values_that_do_not_disable_keep_the_gates_open(idp: Idp, tmp_path: for method, path in (("GET", "/sso/saml/login"), ("POST", "/sso/saml/callback")): saml: Final = proxy.client.request(method, path) assert DISABLED_PAGE_TITLE not in saml.text and saml.status_code != 200, f"{path}: {saml.status_code}" + + +def test_cli_session_token_is_denied_once_its_team_budget_is_exhausted( + one_worker: OneWorkerProxy, provider: SharedProvider +) -> None: + proxy: Final = one_worker.owned.gateway + subject: Final = f"cli-sso-team-budget-{uuid.uuid4().hex[:12]}" + budget: Final = 0.0000000005 + with proxy.scenario() as scenario: + team: Final = scenario.team(max_budget=budget, models=[MESSAGE_MODEL]) + scenario.user(user_id=subject, user_email=f"{subject}@example.com", user_role="internal_user") + added: Final = proxy.request( + "POST", "/team/member_add", {"team_id": team, "member": {"user_id": subject, "role": "user"}} + ) + assert added.status_code == 200, added.text + session: Final = _start_lite_login(proxy) + with _browser() as browser: + _sign_in(proxy, one_worker.idp, browser, session, subject=subject) + ready: Final = proxy.client.get( + f"/sso/cli/poll/{session.login_id}", + params={"team_id": team}, + headers={POLL_SECRET_HEADER: session.poll_secret}, + ) + assert ready.status_code == 200, f"{ready.status_code} {ready.text}" + body: Final = JSON_OBJECT.validate_json(ready.content) + assert body["status"] == "ready" and body["user_id"] == subject, ready.text + key: Final = string_value(body["key"]) + assert not key.startswith("sk-"), key + _send_message(proxy, provider, key) + eventually( + lambda: read_rows('SELECT spend FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team,)), + lambda rows: len(rows) == 1 and float(str(rows[0]["spend"])) > budget, + seconds=70, + ) + refused: Final = proxy.request( + "POST", + "/v1/messages", + {"model": MESSAGE_MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "over team budget"}]}, + key=key, + ) + assert refused.status_code == 422, f"{refused.status_code} {refused.text}" + error: Final = object_value(JSON_OBJECT.validate_json(refused.content)["error"]) + assert error["type"] == "budget_exceeded", refused.text + assert error["code"] == "422", refused.text + message: Final = string_value(error["message"]) + assert "Budget has been exceeded!" in message, refused.text + assert f"Team={team}" in message, refused.text + assert "Current cost:" in message and f"Max budget: {budget}" in message, refused.text + assert provider.received() == () diff --git a/tests/integration/authorization/test_model_access_allow_lists.py b/tests/integration/authorization/test_model_access_allow_lists.py new file mode 100644 index 00000000000..cdefc8c8a3e --- /dev/null +++ b/tests/integration/authorization/test_model_access_allow_lists.py @@ -0,0 +1,171 @@ +from __future__ import annotations + +import json +import uuid +from collections.abc import Sequence +from dataclasses import dataclass +from typing import Final, Literal + +import pytest + +from tests.integration._support.client import JSON_OBJECT, Gateway, Scenario, object_value, string_value +from tests.integration._support.provider import PROVIDER_URL, SharedProvider +from tests.integration._support.wire import Reply + +_Denial = Literal["key_model_access_denied", "team_model_access_denied"] + + +@dataclass(frozen=True, slots=True) +class _Models: + gpt: str + gpt_mini: str + claude: str + bedrock_claude: str + bedrock_titan: str + + +@dataclass(frozen=True, slots=True) +class _AccessCase: + name: str + allowed: Sequence[str] | None + requested: str + served: bool + + +_CASES: Final = ( + _AccessCase("openai_wildcard_denies_anthropic", ["openai/*"], "claude", False), + _AccessCase("exact_name_allows_itself", ["gpt"], "gpt", True), + _AccessCase("provider_wildcard_allows_bedrock", ["bedrock/*"], "bedrock_claude", True), + _AccessCase("family_wildcard_allows_its_family", ["bedrock/anthropic.*"], "bedrock_claude", True), + _AccessCase("family_wildcard_denies_another_family", ["bedrock/anthropic.*"], "bedrock_titan", False), + _AccessCase("unset_models_allow_everything", None, "gpt", True), + _AccessCase("empty_models_allow_everything", [], "gpt", True), +) + + +def _deployment(scenario: Scenario, gateway: Gateway, name: str) -> str: + created: Final = gateway.post( + "/model/new", + { + "model_name": name, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{PROVIDER_URL}/v1", + "api_key": "sk-fixture", + }, + "model_info": {"id": f"access-{uuid.uuid4().hex}"}, + }, + ) + scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"])) + return name + + +def _models(scenario: Scenario, gateway: Gateway) -> _Models: + tag: Final = uuid.uuid4().hex[:10] + return _Models( + gpt=_deployment(scenario, gateway, f"openai/gpt-{tag}"), + gpt_mini=_deployment(scenario, gateway, f"openai/gpt-mini-{tag}"), + claude=_deployment(scenario, gateway, f"anthropic/claude-{tag}"), + bedrock_claude=_deployment(scenario, gateway, f"bedrock/anthropic.claude-{tag}"), + bedrock_titan=_deployment(scenario, gateway, f"bedrock/amazon.titan-{tag}"), + ) + + +def _pick(models: _Models, alias: str) -> str: + return string_value(getattr(models, alias)) + + +def _completion() -> Reply: + return Reply( + body=json.dumps( + { + "id": f"chatcmpl-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "scripted"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, + } + ).encode() + ) + + +def _served(gateway: Gateway, provider: SharedProvider, key: str, model: str) -> None: + provider.expect(_completion()) + body: Final = gateway.chat(model, key=key, text=f"access {uuid.uuid4().hex}") + assert object_value(object_value(body["choices"][0])["message"])["content"] == "scripted", body + assert len(provider.received()) == 1 + + +def _denied(gateway: Gateway, provider: SharedProvider, key: str, model: str, denial: _Denial) -> None: + refused: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": "hi"}]}, key=key + ) + assert refused.status_code == 403, f"{refused.status_code} {refused.text}" + error: Final = object_value(JSON_OBJECT.validate_json(refused.content)["error"]) + assert error["type"] == denial, refused.text + assert error["param"] == "model", refused.text + assert error["code"] == "403", refused.text + message: Final = string_value(error["message"]) + assert "is not available for this API key" in message, message + assert "not allowed to access model" not in message, message + assert provider.received() == () + + +@pytest.mark.parametrize("case", _CASES, ids=[case.name for case in _CASES]) +def test_a_key_model_allow_list_decides_which_models_it_reaches( + gateway: Gateway, provider: SharedProvider, case: _AccessCase +) -> None: + with gateway.scenario() as scenario: + models: Final = _models(scenario, gateway) + allowed: Final = ( + None + if case.allowed is None + else [_pick(models, item) if hasattr(models, item) else item for item in case.allowed] + ) + key: Final = scenario.key(models=allowed) + requested: Final = _pick(models, case.requested) + if case.served: + _served(gateway, provider, key, requested) + else: + _denied(gateway, provider, key, requested, "key_model_access_denied") + + +def test_widening_a_key_allow_list_to_a_wildcard_takes_effect_on_the_next_request( + gateway: Gateway, provider: SharedProvider +) -> None: + with gateway.scenario() as scenario: + models: Final = _models(scenario, gateway) + key: Final = scenario.key(models=[models.gpt]) + _served(gateway, provider, key, models.gpt) + _denied(gateway, provider, key, models.gpt_mini, "key_model_access_denied") + gateway.post("/key/update", {"key": key, "models": ["openai/*"]}) + _served(gateway, provider, key, models.gpt) + _served(gateway, provider, key, models.gpt_mini) + _denied(gateway, provider, key, models.claude, "key_model_access_denied") + + +def test_a_team_allow_list_denies_its_keys_a_model_outside_it(gateway: Gateway, provider: SharedProvider) -> None: + with gateway.scenario() as scenario: + models: Final = _models(scenario, gateway) + team: Final = scenario.team(models=["openai/*"]) + key: Final = scenario.key(team_id=team) + _served(gateway, provider, key, models.gpt) + _denied(gateway, provider, key, models.claude, "team_model_access_denied") + + +def test_widening_a_team_allow_list_takes_effect_for_its_keys_on_the_next_request( + gateway: Gateway, provider: SharedProvider +) -> None: + with gateway.scenario() as scenario: + models: Final = _models(scenario, gateway) + team: Final = scenario.team(models=[models.gpt]) + key: Final = scenario.key(team_id=team) + _served(gateway, provider, key, models.gpt) + _denied(gateway, provider, key, models.gpt_mini, "team_model_access_denied") + gateway.post("/team/update", {"team_id": team, "models": ["openai/*"]}) + _served(gateway, provider, key, models.gpt) + _served(gateway, provider, key, models.gpt_mini) + _denied(gateway, provider, key, models.claude, "team_model_access_denied") diff --git a/tests/integration/authorization/test_websocket_rejection_logging.py b/tests/integration/authorization/test_websocket_rejection_logging.py new file mode 100644 index 00000000000..ea0a841a2b5 --- /dev/null +++ b/tests/integration/authorization/test_websocket_rejection_logging.py @@ -0,0 +1,317 @@ +import asyncio +import uuid +from collections.abc import Iterator, Mapping +from dataclasses import dataclass +from itertools import product +from types import MappingProxyType +from typing import Final + +import psutil +import pytest +import websockets +from integration._support.client import eventually, gateway_from_environment +from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process +from pydantic import JsonValue +from websockets.exceptions import InvalidStatus + +pytestmark: Final = pytest.mark.timeout(2 * graceful_stop_seconds() + 120) + +WEBSOCKET_ROUTES: Final = ( + "/v1/responses", + "/responses", + "/v1/realtime", + "/openai/v1/realtime", + "/realtime", + "/openai/v1/responses", + "/openai_passthrough/v1/responses", + "/deepgram/v1/listen", + "/deepgram/listen", + "/vertex_ai/live", +) +HTTP_ONLY_ROUTES: Final = ("/v1/traces", "/v1/logs", "/v1/chat/completions", "/nope") +HANDSHAKE_SHAPES: Final = ( + "no_header", + "empty_bearer", + "basic_scheme", + "lowercase_bearer", + "unknown_key", + "huge_key", + "duplicate_header", + "api_key_header", + "subprotocol_key", + "malformed_key", +) +BURST_SHAPES: Final = ("no_header", "unknown_key", "denied_model") +CRASH_MARKERS: Final = ( + "Exception in ASGI application", + "AttributeError: 'WebSocket' object has no attribute 'method'", + "ERROR: user_api_key_auth.py", +) +OTLP_LENS_MESSAGE: Final = "Send traces and logs directly to the Lens endpoint shown in Lens setup." +REALTIME_TRANSCRIPTION_MODEL: Final = "gpt-realtime-whisper" + +Headers = tuple[tuple[str, str], ...] + + +@dataclass(frozen=True, slots=True) +class Keys: + restricted: str + allowed_model: str + + +@dataclass(frozen=True, slots=True) +class Refusal: + status: int + window: str + + +@pytest.fixture(scope="module") +def owned(tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedProxy]: + with ( + gateway_from_environment() as base, + owned_proxy_process( + base, + tmp_path_factory.mktemp("websocket-rejection"), + {"LITELLM_DISABLE_NO_REDIS_WARNING": "true"}, + workers=2, + ) as proxy, + ): + yield proxy + stopped_log: Final = proxy.log.read_text(errors="replace") + assert _crash_lines(stopped_log) == (), _marker_lines(stopped_log) + + +@pytest.fixture(scope="module") +def keys(owned: OwnedProxy) -> Iterator[Keys]: + with owned.gateway.scenario() as scenario: + allowed: Final = scenario.model() + yield Keys(restricted=scenario.key(models=[allowed]), allowed_model=allowed) + + +def _tag(label: str) -> str: + return f"integration-{label}-{uuid.uuid4().hex[:12]}" + + +def _ws_url(owned: OwnedProxy, path: str, query: str) -> str: + base: Final = str(owned.gateway.client.base_url).rstrip("/").replace("http://", "ws://") + return f"{base}{path}?{query}" + + +def _log_after(owned: OwnedProxy, offset: int) -> str: + return owned.log.read_bytes()[offset:].decode(errors="replace") + + +def _crash_lines(window: str) -> tuple[str, ...]: + return tuple(marker for marker in CRASH_MARKERS if marker in window) + + +def _marker_lines(text: str) -> tuple[str, ...]: + return tuple(line for line in text.splitlines() if any(marker in line for marker in CRASH_MARKERS)) + + +def _shapes(keys: Keys) -> Mapping[str, Headers]: + unknown: Final = f"sk-{uuid.uuid4().hex}" + return MappingProxyType( + { + "no_header": (), + "empty_bearer": (("Authorization", "Bearer "),), + "basic_scheme": (("Authorization", "Basic aW50ZWdyYXRpb246eA=="),), + "lowercase_bearer": (("Authorization", f"bearer {unknown}"),), + "unknown_key": (("Authorization", f"Bearer {unknown}"),), + "huge_key": (("Authorization", f"Bearer sk-{'a' * 5120}"),), + "duplicate_header": (("Authorization", f"Bearer {unknown}"), ("Authorization", f"Bearer {unknown}")), + "api_key_header": (("api-key", unknown),), + "subprotocol_key": (("Sec-WebSocket-Protocol", f"openai-insecure-api-key.{unknown}"),), + "malformed_key": (("Authorization", "Bearer integration-not-a-virtual-key"),), + "denied_model": (("Authorization", f"Bearer {keys.restricted}"),), + } + ) + + +async def _refused_status(owned: OwnedProxy, path: str, headers: Headers, query: str) -> int: + with pytest.raises(InvalidStatus) as refused: + async with websockets.connect(_ws_url(owned, path, query), additional_headers=headers): + pass + return refused.value.response.status_code + + +async def _refused(owned: OwnedProxy, path: str, headers: Headers, query: str) -> Refusal: + offset: Final = owned.log.stat().st_size + status: Final = await _refused_status(owned, path, headers, query) + window: Final = eventually( + lambda: _log_after(owned, offset), + lambda text: f'"WebSocket {path}?{query}" 403' in text, + seconds=30, + ) + return Refusal(status, window) + + +@pytest.mark.parametrize("shape", HANDSHAKE_SHAPES) +@pytest.mark.parametrize("path", WEBSOCKET_ROUTES) +async def test_a_refused_handshake_answers_403_and_logs_no_asgi_crash( + owned: OwnedProxy, keys: Keys, path: str, shape: str +) -> None: + refusal: Final = await _refused(owned, path, _shapes(keys)[shape], f"model={_tag('ws')}") + assert refusal.status == 403, refusal.window + assert _crash_lines(refusal.window) == (), refusal.window + + +@pytest.mark.parametrize("path", WEBSOCKET_ROUTES) +async def test_a_denied_model_handshake_logs_one_warning_line(owned: OwnedProxy, keys: Keys, path: str) -> None: + tag: Final = _tag("denied") + refusal: Final = await _refused(owned, path, _shapes(keys)["denied_model"], f"model={tag}") + assert refusal.status == 403, refusal.window + assert refusal.window.count(f"Tried to access {tag}") == 1, refusal.window + assert _crash_lines(refusal.window) == (), refusal.window + + +async def test_a_realtime_denial_on_the_resolved_transcription_model_logs_one_warning_line( + owned: OwnedProxy, keys: Keys +) -> None: + tag: Final = _tag("transcription") + refusal: Final = await _refused( + owned, "/v1/realtime", _shapes(keys)["denied_model"], f"intent=transcription&tag={tag}" + ) + assert refusal.status == 403, refusal.window + assert refusal.window.count(f"Tried to access {REALTIME_TRANSCRIPTION_MODEL}") == 1, refusal.window + assert _crash_lines(refusal.window) == (), refusal.window + + +async def test_a_realtime_handshake_without_a_model_is_refused_without_a_denial_line( + owned: OwnedProxy, keys: Keys +) -> None: + refusal: Final = await _refused(owned, "/v1/realtime", _shapes(keys)["denied_model"], f"tag={_tag('nomodel')}") + assert refusal.status == 403, refusal.window + assert "Tried to access" not in refusal.window, refusal.window + assert _crash_lines(refusal.window) == (), refusal.window + + +@pytest.mark.parametrize("path", HTTP_ONLY_ROUTES) +async def test_a_websocket_handshake_on_an_http_only_route_is_refused_with_403(owned: OwnedProxy, path: str) -> None: + refusal: Final = await _refused(owned, path, (), f"model={_tag('httponly')}") + assert refusal.status == 403, refusal.window + assert _crash_lines(refusal.window) == (), refusal.window + + +def _http_window(owned: OwnedProxy, offset: int, access_line: str) -> str: + return eventually(lambda: _log_after(owned, offset), lambda text: access_line in text, seconds=30) + + +@pytest.mark.parametrize("content_type", ("application/json", "application/x-protobuf")) +@pytest.mark.parametrize("path", ("/v1/traces", "/v1/logs")) +def test_an_otlp_ingest_post_still_answers_410_in_the_caller_s_encoding( + owned: OwnedProxy, path: str, content_type: str +) -> None: + tag: Final = _tag("otlp") + offset: Final = owned.log.stat().st_size + response: Final = owned.gateway.client.post( + f"{path}?tag={tag}", content=b"", headers={"content-type": content_type} + ) + assert response.status_code == 410, response.text + assert response.headers["content-type"].split(";", 1)[0] == content_type, response.headers + assert OTLP_LENS_MESSAGE.encode() in response.content, response.content + window: Final = _http_window(owned, offset, f'"POST {path}?tag={tag} HTTP/1.1" 410') + assert _crash_lines(window) == (), window + + +def test_a_get_on_the_otlp_path_is_not_otlp_ingest_and_answers_a_plain_401(owned: OwnedProxy) -> None: + tag: Final = _tag("otlpget") + offset: Final = owned.log.stat().st_size + response: Final = owned.gateway.client.get( + f"/v1/traces?tag={tag}", headers={"content-type": "application/x-protobuf"} + ) + assert response.status_code == 401, response.text + assert response.headers["content-type"].split(";", 1)[0] == "application/json", response.headers + assert response.json()["error"]["code"] == "401", response.text + window: Final = _http_window(owned, offset, f'"GET /v1/traces?tag={tag} HTTP/1.1" 401') + assert _crash_lines(window) == (), window + + +def test_an_unknown_http_route_answers_a_plain_404(owned: OwnedProxy) -> None: + tag: Final = _tag("nope") + offset: Final = owned.log.stat().st_size + response: Final = owned.gateway.request("POST", f"/nope?tag={tag}", {}) + assert response.status_code == 404, response.text + assert response.json() == {"detail": "Not Found"}, response.text + window: Final = _http_window(owned, offset, f'"POST /nope?tag={tag} HTTP/1.1" 404') + assert _crash_lines(window) == (), window + + +def test_a_bad_request_on_a_plain_route_keeps_its_detail(owned: OwnedProxy, keys: Keys) -> None: + tag: Final = _tag("tokens") + offset: Final = owned.log.stat().st_size + response: Final = owned.gateway.request("POST", f"/utils/token_counter?tag={tag}", {"model": keys.allowed_model}) + assert response.status_code == 400, response.text + assert response.json() == {"detail": "prompt or messages or contents must be provided"}, response.text + window: Final = _http_window(owned, offset, f'"POST /utils/token_counter?tag={tag} HTTP/1.1" 400') + assert _crash_lines(window) == (), window + + +def _denied_body(path: str, model: str) -> Mapping[str, JsonValue]: + message: Final[dict[str, JsonValue]] = {"role": "user", "content": "integration denial"} + match path: + case "/v1/messages": + return {"model": model, "max_tokens": 4, "messages": [message]} + case "/v1/responses": + return {"model": model, "input": "integration denial"} + case _: + return {"model": model, "messages": [message]} + + +@pytest.mark.parametrize("path", ("/v1/chat/completions", "/v1/messages", "/v1/responses")) +def test_an_http_model_denial_logs_one_warning_line_and_no_traceback(owned: OwnedProxy, keys: Keys, path: str) -> None: + tag: Final = _tag("httpdenied") + offset: Final = owned.log.stat().st_size + response: Final = owned.gateway.request("POST", path, _denied_body(path, tag), key=keys.restricted) + assert response.status_code == 403, response.text + window: Final = eventually( + lambda: _log_after(owned, offset), lambda text: f"Tried to access {tag}" in text, seconds=30 + ) + assert window.count(f"Tried to access {tag}") == 1, window + assert _crash_lines(window) == (), window + + +def _workers(owned: OwnedProxy) -> tuple[psutil.Process, ...]: + return tuple(child for child in psutil.Process(owned.process.pid).children() if _is_worker(child)) + + +def _is_worker(child: psutil.Process) -> bool: + try: + return "spawn_main" in " ".join(child.cmdline()) and child.status() != psutil.STATUS_ZOMBIE + except psutil.Error: + return False + + +async def _burst(owned: OwnedProxy, shapes: Mapping[str, Headers], label: str) -> str: + plan: Final = tuple((path, shape, _tag(label)) for path, shape in product(WEBSOCKET_ROUTES, BURST_SHAPES)) + offset: Final = owned.log.stat().st_size + statuses: Final = await asyncio.gather( + *(_refused_status(owned, path, shapes[shape], f"model={tag}") for path, shape, tag in plan) + ) + assert tuple(statuses) == (403,) * len(plan), statuses + access_lines: Final = tuple(f'"WebSocket {path}?model={tag}" 403' for path, _, tag in plan) + window: Final = eventually( + lambda: _log_after(owned, offset), + lambda text: all(line in text for line in access_lines), + seconds=60, + ) + assert window.count(f"Tried to access integration-{label}-") == len(WEBSOCKET_ROUTES), window + return window + + +async def test_a_rejection_burst_survives_a_killed_worker_without_an_asgi_crash(owned: OwnedProxy, keys: Keys) -> None: + shapes: Final = _shapes(keys) + before: Final = await _burst(owned, shapes, "burst") + assert _crash_lines(before) == (), before + victim: Final = eventually(lambda: _workers(owned), lambda workers: len(workers) == 2)[0] + victim.kill() + during: Final = await _refused(owned, "/v1/responses", shapes["unknown_key"], f"model={_tag('during')}") + assert during.status == 403, during.window + assert _crash_lines(during.window) == (), during.window + eventually( + lambda: frozenset(worker.pid for worker in _workers(owned)), + lambda pids: len(pids) == 2 and victim.pid not in pids, + seconds=graceful_stop_seconds(), + ) + after: Final = await _burst(owned, shapes, "respawned") + assert _crash_lines(after) == (), after diff --git a/tests/integration/caching/test_embedding_file_block_cache.py b/tests/integration/caching/test_embedding_file_block_cache.py new file mode 100644 index 00000000000..9f753e996e8 --- /dev/null +++ b/tests/integration/caching/test_embedding_file_block_cache.py @@ -0,0 +1,216 @@ +from __future__ import annotations + +import asyncio +import json +import threading +import zlib +from collections.abc import Mapping +from hashlib import sha256 +from typing import Final + +import pytest +from openai import AsyncOpenAI, OpenAI +from openai.types import CreateEmbeddingResponse +from pydantic import JsonValue + +from tests.integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from tests.integration._support.database import read_rows +from tests.integration._support.wire import Reply, Request, Wire, wire_server + +BACKEND: Final = "gemini-embedding-001" +TARGET: Final = f"/models/{BACKEND}:batchEmbedContents" +METADATA: Final[dict[str, JsonValue]] = {"fps": 1.0, "start_offset": "0s", "end_offset": "3s"} +WIRE_METADATA: Final[dict[str, JsonValue]] = {"fps": 1.0, "startOffset": "0s", "endOffset": "3s"} +TEXT_PART: Final[dict[str, JsonValue]] = {"text": "a red bus"} +SETTLE_SECONDS: Final = 15 +SPEND_SQL: Final = ( + 'SELECT request_id, cache_hit, spend FROM "LiteLLM_SpendLogs" WHERE api_key = %s ORDER BY "startTime", request_id' +) + + +def clip(name: str) -> str: + return f"gs://scripted-bucket/cache/{name}.mp4" + + +def block(name: str, metadata: Mapping[str, JsonValue] | None = METADATA) -> dict[str, JsonValue]: + file: Final[dict[str, JsonValue]] = {"file_id": clip(name)} + return {"type": "file", "file": file if metadata is None else {**file, "video_metadata": dict(metadata)}} + + +def gcs_part(name: str, metadata: Mapping[str, JsonValue] | None = WIRE_METADATA) -> dict[str, JsonValue]: + part: Final[dict[str, JsonValue]] = {"file_data": {"mime_type": "video/mp4", "file_uri": clip(name)}} + return part if metadata is None else {**part, "video_metadata": dict(metadata)} + + +def list_value(value: JsonValue) -> list[JsonValue]: + assert isinstance(value, list), value + return value + + +def source_of(part: JsonValue) -> str: + item: Final = object_value(part) + if "text" in item: + return string_value(item["text"]) + return string_value(object_value(item["file_data"])["file_uri"]) + + +def vector(parts: list[JsonValue]) -> list[float]: + return [zlib.crc32("|".join(source_of(part) for part in parts).encode()) / 2**32, 0.5] + + +def request_parts(request: Request) -> list[list[JsonValue]]: + body: Final = object_value(json.loads(request.body)) + return [list_value(object_value(object_value(item)["content"])["parts"]) for item in list_value(body["requests"])] + + +def embed_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target.split("?")[0] == TARGET, request.target + embeddings: Final = [{"values": vector(parts)} for parts in request_parts(request)] + return Reply(body=json.dumps({"embeddings": embeddings}).encode()) + + +def wire_parts(wire: Wire) -> list[list[JsonValue]]: + received: Final = wire.drain() + assert len(received) == 1, [item.target for item in received] + return request_parts(received[0]) + + +def embeddings_of(answer: Mapping[str, JsonValue]) -> list[JsonValue]: + data: Final = [object_value(item) for item in list_value(answer["data"])] + assert [item["index"] for item in data] == list(range(len(data))), answer + return [item["embedding"] for item in data] + + +def gemini_model(scenario: Scenario, url: str) -> str: + return scenario.model(model=f"gemini/{BACKEND}", api_key="scripted-gemini-key", api_base=url) + + +def embed_body(model: str, elements: JsonValue) -> dict[str, JsonValue]: + return {"model": model, "input": elements} + + +def spend_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows(SPEND_SQL, (sha256(key.encode()).hexdigest(),)) + + +def served_from_cache( + gateway: Gateway, wire: Wire, body: Mapping[str, JsonValue], key: str +) -> dict[str, JsonValue] | None: + answer: Final = gateway.post("/v1/embeddings", body, key=key) + return None if wire.drain() else answer + + +def await_hit(gateway: Gateway, wire: Wire, body: Mapping[str, JsonValue], key: str) -> dict[str, JsonValue]: + served: Final = eventually( + lambda: served_from_cache(gateway, wire, body, key), lambda answer: answer is not None, SETTLE_SECONDS + ) + assert served is not None + return served + + +def proxy_root(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def test_block_embedding_is_served_from_the_cache_with_a_zero_spend_row(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + body: Final = embed_body(model, [block("one")]) + first: Final = gateway.request("POST", "/v1/embeddings", body, key=key) + assert first.status_code == 200, first.text + first_id: Final = first.headers["x-litellm-call-id"] + assert wire_parts(wire) == [[gcs_part("one")]] + served: Final = await_hit(gateway, wire, body, key) + assert embeddings_of(served) == embeddings_of(object_value(first.json())) == [vector([gcs_part("one")])] + rows: Final = eventually( + lambda: spend_rows(key), lambda found: any(row["cache_hit"] == "True" for row in found), seconds=70 + ) + by_id: Final = {string_value(row["request_id"]): row for row in rows} + assert by_id[first_id]["cache_hit"] != "True", rows + hits: Final = [row for row in rows if row["cache_hit"] == "True"] + assert hits and all(float(str(row["spend"])) == 0 for row in hits), rows + + +def test_cached_block_is_not_resent_when_a_new_text_joins_it(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + cached: Final = embed_body(model, [block("one")]) + gateway.post("/v1/embeddings", cached, key=key) + assert wire_parts(wire) == [[gcs_part("one")]] + await_hit(gateway, wire, cached, key) + mixed: Final = gateway.post("/v1/embeddings", embed_body(model, [block("one"), "a red bus"]), key=key) + assert wire_parts(wire) == [[TEXT_PART]] + assert embeddings_of(mixed) == [vector([gcs_part("one")]), vector([TEXT_PART])] + + +@pytest.mark.parametrize("names", (("a", "b", "c", "d", "e"), ("same", "same")), ids=("five-distinct", "two-identical")) +def test_block_lists_are_cached_whole(gateway: Gateway, names: tuple[str, ...]) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + body: Final = embed_body(model, [block(name) for name in names]) + first: Final = gateway.post("/v1/embeddings", body, key=key) + assert wire_parts(wire) == [[gcs_part(name)] for name in names] + assert embeddings_of(first) == [vector([gcs_part(name)]) for name in names] + assert embeddings_of(await_hit(gateway, wire, body, key)) == embeddings_of(first) + + +def test_provider_failure_is_not_cached(gateway: Gateway) -> None: + healed: Final = threading.Event() + + def respond(request: Request) -> Reply: + if healed.is_set(): + return embed_peer(request) + return Reply(status=500, body=b'{"error": {"message": "scripted outage"}}') + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + body: Final = embed_body(model, [block("flaky")]) + failed: Final = gateway.request("POST", "/v1/embeddings", body, key=key) + assert failed.status_code >= 500, failed.text + assert wire.drain(), "the failing call never reached the provider" + healed.set() + recovered: Final = gateway.post("/v1/embeddings", body, key=key) + assert wire_parts(wire) == [[gcs_part("flaky")]] + assert embeddings_of(await_hit(gateway, wire, body, key)) == embeddings_of(recovered) + + +def test_video_metadata_is_part_of_the_cache_key(gateway: Gateway) -> None: + two_seconds: Final[dict[str, JsonValue]] = {**METADATA, "end_offset": "2s"} + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + settled: Final = embed_body(model, [block("keyed")]) + gateway.post("/v1/embeddings", settled, key=key) + assert wire_parts(wire) == [[gcs_part("keyed")]] + await_hit(gateway, wire, settled, key) + variants: Final = ( + (block("keyed", two_seconds), gcs_part("keyed", {**WIRE_METADATA, "endOffset": "2s"})), + (block("keyed", None), gcs_part("keyed", None)), + ) + for variant, part in variants: + gateway.post("/v1/embeddings", embed_body(model, [variant]), key=key) + assert wire_parts(wire) == [[part]] + + +def test_openai_sdk_clients_fill_and_hit_the_cache(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + key: Final = scenario.key(models=[model]) + body: Final = embed_body(model, [block("sdk")]) + with OpenAI(base_url=f"{proxy_root(gateway)}/v1", api_key=key, max_retries=0) as client: + filled: Final = client.post("/embeddings", body=body, cast_to=CreateEmbeddingResponse) + assert wire_parts(wire) == [[gcs_part("sdk")]] + await_hit(gateway, wire, body, key) + + async def drive() -> CreateEmbeddingResponse: + async with AsyncOpenAI(base_url=f"{proxy_root(gateway)}/v1", api_key=key, max_retries=0) as client: + return await client.post("/embeddings", body=body, cast_to=CreateEmbeddingResponse) + + served: Final = asyncio.run(drive()) + assert served.data[0].embedding == filled.data[0].embedding == vector([gcs_part("sdk")]) + assert wire.drain() == (), "the cached answer reached the provider again" diff --git a/tests/integration/configuration/test_callback_leak_contracts.py b/tests/integration/configuration/test_callback_leak_contracts.py new file mode 100644 index 00000000000..9ae5179dad2 --- /dev/null +++ b/tests/integration/configuration/test_callback_leak_contracts.py @@ -0,0 +1,147 @@ +import json +import re +import uuid +from collections import Counter +from itertools import chain +from pathlib import Path +from typing import Final + +from pydantic import JsonValue + +from tests.integration._support.client import Gateway, object_value +from tests.integration._support.process import owned_proxy + +SAMPLES: Final = 4 +REQUESTS_PER_INTERVAL: Final = 5 +LEAK_MIN_NET_GROWTH: Final = 5 +LEAK_MIN_GROWING_INTERVALS: Final = 2 +ADDRESS: Final = re.compile(r" at 0x[0-9a-fA-F]+") +OBJECT: Final = re.compile(r"<([\w.]+) object") +BOUND_METHOD: Final = re.compile(r"bound method ([\w.]+)") + + +def _callback_type(text: str) -> str: + stripped: Final = ADDRESS.sub("", text) + if (instance := OBJECT.search(stripped)) is not None: + return instance.group(1).split(".")[-1] + if (method := BOUND_METHOD.search(stripped)) is not None: + return method.group(1) + return stripped.strip() + + +def _sample(candidate: Gateway) -> tuple[Counter[str], int]: + response: Final = candidate.request("GET", "/active/callbacks") + assert response.status_code == 200, response.text + body: Final = object_value(response.json()) + callbacks: Final = body["all_litellm_callbacks"] + alerting: Final = body["num_alerting"] + assert isinstance(callbacks, list) and isinstance(alerting, int), body + return Counter(_callback_type(str(callback)) for callback in callbacks), alerting + + +Samples = tuple[Counter[str], ...] + + +def _kinds(samples: Samples) -> frozenset[str]: + return frozenset(chain.from_iterable(samples)) + + +def _series(samples: Samples, kind: str) -> tuple[int, ...]: + return tuple(sample.get(kind, 0) for sample in samples) + + +def _deltas(series: tuple[int, ...]) -> tuple[int, ...]: + return tuple(after - before for before, after in zip(series, series[1:])) + + +def _grows(series: tuple[int, ...]) -> bool: + deltas: Final = _deltas(series) + return ( + all(delta >= 0 for delta in deltas) + and series[-1] - series[0] >= LEAK_MIN_NET_GROWTH + and sum(1 for delta in deltas if delta > 0) >= LEAK_MIN_GROWING_INTERVALS + ) + + +def _grows_only_in_last_interval(series: tuple[int, ...]) -> bool: + deltas: Final = _deltas(series) + return all(delta >= 0 for delta in deltas) and [ + index for index, delta in enumerate(deltas) if delta > 0 + ] == [len(deltas) - 1] + + +def _leaking(samples: Samples) -> dict[str, tuple[int, ...]]: + return {kind: _series(samples, kind) for kind in sorted(_kinds(samples)) if _grows(_series(samples, kind))} + + +def _config(directory: Path, upstream_url: str, model: str, extra: dict[str, JsonValue]) -> Path: + config: Final = directory / f"callback_leak_{uuid.uuid4().hex}.yaml" + general_settings: Final = extra.get("general_settings") + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": model, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{upstream_url}/v1", + "api_key": "integration-provider-key", + }, + } + ], + "router_settings": extra.get("router_settings", {}), + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + **(general_settings if isinstance(general_settings, dict) else {}), + }, + } + ) + ) + return config + + +def _interval(candidate: Gateway, model: str, index: int) -> tuple[Counter[str], int]: + for request in range(REQUESTS_PER_INTERVAL if index else 0): + assert candidate.chat(model, text=f"leak probe {index} {request}")["model"] == model + return _sample(candidate) + + +def _sample_under_traffic(candidate: Gateway, model: str) -> tuple[Samples, tuple[int, ...]]: + taken: Final = tuple(_interval(candidate, model, index) for index in range(SAMPLES)) + callbacks: Final = tuple(sample for sample, _ in taken) + late_growth: Final = any(_grows_only_in_last_interval(_series(callbacks, kind)) for kind in _kinds(callbacks)) + confirmed: Final = taken + ((_interval(candidate, model, SAMPLES),) if late_growth else ()) + return tuple(sample for sample, _ in confirmed), tuple(alerting for _, alerting in confirmed) + + +def test_callback_registry_does_not_grow_with_traffic(gateway: Gateway, tmp_path: Path) -> None: + model: Final = f"integration-callback-leak-{uuid.uuid4().hex}" + config: Final = _config(tmp_path, gateway.upstream_url, model, {}) + with owned_proxy(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=config) as candidate: + samples, _ = _sample_under_traffic(candidate, model) + assert sum(samples[0].values()) > 0 + assert _leaking(samples) == {}, samples + + +def test_callback_registry_does_not_grow_under_latency_routing_with_alerting(gateway: Gateway, tmp_path: Path) -> None: + model: Final = f"integration-callback-leak-{uuid.uuid4().hex}" + config: Final = _config( + tmp_path, + gateway.upstream_url, + model, + { + "router_settings": {"routing_strategy": "latency-based-routing"}, + "general_settings": { + "alert_to_webhook_url": {"llm_exceptions": "http://127.0.0.1:9/integration-alerts"}, + "alert_types": ["llm_exceptions", "db_exceptions"], + }, + }, + ) + with owned_proxy(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=config) as candidate: + samples, alerts = _sample_under_traffic(candidate, model) + assert sum(samples[0].values()) > 0 + assert any(kind.startswith("LowestLatencyLoggingHandler") for kind in samples[0]), samples[0] + assert _leaking(samples) == {}, samples + assert len(set(alerts)) == 1, alerts diff --git a/tests/integration/configuration/test_config_declared_model_behaviours.py b/tests/integration/configuration/test_config_declared_model_behaviours.py new file mode 100644 index 00000000000..6b1f25c35a3 --- /dev/null +++ b/tests/integration/configuration/test_config_declared_model_behaviours.py @@ -0,0 +1,162 @@ +from __future__ import annotations + +import json +import uuid +from collections.abc import Iterator +from pathlib import Path +from typing import Final + +import pytest +import yaml +from pydantic import JsonValue, TypeAdapter + +from tests.integration._support.client import JSON_OBJECT, Gateway, gateway_from_environment, object_value, string_value +from tests.integration._support.process import owned_proxy +from tests.integration._support.provider import PROVIDER_URL, SharedProvider +from tests.integration._support.wire import Reply + +_TAGGED_GROUP: Final = "tag-filtered-group" +_TAGGED_IDS: Final = frozenset({"tag-filtered-team-a", "tag-filtered-team-b"}) +_VISION_MODEL: Final = "llava-hf" +_ITEMS: Final = TypeAdapter(list[dict[str, JsonValue]]) + + +def _config(directory: Path) -> Path: + configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + configuration["model_list"] = [ + *configuration["model_list"], + *( + { + "model_name": _TAGGED_GROUP, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "sk-fixture", + "api_base": f"{PROVIDER_URL}/v1", + "tags": [tag], + }, + "model_info": {"id": identity}, + } + for tag, identity in (("teamA", "tag-filtered-team-a"), ("teamB", "tag-filtered-team-b")) + ), + { + "model_name": _VISION_MODEL, + "litellm_params": { + "model": "openai/llava-hf", + "api_key": "sk-fixture", + "api_base": "http://127.0.0.1:9/v1", + }, + "model_info": {"supports_vision": True}, + }, + ] + configuration.setdefault("router_settings", {})["enable_tag_filtering"] = True + path: Final = directory / "config-declared-models.yaml" + path.write_text(yaml.safe_dump(configuration)) + return path + + +@pytest.fixture(scope="module") +def declared(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("config-declared-models") + with ( + gateway_from_environment() as shared, + owned_proxy( + shared, + directory, + {"OPENAI_API_KEY": "sk-fixture", "OPENAI_BASE_URL": f"{PROVIDER_URL}/v1", "GCS_FLUSH_INTERVAL": "1"}, + config=_config(directory), + remove_environment=("GCS_BUCKET_NAME", "OPENAI_API_BASE"), + ) as owned, + ): + yield owned + + +def test_model_info_reports_the_vision_capability_declared_in_the_config(declared: Gateway) -> None: + listing: Final = _ITEMS.validate_python(declared.get("/model/info")["data"]) + vision: Final = [item for item in listing if item["model_name"] == _VISION_MODEL] + assert len(vision) == 1, [item["model_name"] for item in listing] + assert object_value(vision[0]["model_info"])["supports_vision"] is True, vision[0] + + +def test_an_untagged_request_is_served_by_a_group_whose_deployments_are_all_tagged( + declared: Gateway, provider: SharedProvider +) -> None: + provider.expect( + Reply( + body=json.dumps( + { + "id": f"chatcmpl-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "tagged"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4}, + } + ).encode() + ) + ) + response: Final = declared.request( + "POST", + "/v1/chat/completions", + {"model": _TAGGED_GROUP, "messages": [{"role": "user", "content": f"untagged {uuid.uuid4().hex}"}]}, + ) + assert response.status_code == 200, response.text + assert response.headers["x-litellm-model-id"] in _TAGGED_IDS, dict(response.headers) + message: Final = object_value(object_value(JSON_OBJECT.validate_json(response.content)["choices"][0])["message"]) + assert message["content"] == "tagged" + assert [request.target for request in provider.received()] == ["/v1/chat/completions"] + + +def test_a_moderation_request_without_a_model_reaches_the_provider_default( + declared: Gateway, provider: SharedProvider +) -> None: + provider.expect( + Reply( + body=json.dumps( + { + "id": f"modr-{uuid.uuid4().hex}", + "model": "omni-moderation-latest", + "results": [ + {"flagged": True, "categories": {"violence": True}, "category_scores": {"violence": 0.9}} + ], + } + ).encode() + ) + ) + phrase: Final = f"I want to harm someone {uuid.uuid4().hex}" + response: Final = declared.request("POST", "/moderations", {"input": phrase}) + assert response.status_code == 200, response.text + body: Final = JSON_OBJECT.validate_json(response.content) + assert body["model"] == "omni-moderation-latest", body + assert object_value(_ITEMS.validate_python(body["results"])[0])["flagged"] is True, body + sent: Final = provider.received() + assert [request.target for request in sent] == ["/v1/moderations"] + payload: Final = JSON_OBJECT.validate_json(sent[0].body) + assert payload == {"input": phrase}, payload + + +def test_key_health_reports_an_unconfigured_key_logging_callback_as_unhealthy(declared: Gateway) -> None: + with declared.scenario() as scenario: + key: Final = scenario.key( + metadata={ + "logging": [ + { + "callback_name": "gcs_bucket", + "callback_type": "success_and_failure", + "callback_vars": { + "gcs_bucket_name": "key-logging-project1", + "gcs_path_service_account": "bad-service-account", + }, + } + ] + } + ) + health: Final = declared.request("POST", "/key/health", {}, key=key) + assert health.status_code == 200, health.text + body: Final = JSON_OBJECT.validate_json(health.content) + assert "key" in body, body + status: Final = object_value(body["logging_callbacks"]) + assert status["callbacks"] == ["gcs_bucket"], status + assert status["status"] == "unhealthy", status + assert "GCS_BUCKET_NAME is not set in the environment" in string_value(status["details"]), status diff --git a/tests/integration/configuration/test_multi_worker_slow_import_boot.py b/tests/integration/configuration/test_multi_worker_slow_import_boot.py new file mode 100644 index 00000000000..a90160e5e80 --- /dev/null +++ b/tests/integration/configuration/test_multi_worker_slow_import_boot.py @@ -0,0 +1,45 @@ +import os +import re +from pathlib import Path +from typing import Final + +import pytest +from integration._support.client import Gateway, eventually +from integration._support.process import graceful_stop_seconds, owned_proxy_process + +WORKERS: Final = 2 +DEPLOYMENT_HEALTHCHECK_SECONDS: Final = 5 +WORKER_IMPORT_DELAY_SECONDS: Final = 3 * DEPLOYMENT_HEALTHCHECK_SECONDS +STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +DIED_WORKER: Final = re.compile(r"Child process \[(\d+)\] died") +SLOW_WORKER_HOOK: Final = f"""\ +import sys +import time + +if "--multiprocessing-fork" in sys.argv: + time.sleep({WORKER_IMPORT_DELAY_SECONDS}) +""" + + +def _slow_worker_environment(directory: Path) -> dict[str, str]: + hook: Final = directory / "slow_worker" + hook.mkdir() + (hook / "sitecustomize.py").write_text(SLOW_WORKER_HOOK) + return { + "PYTHONPATH": os.pathsep.join((str(hook), os.environ.get("PYTHONPATH", ""))), + "TIMEOUT_WORKER_HEALTHCHECK": str(DEPLOYMENT_HEALTHCHECK_SECONDS), + } + + +@pytest.mark.timeout(2 * graceful_stop_seconds() + 60) +def test_owned_proxy_workers_outlast_the_deployments_healthcheck_default(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy_process(gateway, tmp_path, _slow_worker_environment(tmp_path), workers=WORKERS) as owned: + started: Final = eventually( + lambda: STARTED_WORKER.findall(owned.log.read_text()), + lambda pids: len(pids) >= WORKERS, + seconds=graceful_stop_seconds(), + ) + assert owned.gateway.request("GET", "/health/readiness").status_code == 200 + log: Final = owned.log.read_text() + assert len(started) == WORKERS, log + assert DIED_WORKER.findall(log) == [], log diff --git a/tests/integration/database/test_lens_repository.py b/tests/integration/database/test_lens_repository.py index a46e37668ab..6349da7ab2a 100644 --- a/tests/integration/database/test_lens_repository.py +++ b/tests/integration/database/test_lens_repository.py @@ -1,6 +1,7 @@ import asyncio import os -from collections.abc import AsyncIterator +from collections.abc import AsyncGenerator, AsyncIterator +from contextlib import asynccontextmanager from datetime import datetime, timedelta, timezone from pathlib import Path from types import SimpleNamespace @@ -15,6 +16,7 @@ from fastapi import HTTPException from prisma import Prisma from psycopg import sql from pydantic import TypeAdapter +from typing_extensions import LiteralString from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.db.prisma_client import PrismaWrapper @@ -35,7 +37,7 @@ from litellm.proxy.lens.models import ( TraceIdentity, Worker, ) -from litellm.proxy.lens.repository import LensRepository, WriterDatabase +from litellm.proxy.lens.repository import Database, LensRepository, Row, WriterDatabase from litellm.proxy.lens.state import cancel_job, claim_job, current_job, due_at, end_job, queue_job, replace_job @@ -379,7 +381,8 @@ async def test_managed_registration_is_atomic_and_keeps_the_original_worker_id(l @pytest.mark.asyncio async def test_claim_pages_only_yield_work_the_worker_can_claim(lens_db: Prisma) -> None: - now: Final = datetime.now(timezone.utc) + clock: Final = datetime.now(timezone.utc) + now: Final = clock.replace(microsecond=clock.microsecond // 1000 * 1000) scope: Final = Scope(team_id=uuid4().hex) repo: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db))) prefix: Final = uuid4().hex @@ -769,6 +772,156 @@ async def test_delayed_progress_cannot_replace_a_newer_checkpoint(lens_db: Prism await lens_db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', claimed.id) +class ProgressInterleavingDatabase: + def __init__( + self, + database: Database, + read: asyncio.Future[int], + resume: asyncio.Event, + committed: asyncio.Event, + ) -> None: + self.database: Final = database + self.read: Final = read + self.resume: Final = resume + self.committed: Final = committed + + async def query_raw(self, query: LiteralString, *args: object) -> object: + rows: Final = await self.database.query_raw(query, *args) + if query == 'SELECT data FROM "LiteLLM_Lens" WHERE id=$1' and not self.read.done(): + backend: Final = TypeAdapter(tuple[Row, ...]).validate_python( + await self.database.query_raw("SELECT to_jsonb(pg_backend_pid()) AS data") + ) + self.read.set_result(TypeAdapter(int).validate_python(backend[0].data)) + await self.resume.wait() + return rows + + async def execute_raw(self, query: LiteralString, *args: object) -> int: + return await self.database.execute_raw(query, *args) + + @asynccontextmanager + async def transaction(self) -> AsyncGenerator[Database]: + async with self.database.transaction() as database: + yield ProgressInterleavingDatabase(database, self.read, self.resume, self.committed) + self.committed.set() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("renewing", (False, True)) +async def test_budget_reservation_survives_competing_progress(lens_db: Prisma, renewing: bool) -> None: + from litellm.proxy.lens.inference import ( + BUDGET_LEASE, + renew_budget_reservation, + reserve_attempt, + wait_for_reservation, + ) + from litellm.proxy.lens.models import BudgetReservation + from tests.unit.proxy.lens.test_state import lens, worker + + now: Final = datetime.now(timezone.utc) + claimed: Final = claim_job(queue_job(lens(), now, uuid4().hex), worker(), now).model_copy( + update={"id": uuid4().hex, "budget_month": now.strftime("%Y-%m")} + ) + job: Final = claimed.jobs[0] + hold: Final = BudgetReservation( + id=uuid4().hex, job_id=job.id, amount=1, month=claimed.budget_month, expires_at=now + BUDGET_LEASE + ) + database: Final = WriterDatabase(PrismaWrapper(lens_db)) + repo: Final = LensRepository(database) + read: Final[asyncio.Future[int]] = asyncio.get_running_loop().create_future() + resume: Final = asyncio.Event() + committed: Final = asyncio.Event() + competing: Final = ProgressInterleavingDatabase(database, read, resume, committed) + await repo.create(claimed.model_copy(update={"reservations": (hold,) if renewing else ()})) + admitted: Final = asyncio.Event() + admitted.set() + operation: Final = asyncio.create_task( + renew_budget_reservation(LensRepository(competing), claimed.id, hold.id, admitted) + if renewing + else wait_for_reservation( + LensRepository(competing), + claimed.id, + hold.id, + lambda current: reserve_attempt(current, job, worker().id, hold, now), + ) + ) + + async def write_progress() -> Lens | None: + await read + return await repo.progress(claimed.id, job, Progress(stage="Reviewing traces concurrently")) + + progress: Final = asyncio.create_task(write_progress()) + try: + async with asyncio.timeout(45): + blocker: Final = await read + async with asyncio.timeout(5): + while not await lens_db.query_raw( + "SELECT pid FROM pg_stat_activity WHERE $1::int=ANY(pg_blocking_pids(pid))", blocker + ): + assert not progress.done(), "Progress committed before the reservation released its lock" + await asyncio.sleep(0.01) + assert not progress.done() + assert not committed.is_set() + resume.set() + await committed.wait() + updated: Final = await progress + assert updated is not None + assert updated.jobs[0].stage == "Reviewing traces concurrently" + assert tuple(reservation.id for reservation in updated.reservations) == (hold.id,) + if not renewing: + await operation + stored: Final = await repo.get(claimed.id) + assert stored is not None + assert tuple(reservation.id for reservation in stored.reservations) == (hold.id,) + assert stored.spent == claimed.spent + assert stored.jobs[0].cost == 0 + assert stored.jobs[0].stage == "Reviewing traces concurrently" + if renewing: + assert stored.reservations[0].expires_at is not None + assert stored.reservations[0].expires_at > now + BUDGET_LEASE + finally: + operation.cancel() + progress.cancel() + await asyncio.gather(operation, progress, return_exceptions=True) + await lens_db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', claimed.id) + + +@pytest.mark.asyncio +async def test_parallel_reservations_release_the_lock_while_waiting_for_budget(lens_db: Prisma) -> None: + from litellm.proxy.lens.inference import reserve_amount, settle_amount, wait_for_reservation + from litellm.proxy.lens.models import BudgetReservation + from tests.unit.proxy.lens.test_state import lens + + now: Final = datetime.now(timezone.utc) + original: Final = lens() + queued: Final = queue_job(original, now, uuid4().hex).model_copy( + update={"id": uuid4().hex, "settings": original.settings.model_copy(update={"monthly_budget": 2})} + ) + repo: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db))) + holds: Final = tuple( + BudgetReservation(id=uuid4().hex, job_id=queued.jobs[0].id, amount=1, month=queued.budget_month) + for _ in range(8) + ) + await repo.create(queued) + try: + + async def analyze(hold: BudgetReservation) -> None: + await wait_for_reservation(repo, queued.id, hold.id, lambda current: reserve_amount(current, hold)) + stored: Final = await repo.get(queued.id) + assert stored is not None and hold in stored.reservations + assert stored.spent + sum(reservation.amount for reservation in stored.reservations) <= 2 + assert await repo.update_locked(queued.id, lambda current: settle_amount(current, hold.id, 0.125, None)) + + async with asyncio.timeout(15): + await asyncio.gather(*(analyze(hold) for hold in holds)) + stored: Final = await repo.get(queued.id) + assert stored is not None + assert stored.spent == 1 + assert stored.jobs[0].cost == 1 + assert stored.reservations == () + finally: + await lens_db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', queued.id) + + @pytest.mark.asyncio async def test_locked_settlement_charges_every_concurrent_call_exactly_once(lens_db: Prisma) -> None: from litellm.proxy.lens.inference import settle_amount diff --git a/tests/integration/management/test_batch_output_file_listing.py b/tests/integration/management/test_batch_output_file_listing.py new file mode 100644 index 00000000000..26d1ae8de46 --- /dev/null +++ b/tests/integration/management/test_batch_output_file_listing.py @@ -0,0 +1,651 @@ +from __future__ import annotations + +import contextlib +import datetime +import json +import socket +import socketserver +import ssl +import threading +import uuid +from collections.abc import Generator +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from typing import Final + +import httpx +from cryptography import x509 +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import ec +from cryptography.x509.oid import NameOID +from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, TextResponse +from pydantic import JsonValue + +from litellm.litellm_core_utils.cloud_storage_security import BEDROCK_MANAGED_S3_OUTPUT_PREFIX + +OUTPUT_BYTES: Final = 4096 +ERROR_BYTES: Final = 512 +BEDROCK_MODEL: Final = "bedrock/anthropic.claude-3-haiku-20240307-v1:0" +BEDROCK_MODEL_ID: Final = "anthropic.claude-3-haiku-20240307-v1:0" +BEDROCK_REGION: Final = "us-east-1" +BEDROCK_AUTHORITY: Final = f"bedrock.{BEDROCK_REGION}.amazonaws.com:443" +BEDROCK_BUCKET: Final = "integration-batch-listing-bucket" +BEDROCK_ROLE_ARN: Final = "arn:aws:iam::123456789012:role/integration-batch-role" +BEDROCK_JOB_ARN_PREFIX: Final = f"arn:aws:bedrock:{BEDROCK_REGION}:123456789012:model-invocation-job/" +BEDROCK_LAST_MODIFIED: Final = "Thu, 02 Oct 2025 12:00:00 GMT" +BEDROCK_OUTPUT_CONTENT: Final = b'{"recordId":"req-1","modelOutput":{}}\n' +_BATCH_PROCESSED_SQL: Final = 'SELECT batch_processed FROM "LiteLLM_ManagedObjectTable" WHERE unified_object_id=%s' + + +@dataclass(frozen=True, slots=True) +class _OpenAIBatch: + model: str + owner_key: str + unrelated_key: str | None + batch_id: str + input_file_id: str + scenario: ScenarioHandle + + +def _output_content(model: str) -> str: + return ( + json.dumps( + { + "id": "batch_req_$REQUEST_ID", + "custom_id": "req-1", + "response": { + "status_code": 200, + "request_id": "$REQUEST_ID", + "body": { + "id": "chatcmpl-$REQUEST_ID", + "object": "chat.completion", + "model": model, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 2, "completion_tokens": 1, "total_tokens": 3}, + }, + }, + "error": None, + }, + separators=(",", ":"), + ) + + "\n" + ) + + +def _error_content() -> str: + return ( + json.dumps( + { + "id": "batch_req_$REQUEST_ID", + "custom_id": "req-2", + "response": {"status_code": 400, "body": {"error": {"message": "rejected"}}}, + "error": {"code": "bad_request", "message": "rejected"}, + }, + separators=(",", ":"), + ) + + "\n" + ) + + +def _batch_routes( + model: str, + *, + output_file: bool = True, + metadata_fails: bool = False, +) -> RoutedResponse: + completed: Final[dict[str, JsonValue]] = { + "id": "batch-$REQUEST_ID", + "object": "batch", + "endpoint": "/v1/chat/completions", + "errors": None, + "input_file_id": "file-in-$REQUEST_ID", + "completion_window": "24h", + "status": "completed", + "output_file_id": "file-out-$REQUEST_ID" if output_file else None, + "error_file_id": "file-err-$REQUEST_ID", + "created_at": 1, + "in_progress_at": 1, + "completed_at": 1, + "expires_at": 1, + "request_counts": {"total": 1 if not output_file else 2, "completed": 1 if output_file else 0, "failed": 1}, + "metadata": None, + } + output_metadata: Final = JsonResponse( + content_type="application/json", + body=( + {"error": "provider metadata unavailable"} + if metadata_fails + else { + "id": "file-out-$REQUEST_ID", + "object": "file", + "purpose": "batch_output", + "bytes": OUTPUT_BYTES, + "created_at": 1, + "filename": "output.jsonl", + "status": "processed", + } + ), + status=500 if metadata_fails else 200, + ) + return RoutedResponse( + content_type="application/x-routed", + routes={ + "POST /files": JsonResponse( + content_type="application/json", + body={ + "id": "file-in-$REQUEST_ID", + "object": "file", + "purpose": "batch", + "bytes": 100, + "created_at": 1, + "filename": "input.jsonl", + "status": "processed", + }, + ), + "POST /batches": JsonResponse( + content_type="application/json", + body={**completed, "status": "validating", "output_file_id": None, "error_file_id": None}, + ), + "GET /batches/batch-$REQUEST_ID": JsonResponse(content_type="application/json", body=completed), + "GET /files/file-out-$REQUEST_ID": output_metadata, + "GET /files/file-err-$REQUEST_ID": JsonResponse( + content_type="application/json", + body={ + "id": "file-err-$REQUEST_ID", + "object": "file", + "purpose": "batch_output", + "bytes": ERROR_BYTES, + "created_at": 1, + "filename": "errors.jsonl", + "status": "processed", + }, + ), + "GET /files/file-out-$REQUEST_ID/content": TextResponse( + content_type="application/jsonl", + body=_output_content(model), + ), + "GET /files/file-err-$REQUEST_ID/content": TextResponse( + content_type="application/jsonl", + body=_error_content(), + ), + }, + ) + + +def _create_batch( + scenario: Scenario, + routes: RoutedResponse, + *, + unrelated_user: bool = False, +) -> _OpenAIBatch: + handle: Final = register_scenario(f"batch-output-listing-{uuid.uuid4().hex}", routes) + scenario.cleanups.callback(delete_scenario, handle) + model: Final = scenario.model(api_base=handle.api_base()) + owner_id: Final = scenario.user(user_role="internal_user") + owner_key: Final = scenario.key(user_id=owner_id, models=[model]) + unrelated_key: Final = ( + scenario.key(user_id=scenario.user(user_role="internal_user"), models=[model]) if unrelated_user else None + ) + input_content: Final = json.dumps( + { + "custom_id": "req-1", + "method": "POST", + "url": "/v1/chat/completions", + "body": {"model": model, "messages": [{"role": "user", "content": "ping"}], "max_tokens": 8}, + } + ).encode() + uploaded: Final = scenario.gateway.request_multipart( + "/v1/files", + {"purpose": "batch", "target_model_names": model}, + {"file": ("input.jsonl", input_content, "application/jsonl")}, + key=owner_key, + ) + assert uploaded.status_code == 200, uploaded.text + input_file_id: Final = string_value(JSON_OBJECT.validate_json(uploaded.content)["id"]) + created: Final = scenario.gateway.request( + "POST", + "/v1/batches", + { + "input_file_id": input_file_id, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": model, + }, + key=owner_key, + ) + assert created.status_code == 200, created.text + batch_id: Final = string_value(JSON_OBJECT.validate_json(created.content)["id"]) + return _OpenAIBatch(model, owner_key, unrelated_key, batch_id, input_file_id, handle) + + +def _retrieve_batch(gateway: Gateway, batch: _OpenAIBatch) -> dict[str, JsonValue]: + response: Final = gateway.request("GET", f"/v1/batches/{batch.batch_id}", key=batch.owner_key) + assert response.status_code == 200, response.text + return JSON_OBJECT.validate_json(response.content) + + +def _list_files( + gateway: Gateway, + key: str, + *, + purpose: str | None = None, +) -> tuple[dict[str, JsonValue], ...]: + response: Final = gateway.request( + "GET", + "/v1/files", + key=key, + params={"purpose": purpose} if purpose is not None else None, + ) + assert response.status_code == 200, response.text + values: Final = JSON_OBJECT.validate_json(response.content)["data"] + assert isinstance(values, list), response.text + return tuple(object_value(value) for value in values) + + +def _metadata_hit_count(gateway: Gateway, scenario: ScenarioHandle) -> int: + response: Final = httpx.get( + f"{gateway.upstream_url}/__observations", + timeout=15, + trust_env=False, + ) + assert response.status_code == 200, response.text + requests: Final = JSON_OBJECT.validate_json(response.content)["requests"] + assert isinstance(requests, list), response.text + return sum(1 for request in requests if _is_metadata_hit(request, scenario.scenario_id)) + + +def _is_metadata_hit(request: JsonValue, scenario_id: str) -> bool: + if not isinstance(request, dict): + return False + path: Final = request.get("path") + return request.get("method") == "GET" and isinstance(path, str) and path.endswith(f"/files/file-out-{scenario_id}") + + +def test_output_file_lists_after_owner_retrieves_completed_batch(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + batch: Final = _create_batch(scenario, _batch_routes(model="gpt-4o-mini")) + retrieved: Final = _retrieve_batch(gateway, batch) + assert retrieved["status"] == "completed", retrieved + output_id: Final = string_value(retrieved["output_file_id"]) + error_id: Final = string_value(retrieved["error_file_id"]) + files: Final = _list_files(gateway, batch.owner_key) + output: Final = next((file for file in files if file.get("id") == output_id), None) + assert output is not None, f"Completed batch output {output_id} is absent from GET /v1/files" + assert (output["purpose"], output["bytes"]) == ("batch_output", OUTPUT_BYTES), output + output_files: Final = _list_files(gateway, batch.owner_key, purpose="batch_output") + assert {string_value(file["id"]) for file in output_files} == {output_id, error_id}, output_files + assert all(file["purpose"] == "batch_output" for file in output_files), output_files + input_files: Final = _list_files(gateway, batch.owner_key, purpose="batch") + assert tuple(string_value(file["id"]) for file in input_files) == (batch.input_file_id,), input_files + + +def test_poller_registers_listable_output_files_without_a_batch_retrieve(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + batch: Final = _create_batch( + scenario, + _batch_routes(model="gpt-4o-mini"), + unrelated_user=True, + ) + assert batch.unrelated_key is not None + output_files: Final = eventually( + lambda: _list_files(gateway, batch.owner_key, purpose="batch_output"), + lambda values: len(values) == 2, + seconds=60, + ) + assert {file["purpose"] for file in output_files} == {"batch_output"}, output_files + unrelated_files: Final = _list_files(gateway, batch.unrelated_key, purpose="batch_output") + assert unrelated_files == (), unrelated_files + retrieved: Final = _retrieve_batch(gateway, batch) + output_id: Final = string_value(retrieved["output_file_id"]) + error_id: Final = string_value(retrieved["error_file_id"]) + assert {string_value(file["id"]) for file in output_files} == {output_id, error_id}, output_files + + +def test_error_file_only_batch_still_lists_its_error_file(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + batch: Final = _create_batch( + scenario, + _batch_routes(model="gpt-4o-mini", output_file=False), + ) + error_files: Final = eventually( + lambda: _list_files(gateway, batch.owner_key, purpose="batch_output"), + lambda values: len(values) == 1, + seconds=60, + ) + retrieved: Final = _retrieve_batch(gateway, batch) + assert retrieved["output_file_id"] is None, retrieved + error_id: Final = string_value(retrieved["error_file_id"]) + assert tuple(string_value(file["id"]) for file in error_files) == (error_id,), error_files + assert error_files[0]["purpose"] == "batch_output", error_files + + +def test_output_file_lists_with_basic_details_when_provider_file_lookup_fails(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + batch: Final = _create_batch( + scenario, + _batch_routes(model="gpt-4o-mini", metadata_fails=True), + ) + retrieved: Final = _retrieve_batch(gateway, batch) + output_id: Final = string_value(retrieved["output_file_id"]) + output: Final = next( + (file for file in _list_files(gateway, batch.owner_key) if file.get("id") == output_id), + None, + ) + assert output is not None, f"Completed batch output {output_id} is absent after provider metadata failure" + assert (output["purpose"], output["filename"]) == ("batch_output", f"file-out-{batch.scenario.scenario_id}"), ( + output + ) + assert "litellm_details_fallback" not in output, output + + +def test_fallback_output_file_details_refresh_once_provider_recovers(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + batch: Final = _create_batch( + scenario, + _batch_routes(model="gpt-4o-mini", metadata_fails=True), + unrelated_user=True, + ) + retrieved: Final = _retrieve_batch(gateway, batch) + output_id: Final = string_value(retrieved["output_file_id"]) + listed_files: Final = _list_files(gateway, batch.owner_key, purpose="batch_output") + basic_output: Final = next((file for file in listed_files if file.get("id") == output_id), None) + assert basic_output is not None, f"Completed batch output {output_id} is absent after provider metadata failure" + assert (basic_output["purpose"], basic_output["filename"]) == ( + "batch_output", + f"file-out-{batch.scenario.scenario_id}", + ), basic_output + assert basic_output["bytes"] != OUTPUT_BYTES, basic_output + assert "litellm_details_fallback" not in basic_output, basic_output + + processed_batch_rows: Final = eventually( + lambda: read_rows(_BATCH_PROCESSED_SQL, (string_value(retrieved["id"]),)), + lambda rows: len(rows) == 1 and rows[0].get("batch_processed") is True, + seconds=60, + ) + assert processed_batch_rows[0]["batch_processed"] is True, processed_batch_rows + metadata_hits_before_recovery: Final = _metadata_hit_count(gateway, batch.scenario) + assert metadata_hits_before_recovery >= 1, "The provider metadata route was not called before recovery" + register_scenario( + batch.scenario.scenario_id, + _batch_routes(model=batch.model), + control_url=batch.scenario.control_url, + ) + details_response: Final = gateway.request("GET", f"/v1/files/{output_id}", key=batch.owner_key) + assert details_response.status_code == 200, details_response.text + details: Final = JSON_OBJECT.validate_json(details_response.content) + assert details["bytes"] == OUTPUT_BYTES, details + assert (details["filename"], details["purpose"]) == ("output.jsonl", "batch_output"), details + assert "litellm_details_fallback" not in details, details + assert _metadata_hit_count(gateway, batch.scenario) == 1 + + refreshed_files: Final = _list_files(gateway, batch.owner_key, purpose="batch_output") + refreshed_output: Final = next((file for file in refreshed_files if file.get("id") == output_id), None) + assert refreshed_output is not None, f"Refreshed batch output {output_id} is absent from the list" + assert (refreshed_output["bytes"], refreshed_output["filename"]) == (OUTPUT_BYTES, "output.jsonl"), ( + refreshed_output + ) + assert "litellm_details_fallback" not in refreshed_output, refreshed_output + + assert batch.unrelated_key is not None + unrelated_details: Final = gateway.request("GET", f"/v1/files/{output_id}", key=batch.unrelated_key) + assert unrelated_details.status_code == 403, unrelated_details.text + + +def _tls_context(directory: Path) -> ssl.SSLContext: + key: Final = ec.generate_private_key(ec.SECP256R1()) + now: Final = datetime.datetime.now(datetime.timezone.utc) + name: Final = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, BEDROCK_AUTHORITY.split(":")[0])]) + certificate: Final = ( + x509.CertificateBuilder() + .subject_name(name) + .issuer_name(name) + .public_key(key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(now - datetime.timedelta(days=1)) + .not_valid_after(now + datetime.timedelta(days=1)) + .sign(key, hashes.SHA256()) + ) + certificate_file: Final = directory / "bedrock.pem" + key_file: Final = directory / "bedrock.key" + certificate_file.write_bytes(certificate.public_bytes(serialization.Encoding.PEM)) + key_file.write_bytes( + key.private_bytes(serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()) + ) + context: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + context.load_cert_chain(certificate_file, key_file) + return context + + +def _pipe(source: socket.socket, sink: socket.socket) -> None: + with contextlib.suppress(OSError): + for chunk in iter(lambda: source.recv(65536), b""): + sink.sendall(chunk) + with contextlib.suppress(OSError): + sink.shutdown(socket.SHUT_WR) + + +@dataclass(frozen=True, slots=True) +class _ConnectProxy: + url: str + authorities: SimpleQueue[str] + + +@contextmanager +def _bedrock_tunnel(destination: Wire) -> Generator[_ConnectProxy, None, None]: + authorities: Final[SimpleQueue[str]] = SimpleQueue() + destination_port: Final = int(destination.url.rsplit(":", 1)[1]) + + class Tunnel(socketserver.StreamRequestHandler): + rbufsize = 0 + request: socket.socket + + def handle(self) -> None: + authority: Final = self.rfile.readline().decode().split()[1] + while self.rfile.readline() not in (b"\r\n", b""): + pass + authorities.put(authority) + if authority != BEDROCK_AUTHORITY: + self.wfile.write(b"HTTP/1.1 403 Forbidden\r\ncontent-length: 0\r\n\r\n") + return + self.wfile.write(b"HTTP/1.1 200 Connection established\r\n\r\n") + self.request.settimeout(10) + with socket.create_connection(("127.0.0.1", destination_port), timeout=10) as upstream: + outbound: Final = threading.Thread(target=_pipe, args=(self.request, upstream)) + outbound.start() + _pipe(upstream, self.request) + outbound.join(timeout=12) + + with socketserver.ThreadingTCPServer(("127.0.0.1", 0), Tunnel) as server: + thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05}) + thread.start() + try: + yield _ConnectProxy(f"http://127.0.0.1:{server.server_address[1]}", authorities) + finally: + server.shutdown() + thread.join(timeout=6) + + +@dataclass(frozen=True, slots=True) +class _BedrockControlPlane: + job_locations: SimpleQueue[tuple[str, str, str]] + job_id: str + job_arn: str + + def __call__(self, request: Request) -> Reply: + if request.method == "POST" and request.target == "/model-invocation-job": + body: Final = JSON_OBJECT.validate_json(request.body) + input_config: Final = object_value(object_value(body["inputDataConfig"])["s3InputDataConfig"]) + output_config: Final = object_value(object_value(body["outputDataConfig"])["s3OutputDataConfig"]) + job_name: Final = string_value(body["jobName"]) + self.job_locations.put( + (string_value(input_config["s3Uri"]), string_value(output_config["s3Uri"]), job_name) + ) + return Reply(body=json.dumps({"jobArn": self.job_arn}).encode()) + if request.method == "GET" and request.target.endswith(self.job_id): + input_uri, output_uri, retrieved_job_name = self.job_locations.get() + self.job_locations.put((input_uri, output_uri, retrieved_job_name)) + return Reply( + body=json.dumps( + { + "jobArn": self.job_arn, + "jobName": retrieved_job_name, + "modelId": BEDROCK_MODEL_ID, + "status": "Completed", + "submitTime": 1700000000, + "lastModifiedTime": 1700000001, + "endTime": 1700000002, + "inputDataConfig": {"s3InputDataConfig": {"s3Uri": input_uri}}, + "outputDataConfig": {"s3OutputDataConfig": {"s3Uri": output_uri}}, + "totalRecordCount": 1, + "successRecordCount": 1, + "errorRecordCount": 0, + } + ).encode() + ) + return Reply(status=404, body=b'{"message":"not scripted"}') + + +def _bedrock_s3_peer(request: Request) -> Reply: + if request.method == "PUT" and request.target.startswith(f"/{BEDROCK_BUCKET}/"): + return Reply(body=b"") + if request.method == "GET" and request.target.startswith(f"/{BEDROCK_BUCKET}/{BEDROCK_MANAGED_S3_OUTPUT_PREFIX}"): + if request.headers.get("range") == "bytes=0-0": + return Reply( + status=206, + body=BEDROCK_OUTPUT_CONTENT[:1], + headers={ + "Content-Range": f"bytes 0-0/{len(BEDROCK_OUTPUT_CONTENT)}", + "Last-Modified": BEDROCK_LAST_MODIFIED, + }, + ) + return Reply(status=200, body=BEDROCK_OUTPUT_CONTENT, headers={"Last-Modified": BEDROCK_LAST_MODIFIED}) + return Reply(status=404, body=b'{"message":"not scripted"}') + + +def test_bedrock_batch_output_lists_and_retrieves_details(gateway: Gateway, tmp_path: Path) -> None: + job_id: Final = f"integration-batch-listing-{uuid.uuid4().hex}" + job_arn: Final = BEDROCK_JOB_ARN_PREFIX + job_id + environment: Final = { + "AWS_CA_BUNDLE": str(tmp_path / "bedrock.pem"), + "AWS_EC2_METADATA_DISABLED": "true", + "SSL_VERIFY": "False", + } + with ( + wire_server(_bedrock_s3_peer) as s3, + wire_server(_BedrockControlPlane(SimpleQueue(), job_id, job_arn), tls=_tls_context(tmp_path)) as bedrock, + _bedrock_tunnel(bedrock) as tunnel, + owned_proxy(gateway, tmp_path, {**environment, "HTTPS_PROXY": tunnel.url}) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model=BEDROCK_MODEL, + api_key=None, + api_base=None, + aws_access_key_id="AKIAIOSFODNN7EXAMPLE", + aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + aws_region_name=BEDROCK_REGION, + s3_bucket_name=BEDROCK_BUCKET, + s3_endpoint_url=s3.url, + aws_batch_role_arn=BEDROCK_ROLE_ARN, + ) + owner_id: Final = scenario.user(user_role="internal_user") + owner_key: Final = scenario.key(user_id=owner_id, models=[model]) + input_line: Final = { + "custom_id": "req-1", + "method": "POST", + "url": "/v1/chat/completions", + "body": {"model": model, "messages": [{"role": "user", "content": "ping"}], "max_tokens": 8}, + } + uploaded: Final = candidate.request_multipart( + "/v1/files", + {"purpose": "batch", "target_model_names": model}, + {"file": ("input.jsonl", (json.dumps(input_line) + "\n").encode(), "application/jsonl")}, + key=owner_key, + ) + assert uploaded.status_code == 200, uploaded.text + created: Final = candidate.request( + "POST", + "/v1/batches", + { + "input_file_id": string_value(JSON_OBJECT.validate_json(uploaded.content)["id"]), + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + key=owner_key, + ) + assert created.status_code == 200, created.text + batch_id: Final = string_value(JSON_OBJECT.validate_json(created.content)["id"]) + retrieved: Final = candidate.request("GET", f"/v1/batches/{batch_id}", key=owner_key) + assert retrieved.status_code == 200, retrieved.text + batch_object: Final = JSON_OBJECT.validate_json(retrieved.content) + assert batch_object["status"] == "completed", retrieved.text + output_id: Final = string_value(batch_object["output_file_id"]) + output_files: Final = eventually( + lambda: _list_files(candidate, owner_key, purpose="batch_output"), + lambda values: any(file.get("id") == output_id for file in values), + seconds=30, + ) + listed_output: Final = next(file for file in output_files if file.get("id") == output_id) + details: Final = candidate.request("GET", f"/v1/files/{output_id}", key=owner_key) + assert details.status_code == 200, details.text + detail_object: Final = JSON_OBJECT.validate_json(details.content) + assert (detail_object["id"], detail_object["bytes"], detail_object["purpose"]) == ( + output_id, + len(BEDROCK_OUTPUT_CONTENT), + "batch_output", + ), detail_object + assert listed_output["id"] == output_id, listed_output + s3_requests: Final = s3.drain() + uploads: Final = tuple(request for request in s3_requests if request.method == "PUT") + assert len(uploads) == 1, f"Expected one input S3 upload, saw {[request.target for request in uploads]}" + ranged_metadata: Final = tuple( + request + for request in s3_requests + if request.method == "GET" and request.headers.get("range") == "bytes=0-0" + ) + assert ranged_metadata, "Bedrock file metadata retrieval did not issue a ranged S3 GET" + assert all( + request.headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ") for request in ranged_metadata + ), ranged_metadata + authorities: Final = tuple(tunnel.authorities.get_nowait() for _ in range(tunnel.authorities.qsize())) + assert BEDROCK_AUTHORITY in authorities, authorities + bedrock_requests: Final = bedrock.drain() + assert any(request.method == "POST" for request in bedrock_requests), bedrock_requests + assert any(request.method == "GET" for request in bedrock_requests), bedrock_requests + + +def test_repeated_batch_retrieve_does_not_refetch_saved_file_details(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + batch: Final = _create_batch(scenario, _batch_routes(model="gpt-4o-mini")) + first: Final = _retrieve_batch(gateway, batch) + output_id: Final = string_value(first["output_file_id"]) + output_before: Final = next( + (file for file in _list_files(gateway, batch.owner_key) if file.get("id") == output_id), + None, + ) + assert output_before is not None, f"Completed batch output {output_id} is absent from GET /v1/files" + hits_before: Final = _metadata_hit_count(gateway, batch.scenario) + assert hits_before >= 1, "The provider file metadata route was not called for the batch output" + second: Final = _retrieve_batch(gateway, batch) + third: Final = _retrieve_batch(gateway, batch) + assert second["status"] == third["status"] == "completed", (second, third) + output_after: Final = next( + (file for file in _list_files(gateway, batch.owner_key) if file.get("id") == output_id), + None, + ) + assert output_after == output_before, output_after + additional_hits: Final = _metadata_hit_count(gateway, batch.scenario) + assert additional_hits == 0, f"Repeated batch retrieval fetched metadata {additional_hits} more times" diff --git a/tests/integration/management/test_fallback_management_public_names.py b/tests/integration/management/test_fallback_management_public_names.py new file mode 100644 index 00000000000..b54465c1fd0 --- /dev/null +++ b/tests/integration/management/test_fallback_management_public_names.py @@ -0,0 +1,379 @@ +import uuid +from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor +from functools import partial +from itertools import chain +from pathlib import Path +from typing import Final + +import httpx +import psutil +import pytest +from pydantic import JsonValue + +from tests.integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from tests.integration._support.database import read_rows +from tests.integration._support.process import OwnedProxy, owned_proxy_process +from tests.integration._support.redis_process import owned_redis + +pytestmark: Final = pytest.mark.timeout(300) + +PROVIDER_KEY: Final = "integration-provider-key" +EVICTION_WARNING: Final = "config cache eviction of router_settings failed" +REFUSED: Final = frozenset({401, 403}) + + +@pytest.fixture +def upstream(gateway: Gateway) -> Iterator[httpx.Client]: + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as client: + client.get("/__observations").raise_for_status() + yield client + + +def _observed_requests(upstream: httpx.Client) -> list[JsonValue]: + observed: Final = upstream.get("/__observations") + observed.raise_for_status() + requests: Final = object_value(observed.json())["requests"] + assert isinstance(requests, list) + return requests + + +def _calls_to(observed: list[JsonValue], provider_model: str) -> int: + return sum(object_value(object_value(request)["body"]).get("model") == provider_model for request in observed) + + +def _fallback_body(model: str, fallback_models: list[str], fallback_type: str = "general") -> dict[str, JsonValue]: + return {"model": model, "fallback_models": list(fallback_models), "fallback_type": fallback_type} + + +def _forget_fallback(gateway: Gateway, model: str, fallback_type: str) -> None: + gateway.request("DELETE", f"/fallback/{model}", params={"fallback_type": fallback_type}) + + +def _create_fallback( + gateway: Gateway, scenario: Scenario, model: str, fallback_models: list[str], fallback_type: str = "general" +) -> httpx.Response: + scenario.cleanups.callback(_forget_fallback, gateway, model, fallback_type) + return eventually( + lambda: gateway.request("POST", "/fallback", _fallback_body(model, fallback_models, fallback_type)), + lambda response: response.status_code == 200, + seconds=30, + return_last_on_timeout=True, + ) + + +def _fallback_models(response: httpx.Response) -> list[str]: + body: Final = response.json() + models: Final = body.get("fallback_models") if isinstance(body, dict) else None + return [string_value(entry) for entry in models] if isinstance(models, list) else [] + + +def _available_models(response: httpx.Response) -> list[str]: + body: Final = response.json() + if not isinstance(body, dict): + return [] + detail: Final = body.get("detail") + source: Final = detail if isinstance(detail, dict) else body + models: Final = source.get("available_models") + return [string_value(entry) for entry in models] if isinstance(models, list) else [] + + +def _stored_fallbacks(fallback_key: str = "fallbacks") -> list[JsonValue]: + rows: Final = read_rows('SELECT param_value FROM "LiteLLM_Config" WHERE param_name = %s', ("router_settings",)) + if not rows: + return [] + settings: Final = object_value(rows[0]["param_value"]) + entries: Final = settings.get(fallback_key) or [] + assert isinstance(entries, list), entries + return entries + + +def _covered_models(entries: list[JsonValue]) -> frozenset[str]: + dict_entries: Final = (object_value(entry) for entry in entries if isinstance(entry, dict)) + return frozenset(chain.from_iterable(dict_entries)) + + +def _models_over_a_fresh_connection(gateway: Gateway, _: int) -> frozenset[str]: + with httpx.Client(base_url=gateway.client.base_url, timeout=15, trust_env=False) as client: + listed: Final = client.get("/v1/models", headers={"Authorization": f"Bearer {gateway.key}"}) + assert listed.status_code == 200, listed.text + data: Final = object_value(listed.json())["data"] + assert isinstance(data, list), listed.text + return frozenset(string_value(object_value(entry)["id"]) for entry in data) + + +def _every_worker_serves(gateway: Gateway, model: str) -> bool: + with ThreadPoolExecutor(max_workers=16) as pool: + rounds: Final = tuple( + tuple(pool.map(partial(_models_over_a_fresh_connection, gateway), range(16))) for _ in range(2) + ) + return all(model in seen for seen in chain.from_iterable(rounds)) + + +def _wait_until_served(gateway: Gateway, model: str) -> None: + eventually(lambda: _every_worker_serves(gateway, model), lambda served: served, seconds=90) + + +def _team_model(gateway: Gateway, scenario: Scenario, team: str, provider_model: str) -> tuple[str, str]: + public: Final = f"ipub-{uuid.uuid4().hex}" + created: Final = gateway.post( + "/model/new", + { + "model_name": public, + "litellm_params": { + "model": f"openai/{provider_model}", + "api_key": PROVIDER_KEY, + "api_base": f"{gateway.upstream_url}/v1", + "num_retries": 0, + }, + "model_info": {"team_id": team}, + }, + ) + info: Final = object_value(created["model_info"]) + scenario.cleanups.callback(scenario.delete_model, string_value(info["id"])) + return public, string_value(created["model_name"]) + + +def _workers(owned: OwnedProxy) -> tuple[psutil.Process, ...]: + return tuple(child for child in psutil.Process(owned.process.pid).children() if _is_worker(child)) + + +def _is_worker(child: psutil.Process) -> bool: + try: + return "spawn_main" in " ".join(child.cmdline()) and child.status() != psutil.STATUS_ZOMBIE + except psutil.Error: + return False + + +def test_post_by_team_public_name_creates_reads_back_and_peer_converges(gateway: Gateway, peer: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team() + primary: Final = scenario.model(model=f"openai/n1p-{uuid.uuid4().hex}", model_info={"team_id": team}) + fallback: Final = scenario.model(model=f"openai/n1f-{uuid.uuid4().hex}", model_info={"team_id": team}) + created: Final = _create_fallback(gateway, scenario, primary, [fallback]) + assert created.status_code == 200, created.text + here: Final = eventually( + lambda: gateway.request("GET", f"/fallback/{primary}", params={"fallback_type": "general"}), + lambda response: response.status_code == 200 and fallback in _fallback_models(response), + seconds=30, + return_last_on_timeout=True, + ) + assert here.status_code == 200 and fallback in _fallback_models(here), here.text + there: Final = eventually( + lambda: peer.request("GET", f"/fallback/{primary}", params={"fallback_type": "general"}), + lambda response: response.status_code == 200 and fallback in _fallback_models(response), + seconds=30, + return_last_on_timeout=True, + ) + assert there.status_code == 200 and fallback in _fallback_models(there), there.text + + +def test_rule_by_team_public_name_fires_on_chat_completions(gateway: Gateway, upstream: httpx.Client) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team() + primary_provider: Final = f"n2prim-{uuid.uuid4().hex}" + fallback_provider: Final = f"n2fbk-{uuid.uuid4().hex}" + primary: Final = scenario.model(model=f"openai/{primary_provider}", num_retries=0, model_info={"team_id": team}) + fallback: Final = scenario.model(model=f"openai/{fallback_provider}", model_info={"team_id": team}) + team_key: Final = scenario.key(team_id=team) + created: Final = _create_fallback(gateway, scenario, primary, [fallback]) + assert created.status_code == 200, created.text + upstream.post(f"/__scripts/{primary_provider}", json={"statuses": [500]}).raise_for_status() + upstream.get("/__observations").raise_for_status() + answered: Final = eventually( + lambda: gateway.request( + "POST", + "/v1/chat/completions", + {"model": primary, "messages": [{"role": "user", "content": f"fire {uuid.uuid4().hex}"}]}, + key=team_key, + ), + lambda response: response.status_code == 200, + seconds=60, + return_last_on_timeout=True, + ) + assert answered.status_code == 200, answered.text + observed: Final = _observed_requests(upstream) + assert _calls_to(observed, primary_provider) >= 1, observed + assert _calls_to(observed, fallback_provider) >= 1, observed + + +def test_consecutive_creates_within_the_cache_window_keep_every_rule(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + target: Final = scenario.model(model=f"openai/n5t-{uuid.uuid4().hex}") + primaries: Final = tuple(scenario.model(model=f"openai/n5{tag}-{uuid.uuid4().hex}") for tag in "abc") + for model in (target, *primaries): + _wait_until_served(gateway, model) + created: Final = tuple(_create_fallback(gateway, scenario, primary, [target]) for primary in primaries) + assert all(response.status_code == 200 for response in created), [response.text for response in created] + covered: Final = _covered_models(_stored_fallbacks()) + assert frozenset(primaries) <= covered, (primaries, covered) + + +def test_unknown_model_404_lists_team_public_and_internal_names(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team() + public, internal = _team_model(gateway, scenario, team, f"n6-{uuid.uuid4().hex}") + unknown: Final = f"n6-unknown-{uuid.uuid4().hex}" + refused: Final = eventually( + lambda: gateway.request("POST", "/fallback", _fallback_body(unknown, [public])), + lambda response: response.status_code == 404 and internal in _available_models(response), + seconds=30, + return_last_on_timeout=True, + ) + assert refused.status_code == 404, refused.text + available: Final = _available_models(refused) + assert internal in available, (internal, available) + assert public in available, (public, available) + + +def test_create_by_generated_internal_name_still_works(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team() + _, primary_internal = _team_model(gateway, scenario, team, f"c1p-{uuid.uuid4().hex}") + _, fallback_internal = _team_model(gateway, scenario, team, f"c1f-{uuid.uuid4().hex}") + created: Final = _create_fallback(gateway, scenario, primary_internal, [fallback_internal]) + assert created.status_code == 200, created.text + here: Final = eventually( + lambda: gateway.request("GET", f"/fallback/{primary_internal}", params={"fallback_type": "general"}), + lambda response: response.status_code == 200 and fallback_internal in _fallback_models(response), + seconds=30, + return_last_on_timeout=True, + ) + assert here.status_code == 200 and fallback_internal in _fallback_models(here), here.text + + +def test_delete_reads_fresh_rules_and_keeps_the_others(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + target: Final = scenario.model(model=f"openai/c2t-{uuid.uuid4().hex}") + keep: Final = scenario.model(model=f"openai/c2k-{uuid.uuid4().hex}") + drop: Final = scenario.model(model=f"openai/c2d-{uuid.uuid4().hex}") + _wait_until_served(gateway, target) + _wait_until_served(gateway, keep) + _wait_until_served(gateway, drop) + assert _create_fallback(gateway, scenario, keep, [target]).status_code == 200 + assert _create_fallback(gateway, scenario, drop, [target]).status_code == 200 + removed: Final = gateway.request("DELETE", f"/fallback/{drop}", params={"fallback_type": "general"}) + assert removed.status_code == 200, removed.text + covered: Final = _covered_models(_stored_fallbacks()) + assert keep in covered, covered + assert drop not in covered, covered + gone: Final = gateway.request("GET", f"/fallback/{drop}", params={"fallback_type": "general"}) + assert gone.status_code == 404, gone.text + + +def test_self_fallback_and_unknown_fallback_models_are_rejected(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model(model=f"openai/c3-{uuid.uuid4().hex}") + _wait_until_served(gateway, model) + itself: Final = gateway.request("POST", "/fallback", _fallback_body(model, [model])) + assert itself.status_code == 400, itself.text + unknown_target: Final = gateway.request( + "POST", "/fallback", _fallback_body(model, [f"c3-missing-{uuid.uuid4().hex}"]) + ) + assert unknown_target.status_code == 400, unknown_target.text + + +def test_malformed_requests_are_rejected_and_the_proxy_stays_up(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model(model=f"openai/c4-{uuid.uuid4().hex}") + fallback: Final = scenario.model(model=f"openai/c4f-{uuid.uuid4().hex}") + _wait_until_served(gateway, model) + _wait_until_served(gateway, fallback) + malformed: Final = ( + {"model": model, "fallback_models": [], "fallback_type": "general"}, + {"model": model, "fallback_type": "general"}, + {"model": model, "fallback_models": [fallback], "fallback_type": "nonsense"}, + {"model": 123, "fallback_models": [fallback], "fallback_type": "general"}, + ) + for body in malformed: + assert gateway.request("POST", "/fallback", body).status_code in (400, 422), body + healthy: Final = _create_fallback(gateway, scenario, model, [fallback]) + assert healthy.status_code == 200, healthy.text + listed: Final = gateway.request("GET", "/v1/models") + assert listed.status_code == 200, listed.text + + +def test_team_key_is_forbidden_on_create_and_delete(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team() + model: Final = scenario.model(model=f"openai/c5-{uuid.uuid4().hex}", model_info={"team_id": team}) + fallback: Final = scenario.model(model=f"openai/c5f-{uuid.uuid4().hex}", model_info={"team_id": team}) + team_key: Final = scenario.key(team_id=team) + creating: Final = gateway.request("POST", "/fallback", _fallback_body(model, [fallback]), key=team_key) + assert creating.status_code in REFUSED, creating.text + assert model not in _covered_models(_stored_fallbacks()) + admitted: Final = _create_fallback(gateway, scenario, model, [fallback]) + assert admitted.status_code == 200, admitted.text + deleting: Final = gateway.request( + "DELETE", f"/fallback/{model}", params={"fallback_type": "general"}, key=team_key + ) + assert deleting.status_code in REFUSED, deleting.text + assert model in _covered_models(_stored_fallbacks()) + + +def test_context_window_and_content_policy_types_persist(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + target: Final = scenario.model(model=f"openai/c6t-{uuid.uuid4().hex}") + windowed: Final = scenario.model(model=f"openai/c6w-{uuid.uuid4().hex}") + policy: Final = scenario.model(model=f"openai/c6p-{uuid.uuid4().hex}") + _wait_until_served(gateway, target) + _wait_until_served(gateway, windowed) + _wait_until_served(gateway, policy) + window_create: Final = _create_fallback(gateway, scenario, windowed, [target], "context_window") + assert window_create.status_code == 200, window_create.text + window_read: Final = eventually( + lambda: gateway.request("GET", f"/fallback/{windowed}", params={"fallback_type": "context_window"}), + lambda response: response.status_code == 200 and target in _fallback_models(response), + seconds=30, + return_last_on_timeout=True, + ) + assert window_read.status_code == 200 and target in _fallback_models(window_read), window_read.text + policy_create: Final = _create_fallback(gateway, scenario, policy, [target], "content_policy") + assert policy_create.status_code == 200, policy_create.text + covered: Final = _covered_models(_stored_fallbacks("content_policy_fallbacks")) + assert policy in covered, covered + + +def test_fallback_write_survives_a_redis_outage(gateway: Gateway, tmp_path: Path) -> None: + with owned_redis(tmp_path) as cache: + overrides: Final = {"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)} + with owned_proxy_process(gateway, tmp_path, overrides, workers=2, database_setup=()) as owned: + with owned.gateway.scenario() as scenario: + target: Final = scenario.model(model=f"openai/x1t-{uuid.uuid4().hex}") + first: Final = scenario.model(model=f"openai/x1a-{uuid.uuid4().hex}") + second: Final = scenario.model(model=f"openai/x1b-{uuid.uuid4().hex}") + third: Final = scenario.model(model=f"openai/x1c-{uuid.uuid4().hex}") + for model in (target, first, second, third): + _wait_until_served(owned.gateway, model) + assert _create_fallback(owned.gateway, scenario, first, [target]).status_code == 200 + cache.stop() + degraded: Final = tuple( + _create_fallback(owned.gateway, scenario, model, [target]) for model in (second, third) + ) + assert all(response.status_code == 200 for response in degraded), [r.text for r in degraded] + covered: Final = _covered_models(_stored_fallbacks()) + assert {first, second, third} <= covered, covered + assert EVICTION_WARNING in owned.log.read_text(), owned.log.read_text()[-3000:] + cache.start() + + +def test_fallback_write_survives_a_worker_kill(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy_process(gateway, tmp_path, {}, workers=2, database_setup=()) as owned: + with owned.gateway.scenario() as scenario: + target: Final = scenario.model(model=f"openai/x2t-{uuid.uuid4().hex}") + before: Final = scenario.model(model=f"openai/x2a-{uuid.uuid4().hex}") + after: Final = scenario.model(model=f"openai/x2b-{uuid.uuid4().hex}") + last: Final = scenario.model(model=f"openai/x2c-{uuid.uuid4().hex}") + for model in (target, before, after, last): + _wait_until_served(owned.gateway, model) + assert _create_fallback(owned.gateway, scenario, before, [target]).status_code == 200 + victim: Final = eventually(lambda: _workers(owned), lambda workers: len(workers) == 2)[0] + victim.kill() + fresh: Final = scenario.cleanups.enter_context( + httpx.Client(base_url=owned.gateway.client.base_url, timeout=15, trust_env=False) + ) + survivor: Final = Gateway(fresh, owned.gateway.key, owned.gateway.upstream_url) + degraded: Final = tuple(_create_fallback(survivor, scenario, model, [target]) for model in (after, last)) + assert all(response.status_code == 200 for response in degraded), [r.text for r in degraded] + covered: Final = _covered_models(_stored_fallbacks()) + assert {before, after, last} <= covered, covered diff --git a/tests/integration/management/test_key_route_contracts.py b/tests/integration/management/test_key_route_contracts.py new file mode 100644 index 00000000000..977f60adf62 --- /dev/null +++ b/tests/integration/management/test_key_route_contracts.py @@ -0,0 +1,152 @@ +import uuid +from concurrent.futures import ThreadPoolExecutor +from hashlib import sha256 +from typing import Final, Literal + +import pytest +from pydantic import JsonValue + +from tests.integration._support.client import Gateway, Scenario, delete_key_if_present, object_value, string_value +from tests.integration._support.database import read_rows + +KEY_ALIAS: Final = "mistral-7b" +ModelAccess = Literal["all-team-models", "single-model"] +AccessLevel = Literal["key", "team"] +ModelEndpoint = Literal["/v1/models", "/model/info"] + + +def _owned_key( + gateway: Gateway, scenario: Scenario, fields: dict[str, JsonValue], *, key: str | None = None +) -> dict[str, JsonValue]: + response: Final = gateway.request("POST", "/key/generate", fields, key=key) + assert response.status_code == 200, response.text + created: Final = object_value(response.json()) + scenario.cleanups.callback(delete_key_if_present, gateway, string_value(created["key"])) + return created + + +def _data(gateway: Gateway, path: str, key: str, params: dict[str, str] | None = None) -> list[JsonValue]: + response: Final = gateway.request("GET", path, key=key, params=params) + assert response.status_code == 200, response.text + data: Final = object_value(response.json())["data"] + assert isinstance(data, list) + return data + + +def _verification_row(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT user_id FROM "LiteLLM_VerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + +def test_concurrent_key_generation_persists_every_key(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + with ThreadPoolExecutor(max_workers=10) as pool: + keys: Final = list(pool.map(lambda _: scenario.key(models=[model]), range(10))) + assert len(set(keys)) == 10 + for key in keys: + assert len(_verification_row(key)) == 1 + + +def test_generated_key_exposes_hashed_token_and_timestamps(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + created: Final = _owned_key(gateway, scenario, {}) + key: Final = string_value(created["key"]) + assert created["token"] is not None + assert created["token"] != key + assert created["token"] == sha256(key.encode()).hexdigest() + assert created["token_id"] is not None + assert created["created_at"] is not None + assert created["updated_at"] is not None + + +def test_key_info_serves_admin_and_self_and_hides_unknown_keys(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + key: Final = scenario.key() + digest: Final = sha256(key.encode()).hexdigest() + admin: Final = gateway.request("GET", "/key/info", params={"key": key}) + assert admin.status_code == 200, admin.text + assert object_value(admin.json())["key"] == key + explicit: Final = gateway.request("GET", "/key/info", params={"key": key}, key=key) + assert explicit.status_code == 200, explicit.text + assert object_value(explicit.json())["key"] == key + implicit: Final = gateway.request("GET", "/key/info", key=key) + assert implicit.status_code == 200, implicit.text + assert object_value(implicit.json())["key"] == digest + unknown: Final = gateway.request("GET", "/key/info", params={"key": f"sk-{uuid.uuid4()}"}, key=key) + assert unknown.status_code == 404, unknown.text + + +def test_model_info_is_filtered_to_the_models_a_key_can_use(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = scenario.key(models=[model]) + admin_models: Final = _data(gateway, "/model/info", gateway.key) + user_models: Final = _data(gateway, "/model/info", key) + assert len(admin_models) > len(user_models) + assert [object_value(entry)["model_name"] for entry in user_models] == [model] + + +def test_proxy_admin_user_key_deletes_a_key_owned_by_another_user(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + admin_user: Final = scenario.user(user_role="proxy_admin") + owner: Final = scenario.user(user_role="internal_user") + admin_key: Final = string_value(_owned_key(gateway, scenario, {"user_id": admin_user})["key"]) + victim: Final = string_value(_owned_key(gateway, scenario, {"user_id": owner})["key"]) + deleted: Final = gateway.request("POST", "/key/delete", {"keys": [victim]}, key=admin_key) + assert deleted.status_code == 200, deleted.text + assert _verification_row(victim) == [] + + +@pytest.mark.parametrize("model_endpoint", ["/v1/models", "/model/info"]) +@pytest.mark.parametrize("access_level", ["key", "team"]) +@pytest.mark.parametrize("model_access", ["all-team-models", "single-model"]) +def test_key_model_list_follows_key_or_team_access( + gateway: Gateway, model_access: ModelAccess, access_level: AccessLevel, model_endpoint: ModelEndpoint +) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + granted: Final[list[JsonValue]] = [] if model_access == "all-team-models" else [model] + team: Final = scenario.team(models=granted if access_level == "team" else []) + key: Final = scenario.key( + team_id=team, + models=granted if access_level == "key" else [], + aliases={KEY_ALIAS: model}, + ) + data: Final = _data(gateway, model_endpoint, key) + if model_access == "all-team-models": + assert len(data) > 1 + if model_endpoint == "/v1/models": + assert all(isinstance(object_value(entry)["id"], str) for entry in data) + assert model in {object_value(entry)["id"] for entry in data} + else: + assert model in {object_value(entry)["model_name"] for entry in data} + elif model_endpoint == "/v1/models": + assert {object_value(entry)["id"] for entry in data} == {model, KEY_ALIAS} + else: + assert [object_value(entry)["model_name"] for entry in data] == [model] + + +def test_internal_user_cannot_reassign_its_key_to_another_user(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + first: Final = scenario.user(user_role="internal_user") + second: Final = scenario.user(user_role="internal_user") + key: Final = string_value(_owned_key(gateway, scenario, {"user_id": first})["key"]) + own_key: Final = string_value(_owned_key(gateway, scenario, {}, key=key)["key"]) + assert _verification_row(own_key) == [{"user_id": first}] + update: Final = gateway.request("POST", "/key/update", {"key": own_key, "user_id": second}, key=key) + assert update.status_code == 403, update.text + assert _verification_row(own_key) == [{"user_id": first}] + + +def test_internal_user_cannot_delete_another_users_key(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + first: Final = scenario.user(user_role="internal_user") + second: Final = scenario.user(user_role="internal_user") + victim: Final = string_value(_owned_key(gateway, scenario, {"user_id": first})["key"]) + attacker: Final = string_value(_owned_key(gateway, scenario, {"user_id": second})["key"]) + deleted: Final = gateway.request("POST", "/key/delete", {"keys": [victim]}, key=attacker) + assert deleted.status_code == 403, deleted.text + assert _verification_row(victim) == [{"user_id": first}] diff --git a/tests/integration/management/test_listing_and_health_contracts.py b/tests/integration/management/test_listing_and_health_contracts.py new file mode 100644 index 00000000000..8492dc50adb --- /dev/null +++ b/tests/integration/management/test_listing_and_health_contracts.py @@ -0,0 +1,138 @@ +import uuid +from datetime import datetime, timezone +from typing import Final + +from pydantic import JsonValue + +from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time +from tests.integration._support.client import Gateway, Scenario, delete_key_if_present, object_value, string_value +from tests.integration._support.database import read_rows + + +def _utc(text: JsonValue) -> datetime: + parsed: Final = datetime.fromisoformat(string_value(text)) + return parsed if parsed.tzinfo is not None else parsed.replace(tzinfo=timezone.utc) + + +def _model_id(model_name: str) -> str: + rows: Final = read_rows('SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_name = %s', (model_name,)) + assert len(rows) == 1, rows + return string_value(rows[0]["model_id"]) + + +def _data(gateway: Gateway, path: str, key: str, params: dict[str, str] | None = None) -> list[dict[str, JsonValue]]: + response: Final = gateway.request("GET", path, key=key, params=params) + assert response.status_code == 200, response.text + data: Final = object_value(response.json())["data"] + assert isinstance(data, list) + return [object_value(entry) for entry in data] + + +def _wildcard_model(gateway: Gateway, scenario: Scenario, prefix: str) -> None: + created: Final = gateway.post( + "/model/new", + { + "model_name": f"{prefix}/*", + "litellm_params": { + "model": "anthropic/*", + "api_key": "integration-provider-key", + "api_base": gateway.upstream_url, + }, + }, + ) + scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"])) + + +def test_budget_duration_schedules_reset_at_the_next_standardized_boundary(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + budget_id: Final = scenario.budget(max_budget=10.0, budget_duration="1d") + rows: Final = read_rows( + 'SELECT created_at::text AS created_at, budget_reset_at::text AS reset_at FROM "LiteLLM_BudgetTable" ' + "WHERE budget_id = %s", + (budget_id,), + ) + assert len(rows) == 1, rows + assert rows[0]["reset_at"] is not None, rows + expected: Final = get_next_standardized_reset_time("1d", _utc(rows[0]["created_at"]), "UTC") + assert abs((_utc(rows[0]["reset_at"]) - expected).total_seconds()) <= 3, (rows, expected) + + +def test_admin_health_counts_every_deployment_and_reports_a_live_one_healthy(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + model_id: Final = _model_id(model) + report: Final = gateway.get("/health", {"model": model}) + assert (report["healthy_count"], report["unhealthy_count"]) == (1, 0), report + healthy: Final = report["healthy_endpoints"] + assert isinstance(healthy, list) and len(healthy) == 1, report + assert object_value(healthy[0]).get("model_id") == model_id, report + + +def test_routes_listing_is_served_without_credentials(gateway: Gateway) -> None: + response: Final = gateway.client.get("/routes") + assert response.status_code == 200, response.text + routes: Final = object_value(response.json())["routes"] + assert isinstance(routes, list) + assert {"/routes", "/key/generate"} <= {object_value(route)["path"] for route in routes} + + +def test_unrestricted_key_lists_models_and_none_when_only_access_groups_are_requested(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + grouped: Final = scenario.model(model_info={"access_groups": [f"integration-{uuid.uuid4().hex}"]}) + plain: Final = scenario.model() + key: Final = scenario.key() + listed: Final = {entry["id"] for entry in _data(gateway, "/models", key)} + assert {grouped, plain} <= listed + assert _data(gateway, "/models", key, {"only_model_access_groups": "True"}) == [] + + +def test_model_info_by_id_matches_the_entry_in_the_keys_listing(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + scenario.model() + model_id: Final = _model_id(model) + key: Final = scenario.key(models=[model]) + listing: Final = _data(gateway, "/model/info", key) + assert {entry["model_name"] for entry in listing} == {model} + listed: Final = [entry for entry in listing if object_value(entry["model_info"])["id"] == model_id] + assert len(listed) == 1 + by_id: Final = _data(gateway, "/model/info", key, {"litellm_model_id": model_id}) + assert by_id == listed + admin_by_id: Final = _data(gateway, "/model/info", gateway.key, {"litellm_model_id": model_id}) + assert [object_value(entry["model_info"])["id"] for entry in admin_by_id] == [model_id] + + +def test_model_group_info_for_a_personal_user_key_lists_only_the_users_model(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + scenario.model() + created: Final = gateway.post("/user/new", {"user_id": f"integration-{uuid.uuid4().hex}", "models": [model]}) + scenario.cleanups.callback(scenario.delete_user, string_value(created["user_id"])) + key: Final = string_value(created["key"]) + scenario.cleanups.callback(delete_key_if_present, gateway, key) + groups: Final = [entry["model_group"] for entry in _data(gateway, "/model_group/info", key)] + assert groups == [model] + + +def test_model_group_info_expands_a_wildcard_deployment_into_concrete_groups(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + prefix: Final = f"integration{uuid.uuid4().hex}" + _wildcard_model(gateway, scenario, prefix) + groups: Final = { + string_value(entry["model_group"]) for entry in _data(gateway, "/model_group/info", gateway.key) + } + assert f"{prefix}/*" not in groups + assert any(group.startswith(f"{prefix}/claude") for group in groups), sorted(groups)[:20] + + +def test_azure_deployment_route_denies_a_model_outside_the_keys_list(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + allowed: Final = scenario.model() + other: Final = scenario.model() + key: Final = scenario.key(models=[allowed]) + body: Final[dict[str, JsonValue]] = {"messages": [{"role": "user", "content": "integration control"}]} + served: Final = gateway.request("POST", f"/openai/deployments/{allowed}/chat/completions", body, key=key) + assert served.status_code == 200, served.text + denied: Final = gateway.request("POST", f"/openai/deployments/{other}/chat/completions", body, key=key) + assert denied.status_code == 403, denied.text + assert "is not available for this API key" in denied.text diff --git a/tests/integration/management/test_team_member_routes.py b/tests/integration/management/test_team_member_routes.py new file mode 100644 index 00000000000..3dc126b3e98 --- /dev/null +++ b/tests/integration/management/test_team_member_routes.py @@ -0,0 +1,205 @@ +import uuid +from concurrent.futures import ThreadPoolExecutor +from typing import Final, Literal + +import pytest +from pydantic import JsonValue + +from tests.integration._support.client import Gateway, Scenario, delete_key_if_present, object_value, string_value +from tests.integration._support.database import read_rows + +MemberDimension = Literal["user_id", "user_email"] +UNCHANGED_BY_ALIAS_UPDATE_SKIP: Final = frozenset( + { + "team_alias", + "members_with_roles", + "created_at", + "updated_at", + "model_spend", + "model_max_budget", + "model_id", + "litellm_organization_table", + "object_permission_id", + "object_permission", + "litellm_model_table", + "policies", + "allow_team_guardrail_config", + "projects", + } +) + + +def _team_info(gateway: Gateway, team_id: str) -> dict[str, JsonValue]: + return object_value(gateway.get("/team/info", {"team_id": team_id})["team_info"]) + + +def _member_ids(gateway: Gateway, team_id: str) -> list[str | None]: + members: Final = _team_info(gateway, team_id)["members_with_roles"] + assert isinstance(members, list) + return [ + member_id if isinstance(member_id := object_value(member).get("user_id"), str) else None for member in members + ] + + +def _user_team_ids(gateway: Gateway, user_id: str) -> list[str]: + teams: Final = gateway.get("/user/info", {"user_id": user_id})["teams"] + assert isinstance(teams, list) + return [string_value(object_value(team)["team_id"]) for team in teams] + + +def _delete_users_if_present(gateway: Gateway, user_ids: list[str]) -> None: + present: Final = [ + user_id + for user_id in user_ids + if read_rows('SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id = %s', (user_id,)) + ] + if present: + response: Final = gateway.request("POST", "/user/delete", {"user_ids": list[JsonValue](present)}) + assert response.status_code == 200, response.text + + +def _member_delete(gateway: Gateway, team_id: str, user_id: str) -> int: + return gateway.request("POST", "/team/member_delete", {"team_id": team_id, "user_id": user_id}).status_code + + +def _owned_email_member(gateway: Gateway, scenario: Scenario) -> str: + user_id: Final = f"integration_{uuid.uuid4().hex}@example.com" + scenario.cleanups.callback(_delete_users_if_present, gateway, [user_id]) + return user_id + + +def test_concurrent_team_creation_records_the_named_member(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + user: Final = scenario.user() + with ThreadPoolExecutor(max_workers=10) as pool: + teams: Final = list( + pool.map(lambda _: scenario.team(members_with_roles=[{"role": "user", "user_id": user}]), range(10)) + ) + assert len(set(teams)) == 10 + for team in teams: + assert user in _member_ids(gateway, team) + + +def test_team_info_serves_admin_and_team_keys_and_denies_other_keys(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team() + assert _team_info(gateway, team)["team_id"] == team + team_key: Final = scenario.key(team_id=team) + as_team: Final = gateway.request("GET", "/team/info", params={"team_id": team}, key=team_key) + assert as_team.status_code == 200, as_team.text + assert object_value(object_value(as_team.json())["team_info"])["team_id"] == team + outsider: Final = scenario.key() + denied: Final = gateway.request("GET", "/team/info", params={"team_id": team}, key=outsider) + assert denied.status_code in {401, 403}, denied.text + + +def test_team_alias_update_keeps_every_other_team_field(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + admin: Final = scenario.user() + created: Final = gateway.post( + "/team/new", + { + "team_alias": f"integration-{uuid.uuid4().hex}", + "members_with_roles": [{"role": "admin", "user_id": admin}], + }, + ) + team: Final = string_value(created["team_id"]) + scenario.cleanups.callback(scenario.delete_team, team) + initial_size: Final = len(_member_ids(gateway, team)) + new_members: Final = [_owned_email_member(gateway, scenario) for _ in range(10)] + gateway.post( + "/team/member_add", + {"team_id": team, "member": [{"role": "user", "user_id": member} for member in new_members]}, + ) + members: Final = _member_ids(gateway, team) + assert len(members) == initial_size + 10 + assert set(new_members) <= set(members) + new_alias: Final = f"integration-{uuid.uuid4().hex}" + updated: Final = object_value(gateway.post("/team/update", {"team_id": team, "team_alias": new_alias})["data"]) + assert updated["team_alias"] == new_alias + updated_members: Final = updated["members_with_roles"] + assert isinstance(updated_members, list) + assert len(updated_members) == len(members) + compared: Final = (set(created) | set(updated)) - UNCHANGED_BY_ALIAS_UPDATE_SKIP + for field in compared: + assert updated.get(field) == created.get(field), field + assert compared + + +def test_member_added_by_email_sees_the_team_in_user_info(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + email: Final = f"integration-{uuid.uuid4().hex}@example.com" + user: Final = scenario.user(user_email=email) + team: Final = scenario.team() + gateway.post("/team/member_add", {"team_id": team, "member": {"role": "user", "user_email": email}}) + assert team in _user_team_ids(gateway, user) + + +def test_team_delete_detaches_members_and_hides_the_team(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + first: Final = scenario.user() + second: Final = scenario.user() + created: Final = gateway.post( + "/team/new", + { + "team_alias": f"integration-{uuid.uuid4().hex}", + "members_with_roles": [{"role": "admin", "user_id": first}, {"role": "user", "user_id": second}], + }, + ) + team: Final = string_value(created["team_id"]) + scenario.cleanups.callback(gateway.request, "POST", "/team/delete", {"team_ids": [team]}) + team_key: Final = string_value(gateway.post("/key/generate", {"team_id": team, "user_id": second})["key"]) + scenario.cleanups.callback(delete_key_if_present, gateway, team_key) + assert _user_team_ids(gateway, second) == [team] + gateway.post("/team/delete", {"team_ids": [team]}) + assert _user_team_ids(gateway, second) == [] + missing: Final = gateway.request("GET", "/team/info", params={"team_id": team}) + assert missing.status_code == 404, missing.text + + +@pytest.mark.parametrize("dimension", ["user_id", "user_email"]) +def test_member_delete_removes_the_member_by_id_or_email(gateway: Gateway, dimension: MemberDimension) -> None: + with gateway.scenario() as scenario: + email: Final = f"integration-{uuid.uuid4().hex}@example.com" + user: Final = scenario.user(user_email=email) + selector: Final[dict[str, JsonValue]] = {"user_id": user} if dimension == "user_id" else {"user_email": email} + team: Final = scenario.team(members_with_roles=[{"role": "user", **selector}]) + assert user in _member_ids(gateway, team) + deleted: Final = gateway.request("POST", "/team/member_delete", {"team_id": team, **selector}) + assert deleted.status_code == 200, deleted.text + assert user not in _member_ids(gateway, team) + assert team not in _user_team_ids(gateway, user) + + +def test_member_add_with_email_shaped_user_id_grows_team_by_one(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team() + initial: Final = _member_ids(gateway, team) + member: Final = _owned_email_member(gateway, scenario) + gateway.post("/team/member_add", {"team_id": team, "member": {"role": "user", "user_id": member}}) + after: Final = _member_ids(gateway, team) + assert member in after + assert len(after) == len(initial) + 1 + + +def test_member_delete_removes_once_and_rejects_repeats(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team() + member: Final = f"integration-{uuid.uuid4().hex}" + scenario.cleanups.callback(_delete_users_if_present, gateway, [member]) + gateway.post("/team/member_add", {"team_id": team, "member": {"role": "user", "user_id": member}}) + before: Final = _member_ids(gateway, team) + assert member in before + assert _member_delete(gateway, team, member) == 200 + assert [_member_delete(gateway, team, member) for _ in range(4)] == [400, 400, 400, 400] + after: Final = _member_ids(gateway, team) + assert member not in after + assert len(after) == len(before) - 1 + + +def test_member_delete_of_a_non_member_is_a_bad_request(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team() + outsider: Final = f"integration-{uuid.uuid4().hex}" + assert outsider not in _member_ids(gateway, team) + assert _member_delete(gateway, team, outsider) == 400 diff --git a/tests/integration/management/test_user_routes.py b/tests/integration/management/test_user_routes.py new file mode 100644 index 00000000000..feb9ff2c3bf --- /dev/null +++ b/tests/integration/management/test_user_routes.py @@ -0,0 +1,34 @@ +import uuid +from concurrent.futures import ThreadPoolExecutor +from typing import Final + +from tests.integration._support.client import Gateway, object_value +from tests.integration._support.database import read_rows + + +def test_concurrent_user_creation_persists_models_and_aliases(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + alias: Final = f"alias-{uuid.uuid4().hex}" + with ThreadPoolExecutor(max_workers=10) as pool: + users: Final = list(pool.map(lambda _: scenario.user(models=[model], aliases={alias: model}), range(10))) + assert len(set(users)) == 10 + rows: Final = [ + read_rows('SELECT models FROM "LiteLLM_UserTable" WHERE user_id = %s', (user,)) for user in users + ] + assert rows == [[{"models": [model]}]] * len(users) + + +def test_user_info_serves_admin_and_self_and_denies_other_users(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + user: Final = scenario.user(user_role="internal_user") + own_key: Final = scenario.key(user_id=user) + other: Final = scenario.user(user_role="internal_user") + other_key: Final = scenario.key(user_id=other) + admin: Final = object_value(gateway.get("/user/info", {"user_id": user})["user_info"]) + assert admin["user_id"] == user + own: Final = gateway.request("GET", "/user/info", params={"user_id": user}, key=own_key) + assert own.status_code == 200, own.text + assert object_value(object_value(own.json())["user_info"])["user_id"] == user + denied: Final = gateway.request("GET", "/user/info", params={"user_id": user}, key=other_key) + assert denied.status_code == 403, denied.text diff --git a/tests/integration/mcp/test_mcp_akto_logging_only.py b/tests/integration/mcp/test_mcp_akto_logging_only.py new file mode 100644 index 00000000000..5b3fd84aa83 --- /dev/null +++ b/tests/integration/mcp/test_mcp_akto_logging_only.py @@ -0,0 +1,133 @@ +"""Akto guardrail in `mode: logging_only` on MCP tool calls through a real proxy with two workers. + +The scripted MCP server and the Akto `/api/http-proxy` service are the only doubles. +""" + +import json +import uuid +from collections.abc import Iterator +from dataclasses import dataclass +from hashlib import sha256 +from pathlib import Path +from typing import Final +from urllib.parse import parse_qs, urlsplit + +import pytest +import yaml +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment +from integration._support.database import read_rows +from integration._support.mcp import EntryPoint, McpCaller, McpPeer, echo_tool, register_mcp, scripted_peer, tool_calls +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server + +AKTO_KEY: Final = "synthetic-akto-mcp-log-key" +BLOCK_MARK: Final = "SYNTHETIC-AKTO-BLOCK" +RESPONSE_CHECK: Final = {"akto_connector": "litellm", "response_guardrails": "true", "ingest_data": "true"} +MCP_SPEND_ROWS: Final = 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE api_key = %s AND call_type = %s' + + +@dataclass(frozen=True, slots=True) +class AktoCall: + path: str + flags: dict[str, str] + authorization: str + payload: dict[str, object] + + +def _akto_call(request: Request) -> AktoCall: + target: Final = urlsplit(request.target) + return AktoCall( + path=target.path, + flags={name: values[0] for name, values in parse_qs(target.query).items()}, + authorization=request.headers.get("authorization", ""), + payload=json.loads(request.body), + ) + + +def _akto_verdict(request: Request) -> Reply: + verdict: Final = ( + {"Allowed": False, "Behaviour": "block", "Reason": "Synthetic Akto policy block"} + if BLOCK_MARK.encode() in request.body + else {"Allowed": True} + ) + return Reply(body=json.dumps({"data": {"guardrailsResult": verdict}}).encode()) + + +def _config(akto_url: str, root: Path) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": "akto-mcp-log", + "litellm_params": { + "guardrail": "akto", + "mode": "logging_only", + "default_on": True, + "akto_base_url": akto_url, + "akto_api_key": AKTO_KEY, + }, + } + ] + path: Final = root / "akto-mcp-logging-only.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@dataclass(frozen=True, slots=True) +class Rig: + proxy: Gateway + akto: Wire + peer: McpPeer + scenario: Scenario + alias: str + server_id: str + + def key(self) -> str: + return self.scenario.key(object_permission={"mcp_servers": [self.server_id]}) + + def settled_akto_calls(self, key: str, marker: str) -> tuple[AktoCall, ...]: + eventually( + lambda: read_rows(MCP_SPEND_ROWS, (sha256(key.encode()).hexdigest(), "call_mcp_tool")), + lambda rows: len(rows) >= 1, + seconds=70, + ) + eventually(lambda: self.akto.received.qsize(), lambda count: count >= 1, seconds=30) + calls: Final = tuple(_akto_call(request) for request in self.akto.drain() if marker.encode() in request.body) + return tuple(call for call in calls if call.authorization == AKTO_KEY) + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]: + root: Final = tmp_path_factory.mktemp("akto-mcp-logging-only") + alias: Final = "aktolog" + uuid.uuid4().hex[:8] + with ( + gateway_from_environment() as gateway, + wire_server(_akto_verdict) as akto, + scripted_peer(echo_tool("echo")) as peer, + owned_proxy_process(gateway, root, {}, config=_config(akto.url, root), workers=2) as owned, + owned.gateway.scenario() as scenario, + ): + identity: Final = register_mcp(scenario, peer, alias) + peer.drain() + akto.drain() + yield Rig(owned.gateway, akto, peer, scenario, alias, identity) + + +@pytest.mark.parametrize("entry", ("mcp", "rest")) +def test_logging_only_akto_block_verdict_never_blocks_an_mcp_tool_call_and_still_checks_it( + rig: Rig, entry: EntryPoint +) -> None: + marker: Final = "mark-" + uuid.uuid4().hex + arguments: Final = {"note": f"{BLOCK_MARK} {marker}"} + key: Final = rig.key() + + outcome: Final = McpCaller(rig.proxy, key, entry, rig.alias).call( + f"{rig.alias}-echo", arguments, server_id=rig.server_id + ) + + assert (outcome.error, outcome.text) == (None, json.dumps(arguments, sort_keys=True)), outcome.raw + upstream: Final = [call["body"]["params"]["arguments"] for call in tool_calls(rig.peer.drain())] + assert upstream == [arguments], upstream + calls: Final = rig.settled_akto_calls(key, marker) + assert all(call.path == "/api/http-proxy" for call in calls), calls + checked: Final = [call for call in calls if call.flags == RESPONSE_CHECK] + assert [marker in str(call.payload.get("responsePayload")) for call in checked] == [True], f"{entry}: {calls}" diff --git a/tests/integration/messages_endpoint/providers/anthropic/test_passthrough_migration_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_passthrough_migration_wire.py new file mode 100644 index 00000000000..834e10e32b8 --- /dev/null +++ b/tests/integration/messages_endpoint/providers/anthropic/test_passthrough_migration_wire.py @@ -0,0 +1,584 @@ +import json +import uuid +from collections.abc import Callable, Generator, Mapping +from contextlib import contextmanager +from hashlib import sha256 +from pathlib import Path +from typing import Final + +from integration._support.client import Gateway, JsonValue, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, Wire, wire_server + +_MODEL: Final = "claude-sonnet-4-5-20250929" +_KEY: Final = "synthetic-anthropic-key" + +_PROXY_CONFIG: Final = ( + "model_list: []\n" + "general_settings:\n" + " master_key: os.environ/LITELLM_MASTER_KEY\n" + " database_url: os.environ/DATABASE_URL\n" + " store_model_in_db: true\n" + " disable_spend_logs: false\n" + " proxy_batch_write_at: 1\n" +) + + +def _message(message_id: str, model: str = _MODEL) -> dict[str, JsonValue]: + return { + "id": message_id, + "type": "message", + "role": "assistant", + "model": model, + "content": [{"type": "text", "text": "hello test"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 11, "output_tokens": 7}, + } + + +def _thinking_message(message_id: str) -> dict[str, JsonValue]: + return { + **_message(message_id, "claude-haiku-4-5-20251001"), + "content": [ + {"type": "thinking", "thinking": "pondering the joke", "signature": "sig1"}, + {"type": "text", "text": "hello thinking"}, + ], + "usage": {"input_tokens": 11, "output_tokens": 30}, + } + + +def _sse(event: str, payload: Mapping[str, JsonValue]) -> bytes: + return f"event: {event}\ndata: {json.dumps(payload)}\n\n".encode() + + +def _stream_chunks(message_id: str) -> tuple[bytes, ...]: + return ( + _sse( + "message_start", + { + "type": "message_start", + "message": { + "id": message_id, + "type": "message", + "role": "assistant", + "model": _MODEL, + "content": [], + "usage": {"input_tokens": 11, "output_tokens": 1}, + }, + }, + ), + _sse( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + _sse( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hello stream"}}, + ), + _sse("content_block_stop", {"type": "content_block_stop", "index": 0}), + _sse( + "message_delta", + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 7}}, + ), + _sse("message_stop", {"type": "message_stop"}), + ) + + +def _bad_request_reply() -> Reply: + return Reply( + status=400, + body=json.dumps( + {"type": "error", "error": {"type": "invalid_request_error", "message": "messages must be objects"}} + ).encode(), + ) + + +def _messages_are_objects(body: Mapping[str, JsonValue]) -> bool: + messages: Final = body.get("messages") + return isinstance(messages, list) and all(isinstance(message, dict) for message in messages) + + +_SPEND_COLUMNS: Final = ( + "SELECT request_id, status, call_type, prompt_tokens, completion_tokens, total_tokens, spend, request_tags, " + 'end_user, api_base, custom_llm_provider, model, cache_hit, ("startTime" <= "endTime") AS times_ordered ' + 'FROM "LiteLLM_SpendLogs" ' +) + + +def _spend_row(request_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows(_SPEND_COLUMNS + "WHERE request_id=%s", (request_id,)), + lambda values: len(values) == 1, + seconds=90, + ) + return rows[0] + + +def _key_spend_row(key: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + _SPEND_COLUMNS + "WHERE api_key=%s AND call_type=%s", + (sha256(key.encode()).hexdigest(), "pass_through_endpoint"), + ), + lambda values: len(values) == 1, + seconds=90, + ) + return rows[0] + + +_MODEL_LIST_PROBE: Final = ("GET", "/v1/models") + + +def _is_model_list_probe(request: Request) -> bool: + return (request.method, request.target) == _MODEL_LIST_PROBE + + +@contextmanager +def _upstream(respond: Callable[[Request], Reply]) -> Generator[Wire]: + with wire_server( + lambda request: ( + Reply(body=b'{"object":"list","data":[]}') if _is_model_list_probe(request) else respond(request) + ) + ) as wire: + yield wire + + +def _provider_calls(wire: Wire) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if not _is_model_list_probe(request)) + + +def _tags(row: Mapping[str, JsonValue]) -> list[JsonValue]: + raw: Final = row["request_tags"] + tags: Final = json.loads(raw) if isinstance(raw, str) else raw + assert isinstance(tags, list), row + return [tag for tag in tags if not (isinstance(tag, str) and tag.startswith("User-Agent: "))] + + +def _assert_usage_row(row: Mapping[str, JsonValue], call_type: str, tags: list[str]) -> None: + assert row["status"] == "success", row + assert row["call_type"] == call_type, row + assert row["prompt_tokens"] == 11, row + assert row["completion_tokens"] == 7, row + assert row["total_tokens"] == 18, row + spend: Final = row["spend"] + assert isinstance(spend, (int, float)) and spend > 0, row + assert _tags(row) == tags, row + assert row["custom_llm_provider"] == "anthropic", row + assert str(row["cache_hit"]).lower() != "true", row + assert row["times_ordered"] is True, row + + +def _stream_text(gateway: Gateway, path: str, body: Mapping[str, JsonValue], key: str | None = None) -> str: + with gateway.client.stream( + "POST", + f"{gateway.client.base_url}{path}", + json=body, + headers={"Authorization": f"Bearer {key or gateway.key}"}, + ) as stream: + assert stream.status_code == 200, stream.read() + return "".join(stream.iter_text()) + + +def _owned_config(tmp_path: Path, text: str) -> Path: + config: Final = tmp_path / "proxy_config.yaml" + config.write_text(text) + return config + + +def test_passthrough_basic_completion_spend_row_v1_messages(gateway: Gateway) -> None: + marker: Final = "pt-basic-" + uuid.uuid4().hex + prompt: Final = f"say hello {marker}" + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages", request.target + assert request.headers["x-api-key"] == _KEY + body: Final = json.loads(request.body) + assert body["model"] == _MODEL + assert body["messages"] == [{"role": "user", "content": prompt}] + assert "litellm_metadata" not in body + return Reply(body=json.dumps(_message(f"msg_{marker}")).encode()) + + with _upstream(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_KEY) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 100, + "messages": [{"role": "user", "content": prompt}], + "litellm_metadata": {"tags": [f"{marker}-1", f"{marker}-2"], "user": f"end-user-{marker}"}, + }, + ) + assert response.status_code == 200, response.text + assert response.json()["id"] == f"msg_{marker}" + assert len(_provider_calls(wire)) == 1 + row: Final = _spend_row(f"msg_{marker}") + _assert_usage_row(row, "anthropic_messages", [f"{marker}-1", f"{marker}-2"]) + assert row["end_user"] == f"end-user-{marker}", row + + +def test_passthrough_streaming_spend_row_v1_messages(gateway: Gateway) -> None: + marker: Final = "pt-stream-" + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + assert request.target == "/v1/messages", request.target + body: Final = json.loads(request.body) + assert body["model"] == _MODEL + assert body["stream"] is True + return Reply(content_type="text/event-stream", chunks=_stream_chunks(f"msg_{marker}")) + + with _upstream(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_KEY) + text: Final = _stream_text( + gateway, + "/v1/messages", + { + "model": model, + "max_tokens": 100, + "stream": True, + "messages": [{"role": "user", "content": f"say hello {marker}"}], + "litellm_metadata": {"tags": [f"{marker}-1", f"{marker}-2"], "user": f"end-user-{marker}"}, + }, + ) + assert "hello stream" in text + row: Final = _spend_row(f"msg_{marker}") + _assert_usage_row(row, "anthropic_messages", [f"{marker}-1", f"{marker}-2"]) + assert row["end_user"] == f"end-user-{marker}", row + + +def test_passthrough_wildcard_model_strips_provider_prefix(gateway: Gateway) -> None: + marker: Final = "pt-wildcard-" + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + body: Final = json.loads(request.body) + assert body["model"] == "claude-haiku-4-5-20251001" + return Reply(body=json.dumps(_message(f"msg_{marker}", "claude-haiku-4-5-20251001")).encode()) + + with _upstream(respond) as wire, gateway.scenario() as scenario: + created: Final = gateway.post( + "/model/new", + { + "model_name": "anthropic/*", + "litellm_params": {"model": "anthropic/*", "api_base": wire.url, "api_key": _KEY}, + }, + ) + model_info: Final = created["model_info"] + assert isinstance(model_info, dict), created + identity: Final = model_info["id"] + assert isinstance(identity, str), created + scenario.cleanups.callback(scenario.delete_model, identity) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": "anthropic/claude-haiku-4-5-20251001", + "max_tokens": 100, + "messages": [{"role": "user", "content": f"hello wildcard {marker}"}], + }, + ) + assert response.status_code == 200, response.text + assert response.json()["content"][0]["text"] == "hello test" + assert len(_provider_calls(wire)) == 1 + + +def test_passthrough_thinking_block_round_trips_v1_messages(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + body: Final = json.loads(request.body) + assert body["model"] == "claude-haiku-4-5-20251001" + assert body["thinking"] == {"type": "enabled", "budget_tokens": 16000} + assert body["max_tokens"] == 20000 + return Reply(body=json.dumps(_thinking_message("msg_" + uuid.uuid4().hex)).encode()) + + with _upstream(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model="anthropic/claude-haiku-4-5-20251001", api_base=wire.url, api_key=_KEY) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 20000, + "thinking": {"type": "enabled", "budget_tokens": 16000}, + "messages": [{"role": "user", "content": "Just pinging with thinking enabled"}], + }, + ) + assert response.status_code == 200, response.text + content: Final = response.json()["content"] + assert content[0]["type"] == "thinking" + assert content[0]["thinking"] == "pondering the joke" + + +def test_passthrough_bad_request_returns_400_v1_messages(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + body: Final = json.loads(request.body) + assert not _messages_are_objects(body), body + return _bad_request_reply() + + with _upstream(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_KEY) + responses: Final = tuple( + gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 10, "stream": stream, "messages": ["hi"]}, + ) + for stream in (False, True) + ) + assert [response.status_code for response in responses] == [400, 400], [r.text for r in responses] + + +def test_native_anthropic_route_completion_stream_thinking_and_bad_request(gateway: Gateway, tmp_path: Path) -> None: + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages", request.target + assert request.headers["x-api-key"] == _KEY + body: Final = json.loads(request.body) + if not _messages_are_objects(body): + return _bad_request_reply() + if body.get("stream") is True: + return Reply(content_type="text/event-stream", chunks=_stream_chunks("msg_" + uuid.uuid4().hex)) + if body.get("thinking") is not None: + return Reply(body=json.dumps(_thinking_message("msg_" + uuid.uuid4().hex)).encode()) + return Reply(body=json.dumps(_message("msg_" + uuid.uuid4().hex)).encode()) + + with ( + _upstream(respond) as wire, + owned_proxy( + gateway, + tmp_path, + {"ANTHROPIC_API_BASE": wire.url, "ANTHROPIC_API_KEY": _KEY}, + config=_owned_config(tmp_path, _PROXY_CONFIG), + ) as candidate, + ): + completion: Final = candidate.request( + "POST", + "/anthropic/v1/messages", + {"model": _MODEL, "max_tokens": 100, "messages": [{"role": "user", "content": "say hello native"}]}, + ) + assert completion.status_code == 200, completion.text + assert completion.json()["content"][0]["text"] == "hello test" + thinking: Final = candidate.request( + "POST", + "/anthropic/v1/messages", + { + "model": "claude-haiku-4-5-20251001", + "max_tokens": 20000, + "thinking": {"type": "enabled", "budget_tokens": 16000}, + "messages": [{"role": "user", "content": "ping"}], + }, + ) + assert thinking.status_code == 200, thinking.text + assert thinking.json()["content"][0]["type"] == "thinking" + assert thinking.json()["content"][0]["thinking"] == "pondering the joke" + bad: Final = tuple( + candidate.request( + "POST", + "/anthropic/v1/messages", + {"model": _MODEL, "max_tokens": 10, "stream": stream, "messages": ["hi"]}, + ) + for stream in (False, True) + ) + assert [response.status_code for response in bad] == [400, 400], [r.text for r in bad] + text: Final = _stream_text( + candidate, + "/anthropic/v1/messages", + { + "model": _MODEL, + "max_tokens": 100, + "stream": True, + "messages": [{"role": "user", "content": "hello native stream"}], + }, + ) + assert "hello stream" in text + + +def test_native_passthrough_spend_rows_record_usage_tags_and_spend(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "pt-native-" + uuid.uuid4().hex + completion_id: Final = f"msg_{marker}_completion" + stream_id: Final = f"msg_{marker}_stream" + + def respond(request: Request) -> Reply: + body: Final = json.loads(request.body) + assert "litellm_metadata" not in body + if body.get("stream") is True: + return Reply(content_type="text/event-stream", chunks=_stream_chunks(stream_id)) + return Reply(body=json.dumps(_message(completion_id)).encode()) + + with ( + _upstream(respond) as wire, + owned_proxy( + gateway, + tmp_path, + {"ANTHROPIC_API_BASE": wire.url, "ANTHROPIC_API_KEY": _KEY}, + config=_owned_config(tmp_path, _PROXY_CONFIG), + ) as candidate, + candidate.scenario() as scenario, + ): + completion_key: Final = scenario.key() + stream_key: Final = scenario.key() + response: Final = candidate.request( + "POST", + "/anthropic/v1/messages", + { + "model": _MODEL, + "max_tokens": 10, + "messages": [{"role": "user", "content": "Say 'hello test' and nothing else"}], + "litellm_metadata": {"tags": [f"{marker}-1", f"{marker}-2"]}, + }, + key=completion_key, + ) + assert response.status_code == 200, response.text + assert response.json()["id"] == completion_id + text: Final = _stream_text( + candidate, + "/anthropic/v1/messages", + { + "model": _MODEL, + "max_tokens": 10, + "stream": True, + "messages": [{"role": "user", "content": "Say 'hello stream test' and nothing else"}], + "litellm_metadata": {"tags": [f"{marker}-s1", f"{marker}-s2"], "user": f"end-user-{marker}"}, + }, + key=stream_key, + ) + assert "hello stream" in text + completion_row: Final = _key_spend_row(completion_key) + stream_row: Final = _key_spend_row(stream_key) + assert completion_row["request_id"] == completion_id, completion_row + assert stream_row["request_id"] == stream_id, stream_row + _assert_usage_row(completion_row, "pass_through_endpoint", [f"{marker}-1", f"{marker}-2"]) + assert completion_row["api_base"] == f"{wire.url}/v1/messages", completion_row + assert "claude" in str(completion_row["model"]), completion_row + _assert_usage_row(stream_row, "pass_through_endpoint", [f"{marker}-s1", f"{marker}-s2"]) + assert stream_row["end_user"] == f"end-user-{marker}", stream_row + + +def _openai_responses_stream() -> tuple[bytes, ...]: + response: Final[dict[str, JsonValue]] = { + "id": "resp_pt1", + "object": "response", + "created_at": 1700000000, + "model": "gpt-4o-mini", + "status": "completed", + "output": [ + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "hi from openai"}], + } + ], + "usage": {"input_tokens": 12, "output_tokens": 8, "total_tokens": 20}, + } + return ( + _sse( + "response.created", + {"type": "response.created", "response": {**response, "status": "in_progress", "output": []}}, + ), + _sse( + "response.output_text.delta", + { + "type": "response.output_text.delta", + "item_id": "msg_pto", + "output_index": 0, + "content_index": 0, + "delta": "hi from openai", + }, + ), + _sse("response.completed", {"type": "response.completed", "response": response}), + ) + + +def _openai_chat_stream() -> tuple[bytes, ...]: + chunk: Final[dict[str, JsonValue]] = { + "id": "chatcmpl-pt1", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": "gpt-4o", + } + return ( + f"data: {json.dumps({**chunk, 'choices': [{'index': 0, 'delta': {'role': 'assistant', 'content': 'Hi'}}]})}\n\n".encode(), + f"data: {json.dumps({**chunk, 'choices': [{'index': 0, 'delta': {}, 'finish_reason': 'stop'}]})}\n\n".encode(), + f"data: {json.dumps({**chunk, 'choices': [], 'usage': {'prompt_tokens': 12, 'completion_tokens': 8, 'total_tokens': 20}})}\n\n".encode(), + b"data: [DONE]\n\n", + ) + + +def _delta_usages(text: str) -> list[Mapping[str, JsonValue]]: + events: Final = [json.loads(line[len("data: ") :]) for line in text.splitlines() if line.startswith("data: ")] + return [event["usage"] for event in events if event.get("type") == "message_delta" and "usage" in event] + + +def _cost_config(wire_url: str) -> str: + return ( + "model_list:\n" + " - model_name: amsg\n" + " litellm_params:\n" + f" model: anthropic/{_MODEL}\n" + f" api_base: {wire_url}\n" + f" api_key: {_KEY}\n" + " - model_name: omsg\n" + " litellm_params:\n" + " model: openai/gpt-4o-mini\n" + f" api_base: {wire_url}\n" + " api_key: synthetic-openai-key\n" + "litellm_settings:\n" + " include_cost_in_streaming_usage: true\n" + "general_settings:\n" + " master_key: os.environ/LITELLM_MASTER_KEY\n" + " database_url: os.environ/DATABASE_URL\n" + " store_model_in_db: true\n" + " disable_spend_logs: false\n" + " proxy_batch_write_at: 1\n" + ) + + +def _assert_cost_in_every_delta(gateway: Gateway, model: str) -> None: + text: Final = _stream_text( + gateway, + "/v1/messages", + {"model": model, "max_tokens": 20, "stream": True, "messages": [{"role": "user", "content": "Say 'Hi'"}]}, + ) + usages: Final = _delta_usages(text) + assert usages, (model, text) + costs: Final = [usage.get("cost") for usage in usages] + assert all(isinstance(cost, (int, float)) and cost > 0 for cost in costs), (model, text) + + +def test_streaming_cost_injected_into_usage_for_anthropic_and_openai_responses( + gateway: Gateway, tmp_path: Path +) -> None: + def respond(request: Request) -> Reply: + if request.target.endswith("/responses"): + return Reply(content_type="text/event-stream", chunks=_openai_responses_stream()) + assert request.target == "/v1/messages", request.target + return Reply(content_type="text/event-stream", chunks=_stream_chunks("msg_" + uuid.uuid4().hex)) + + with ( + _upstream(respond) as wire, + owned_proxy(gateway, tmp_path, {}, config=_owned_config(tmp_path, _cost_config(wire.url))) as candidate, + ): + _assert_cost_in_every_delta(candidate, "amsg") + _assert_cost_in_every_delta(candidate, "omsg") + targets: Final = [request.target for request in _provider_calls(wire)] + assert targets == ["/v1/messages", "/responses"], targets + + +def test_streaming_cost_injected_into_usage_for_openai_chat_completions_bridge( + gateway: Gateway, tmp_path: Path +) -> None: + def respond(request: Request) -> Reply: + assert request.target.endswith("/chat/completions"), request.target + assert json.loads(request.body)["stream"] is True + return Reply(content_type="text/event-stream", chunks=_openai_chat_stream()) + + with ( + _upstream(respond) as wire, + owned_proxy( + gateway, + tmp_path, + {"LITELLM_USE_CHAT_COMPLETIONS_URL_FOR_ANTHROPIC_MESSAGES": "true"}, + config=_owned_config(tmp_path, _cost_config(wire.url)), + ) as candidate, + ): + _assert_cost_in_every_delta(candidate, "omsg") + assert len(_provider_calls(wire)) == 1 diff --git a/tests/integration/observability/_openinference_support.py b/tests/integration/observability/_openinference_support.py new file mode 100644 index 00000000000..ee6dee07c46 --- /dev/null +++ b/tests/integration/observability/_openinference_support.py @@ -0,0 +1,1171 @@ +from __future__ import annotations + +import base64 +import json +import uuid +from collections.abc import Callable, Iterator, Mapping, Sequence +from contextlib import AbstractContextManager, ExitStack, contextmanager, nullcontext +from dataclasses import dataclass +from itertools import chain +from pathlib import Path +from typing import Final + +import httpx +import yaml +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.process import OwnedProxy, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest +from opentelemetry.proto.common.v1.common_pb2 import AnyValue +from pydantic import JsonValue, TypeAdapter + +JSON_OBJECT: Final[TypeAdapter[dict[str, JsonValue]]] = TypeAdapter(dict[str, JsonValue]) +JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +JSON_MESSAGES: Final[TypeAdapter[list[dict[str, JsonValue]]]] = TypeAdapter(list[dict[str, JsonValue]]) + +CHAT_TOOLS: Final = [ + { + "type": "function", + "function": { + "name": "lookup_weather", + "description": "Get weather", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string"}}, + }, + }, + } +] + +RESPONSES_TOOLS: Final = [ + { + "type": "function", + "name": "lookup_weather", + "description": "Get weather", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string"}}, + }, + } +] + +ANTHROPIC_TOOLS: Final = [ + { + "name": "lookup_weather", + "description": "Get weather", + "input_schema": { + "type": "object", + "properties": {"city": {"type": "string"}}, + }, + } +] + +_TOOL_PREFIX: Final = "llm.output_messages.{message}.message.tool_calls.{tool}.tool_call." + + +@dataclass(frozen=True, slots=True) +class Rig: + proxy: Gateway + owned: OwnedProxy + model: str + provider: Wire + destination: Wire + + +def _json_object(body: bytes) -> dict[str, JsonValue]: + return JSON_OBJECT.validate_json(body) + + +def _json_object_value(value: JsonValue) -> dict[str, JsonValue]: + return JSON_OBJECT.validate_python(value) + + +def _json_messages(value: str) -> list[dict[str, JsonValue]]: + return JSON_MESSAGES.validate_json(value) + + +def _assert_chat_request( + request: Request, + *, + messages: JsonValue, + tools: JsonValue = CHAT_TOOLS, + tool_choice: JsonValue | None = None, + include_tools: bool = True, + stream: bool | None = None, + stream_options: JsonValue | None = None, + n: int | None = None, +) -> dict[str, JsonValue]: + request_tools: Final = ( + { + "tools": tools, + "tool_choice": ( + tool_choice if tool_choice is not None else {"type": "function", "function": {"name": "lookup_weather"}} + ), + } + if include_tools + else {} + ) + expected: Final[dict[str, JsonValue]] = { + "model": "gpt-4o-mini", + "messages": messages, + **request_tools, + **({"stream": stream} if stream is not None else {}), + **({"stream_options": stream_options} if stream_options is not None else {}), + **({"n": n} if n is not None else {}), + } + observed: Final = _json_object(request.body) + assert observed == expected, (observed, expected) + return observed + + +def _chat_request_marker(request: Request) -> str: + body: Final = _json_object(request.body) + messages: Final = JSON_MESSAGES.validate_python(body["messages"]) + marker: Final = messages[0].get("content") + assert isinstance(marker, str), body + return marker + + +def _assert_responses_request( + request: Request, + *, + marker: str, + input_value: str = "weather in Paris?", + stream: bool = False, +) -> None: + expected: Final[dict[str, JsonValue]] = { + "model": "gpt-4o-mini", + "input": input_value, + "tools": RESPONSES_TOOLS, + "tool_choice": {"type": "function", "name": "lookup_weather"}, + "metadata": {"trace_marker": marker}, + **({"stream": True} if stream else {}), + } + observed: Final = _json_object(request.body) + assert observed == expected, request + + +def _assert_messages_request( + request: Request, + *, + marker: str, + prompt: str = "weather in Paris?", + stream: bool = False, +) -> None: + expected: Final[dict[str, JsonValue]] = { + "model": "claude-opus-5-5", + "max_tokens": 64, + "messages": [{"role": "user", "content": prompt}], + "tools": ANTHROPIC_TOOLS, + "tool_choice": {"type": "auto"}, + "metadata": {}, + "stream": stream, + } + observed: Final = _json_object(request.body) + assert observed == expected, request + + +def _stream_values(reply: Reply) -> tuple[dict[str, JsonValue], ...]: + return tuple( + _json_object(chunk.split(b"data: ", 1)[1].splitlines()[0]) + for chunk in reply.chunks or () + if b"data: " in chunk and b"[DONE]" not in chunk + ) + + +def _sse_json_values(body: bytes) -> tuple[dict[str, JsonValue], ...]: + return tuple( + _json_object(line.removeprefix(b"data: ").strip()) + for line in body.splitlines() + if line.startswith(b"data: ") and line != b"data: [DONE]" + ) + + +def _chat_caller_response(reply: Reply, model: str) -> dict[str, JsonValue]: + body: Final = _json_object(reply.body) + choices: Final = JSON_MESSAGES.validate_python(body["choices"]) + return { + **body, + "model": model, + "choices": [ + { + **choice, + "message": _chat_caller_message(_json_object_value(choice["message"])), + "provider_specific_fields": {}, + } + for choice in choices + ], + } + + +def _chat_caller_message(message: dict[str, JsonValue]) -> dict[str, JsonValue]: + raw_tool_calls: Final = message.get("tool_calls") + tool_calls: Final = JSON_MESSAGES.validate_python(raw_tool_calls) if isinstance(raw_tool_calls, list) else () + return { + **{key: value for key, value in message.items() if key != "tool_calls"}, + **({"tool_calls": [_chat_caller_tool_call(call) for call in tool_calls]} if tool_calls else {}), + "provider_specific_fields": {"refusal": None}, + } + + +def _chat_caller_tool_call(call: dict[str, JsonValue]) -> dict[str, JsonValue]: + function: Final = _json_object_value(call["function"]) + arguments: Final = function.get("arguments") + return { + **call, + "function": { + **function, + **( + {"arguments": json.dumps(arguments)} + if "arguments" in function and not isinstance(arguments, str) + else {} + ), + }, + } + + +def _chat_output_tool_call(call: dict[str, JsonValue]) -> dict[str, JsonValue]: + function: Final = _json_object_value(call["function"]) + return { + **call, + "id": call["id"] if isinstance(call.get("id"), str) else None, + "function": { + **function, + **({"name": None} if "name" not in function else {}), + }, + } + + +def _chat_output_value(reply: Reply) -> str: + body: Final = _json_object(reply.body) + choices: Final = JSON_MESSAGES.validate_python(body["choices"]) + assert len(choices) == 1, body + message: Final = _chat_caller_message(_json_object_value(choices[0]["message"])) + tool_calls: Final = JSON_MESSAGES.validate_python(message["tool_calls"]) if "tool_calls" in message else () + return json.dumps( + [ + { + **{ + key: value + for key, value in message.items() + if key not in {"provider_specific_fields", "tool_calls"} + }, + **({"tool_calls": [_chat_output_tool_call(call) for call in tool_calls]} if tool_calls else {}), + } + ] + ) + + +def _chat_caller_stream(reply: Reply, model: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(chain.from_iterable(_chat_caller_stream_events(event, model) for event in _stream_values(reply))) + + +def _chat_cache_hit_caller_stream( + marker: str, model: str, calls: Sequence[dict[str, JsonValue]] +) -> tuple[dict[str, JsonValue], ...]: + return ( + { + "id": marker, + "object": "chat.completion.chunk", + "created": 1, + "model": model, + "choices": [ + { + "index": 0, + "delta": { + "role": "assistant", + "tool_calls": [{"index": index, **call} for index, call in enumerate(calls)], + }, + } + ], + }, + { + "id": marker, + "object": "chat.completion.chunk", + "created": 1, + "model": model, + "choices": [{"index": 0, "delta": {}, "finish_reason": "tool_calls"}], + }, + ) + + +def _chat_caller_stream_event(event: dict[str, JsonValue], model: str) -> dict[str, JsonValue]: + choices: Final = JSON_MESSAGES.validate_python(event["choices"]) + return { + "id": event["id"], + "object": event["object"], + "created": 1, + "model": model, + "choices": [_chat_caller_stream_choice(choice) for choice in choices], + } + + +def _chat_caller_stream_events(event: dict[str, JsonValue], model: str) -> tuple[dict[str, JsonValue], ...]: + return ( + _chat_caller_stream_event(event, model), + *((_chat_caller_stream_usage_event(event, model),) if "usage" in event else ()), + ) + + +def _chat_caller_stream_usage_event(event: dict[str, JsonValue], model: str) -> dict[str, JsonValue]: + return { + "id": event["id"], + "object": event["object"], + "created": 1, + "model": model, + "choices": [{"index": 0, "delta": {}}], + "usage": {**_json_object_value(event["usage"]), "cost": 4.05e-6}, + } + + +def _chat_caller_stream_choice(choice: dict[str, JsonValue]) -> dict[str, JsonValue]: + delta: Final = _json_object_value(choice["delta"]) + tool_calls: Final = JSON_MESSAGES.validate_python(delta["tool_calls"]) if "tool_calls" in delta else None + finish_reason: Final = choice.get("finish_reason") + return { + **{ + key: value + for key, value in choice.items() + if key not in {"delta", "finish_reason", "index"} and value is not None + }, + "index": choice["index"], + "delta": { + **{key: value for key, value in delta.items() if key != "tool_calls" and value is not None}, + **( + {"tool_calls": [_chat_caller_stream_tool_call(call) for call in tool_calls]} + if tool_calls is not None + else {} + ), + }, + **({"finish_reason": finish_reason} if finish_reason is not None else {}), + } + + +def _chat_caller_stream_tool_call(call: dict[str, JsonValue]) -> dict[str, JsonValue]: + function: Final = _json_object_value(call["function"]) + return { + **{key: value for key, value in call.items() if key not in {"function", "type"}}, + "type": "function", + "function": {**function, **({"arguments": ""} if "arguments" not in function else {})}, + } + + +def _normalize_chat_caller_stream(events: Sequence[JsonValue]) -> tuple[dict[str, JsonValue], ...]: + return tuple({**_json_object_value(event), "created": 1} for event in events) + + +def _responses_caller_body(body: dict[str, JsonValue], model: str) -> dict[str, JsonValue]: + usage: Final = _json_object_value(body["usage"]) + output: Final = JSON_MESSAGES.validate_python(body["output"]) + return { + **body, + "id": "", + "model": model, + "error": None, + "incomplete_details": None, + "instructions": None, + "metadata": None, + "parallel_tool_calls": None, + "temperature": None, + "tool_choice": None, + "tools": None, + "top_p": None, + "max_output_tokens": None, + "previous_response_id": None, + "reasoning": None, + "text": None, + "truncation": None, + "user": None, + "store": None, + "output": [{**item, **({"namespace": None} if item.get("type") == "function_call" else {})} for item in output], + "usage": { + "input_tokens_details": None, + "output_tokens_details": None, + "cost": None, + **usage, + }, + } + + +def _responses_stream_caller_body(body: dict[str, JsonValue]) -> dict[str, JsonValue]: + response_id: Final = body.get("id") + assert isinstance(response_id, str), body + usage: Final = _json_object_value(body["usage"]) if body.get("status") == "completed" else {} + return { + **body, + "id": "", + **({"usage": {**usage, "cost": 4.05e-06}} if usage else {}), + } + + +def _responses_caller_response(reply: Reply, model: str) -> dict[str, JsonValue]: + return _responses_caller_body(_json_object(reply.body), model) + + +def _normalize_responses_caller_body(body: dict[str, JsonValue]) -> dict[str, JsonValue]: + response_id: Final = body.get("id") + assert isinstance(response_id, str) and response_id.startswith("resp_"), body + return {**body, "id": ""} + + +def _responses_caller_stream(reply: Reply, model: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + { + **event, + "model": model, + **( + {"response": _responses_stream_caller_body(_json_object_value(event["response"]))} + if "response" in event + else {} + ), + } + for event in _stream_values(reply) + ) + + +def _normalize_responses_caller_stream( + events: Sequence[dict[str, JsonValue]], +) -> tuple[dict[str, JsonValue], ...]: + return tuple( + { + **event, + **( + {"response": _normalize_responses_caller_body(_json_object_value(event["response"]))} + if "response" in event + else {} + ), + } + for event in events + ) + + +def _messages_caller_response(reply: Reply, model: str) -> dict[str, JsonValue]: + return {**_json_object(reply.body), "model": model} + + +def _messages_caller_raw_stream(reply: Reply, model: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + { + **event, + **( + {"message": {**_json_object_value(event["message"]), "model": model}} + if event.get("type") == "message_start" + else {} + ), + } + for event in _stream_values(reply) + ) + + +def _messages_caller_stream_response(reply: Reply, model: str) -> dict[str, JsonValue]: + body: Final = _messages_caller_response(reply, model) + content: Final = JSON_MESSAGES.validate_python(body["content"]) + return { + **body, + "content": [{**block, **({"caller": None} if block.get("type") == "tool_use" else {})} for block in content], + } + + +def _messages_caller_stream( + reply: Reply, model: str, *, final_message: dict[str, JsonValue] +) -> tuple[dict[str, JsonValue], ...]: + events: Final = tuple( + chain.from_iterable( + _messages_caller_stream_event(event, model, final_message) for event in _stream_values(reply) + ) + ) + return (*events, {"type": "message_stop", "message": final_message}) + + +def _messages_caller_stream_event( + event: dict[str, JsonValue], model: str, final_message: dict[str, JsonValue] +) -> tuple[dict[str, JsonValue], ...]: + if event.get("type") == "message_stop": + return () + if event.get("type") == "content_block_stop": + index: Final = event["index"] + assert isinstance(index, int), event + content: Final = JSON_MESSAGES.validate_python(final_message["content"]) + return ({**event, "content_block": content[index]},) + caller_event: Final = { + **event, + **({"message": {**_json_object_value(event["message"]), "model": model}} if "message" in event else {}), + } + if event.get("type") != "content_block_delta": + return (caller_event,) + delta: Final = _json_object_value(event["delta"]) + if delta.get("type") != "input_json_delta": + return (caller_event,) + partial_json: Final = delta["partial_json"] + assert isinstance(partial_json, str), event + return ( + caller_event, + { + "type": "input_json", + "partial_json": partial_json, + "snapshot": JSON_OBJECT.validate_json(partial_json), + }, + ) + + +def _span_attributes(request: Request) -> Iterator[dict[str, str]]: + if request.headers.get("content-type") != "application/x-protobuf": + return + batch: Final = ExportTraceServiceRequest.FromString(request.body) + for resource_spans in batch.resource_spans: + for scope_spans in resource_spans.scope_spans: + for span in scope_spans.spans: + yield {attribute.key: _attribute_text(attribute.value) for attribute in span.attributes} + + +def _attribute_text(value: object) -> str: + assert isinstance(value, AnyValue) + match value.WhichOneof("value"): + case "string_value": + return value.string_value + case "int_value": + return str(value.int_value) + case "double_value": + return str(value.double_value) + case "bool_value": + return str(value.bool_value) + case _: + return "" + + +def _spans(requests: tuple[Request, ...]) -> Iterator[dict[str, str]]: + for request in requests: + yield from _span_attributes(request) + + +def _canonical_response_id(value: str) -> str: + try: + return base64.b64decode(value.removeprefix("resp_").encode()).decode() + except (ValueError, UnicodeDecodeError): + return value + + +def _matching_llm_spans(requests: tuple[Request, ...], response_id: str) -> Iterator[dict[str, str]]: + wanted: Final = _canonical_response_id(response_id) + for attributes in _spans(requests): + if attributes.get("openinference.span.kind") != "LLM": + continue + observed: Final = attributes.get("gen_ai.response.id", "") + if _canonical_response_id(observed) == wanted or response_id in attributes.values(): + yield attributes + + +def _matching_marker_spans(requests: tuple[Request, ...], marker: str) -> Iterator[dict[str, str]]: + for attributes in _spans(requests): + if attributes.get("openinference.span.kind") == "LLM" and marker in attributes.values(): + yield attributes + + +def _single_span(spans: tuple[dict[str, str], ...]) -> dict[str, str]: + assert len(spans) == 1, spans + return spans[0] + + +def _matching_span(destination: Wire, response_id: str) -> dict[str, str]: + spans: Final = eventually( + lambda: tuple(_matching_llm_spans(destination.drain(), response_id)), + bool, + seconds=30, + ) + return _single_span(spans) + + +def _matching_marker_span(destination: Wire, marker: str) -> dict[str, str]: + spans: Final = eventually( + lambda: tuple(_matching_marker_spans(destination.drain(), marker)), + bool, + seconds=30, + ) + return _single_span(spans) + + +def _matching_output_value_span(destination: Wire, marker: str) -> dict[str, str]: + spans: Final = eventually( + lambda: tuple( + attributes + for attributes in _spans(destination.drain()) + if attributes.get("openinference.span.kind") == "LLM" and marker in attributes.get("output.value", "") + ), + bool, + seconds=30, + ) + return _single_span(spans) + + +def _matching_genai_marker_span(destination: Wire, marker: str) -> dict[str, str]: + spans: Final = eventually( + lambda: tuple( + attributes + for attributes in _spans(destination.drain()) + if attributes.get("gen_ai.operation.name") == "chat" and marker in attributes.values() + ), + bool, + seconds=30, + ) + return _single_span(spans) + + +def _matching_any_marker_span(destination: Wire, marker: str) -> dict[str, str]: + def matches(requests: tuple[Request, ...]) -> tuple[dict[str, str], ...]: + return tuple(attributes for attributes in _spans(requests) if marker in attributes.values()) + + spans: Final = eventually( + lambda: matches(destination.drain()), + bool, + seconds=30, + ) + return _single_span(spans) + + +def _collect_marker_spans( + destination: Wire, markers: tuple[str, ...], *, timeout_seconds: float = 30 +) -> tuple[dict[str, str], ...]: + expected: Final = frozenset(markers) + + def matches(requests: tuple[Request, ...]) -> tuple[dict[str, str], ...]: + return tuple( + attributes + for attributes in _spans(requests) + if attributes.get("openinference.span.kind") == "LLM" + and any(marker in attributes.values() for marker in expected) + ) + + def complete(spans: tuple[dict[str, str], ...]) -> bool: + return all(any(marker in attributes.values() for attributes in spans) for marker in expected) + + def collect(previous: tuple[dict[str, str], ...]) -> tuple[dict[str, str], ...]: + current: Final = eventually(lambda: matches(destination.drain()), bool, seconds=timeout_seconds) + combined: Final = (*previous, *current) + return combined if complete(combined) else collect(combined) + + return collect(()) + + +def _llm_spans_through_markers(destination: Wire, markers: tuple[str, ...]) -> tuple[dict[str, str], ...]: + expected: Final = frozenset(markers) + + def llm_spans(requests: tuple[Request, ...]) -> tuple[dict[str, str], ...]: + return tuple( + attributes for attributes in _spans(requests) if attributes.get("openinference.span.kind") == "LLM" + ) + + def complete(spans: tuple[dict[str, str], ...]) -> bool: + return all(any(marker in attributes.values() for attributes in spans) for marker in expected) + + def collect(previous: tuple[dict[str, str], ...]) -> tuple[dict[str, str], ...]: + current: Final = eventually(lambda: llm_spans(destination.drain()), bool, seconds=30) + combined: Final = (*previous, *current) + return combined if complete(combined) else collect(combined) + + return collect(()) + + +def _write_config( + directory: Path, + *, + callbacks: Sequence[str] = ("arize",), + callback_settings: Mapping[str, JsonValue] | None = None, + litellm_settings: Mapping[str, JsonValue] | None = None, + general_settings: Mapping[str, JsonValue] | None = None, +) -> Path: + loaded: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config: Final = { + **loaded, + "litellm_settings": { + **loaded["litellm_settings"], + "callbacks": list(callbacks), + **(litellm_settings or {}), + }, + **({"callback_settings": dict(callback_settings)} if callback_settings is not None else {}), + "general_settings": { + **loaded["general_settings"], + **(general_settings or {}), + }, + } + path: Final = directory / f"openinference-{uuid.uuid4().hex}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _environment(destination: Wire, extra: Mapping[str, str] | None = None) -> dict[str, str]: + return { + "LITELLM_OTEL_V2": "1", + "OTEL_BSP_SCHEDULE_DELAY": "100", + "OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT": "span_only", + "LITELLM_OTEL_BAGGAGE_METADATA_KEYS": "requester_metadata.trace_marker", + "ARIZE_HTTP_ENDPOINT": destination.url + "/v1/traces", + "ARIZE_SPACE_ID": "integration-space", + "ARIZE_API_KEY": "integration-arize-key", + **(extra or {}), + } + + +def _collector(_request: Request) -> Reply: + return Reply(body=b"", content_type="application/x-protobuf") + + +def _owned_sink_handler(sink: Callable[[Request], Reply]) -> Callable[[Request], Reply]: + def handle(request: Request) -> Reply: + if ( + request.method != "POST" + or not request.body + or request.headers.get("content-type", "").split(";", maxsplit=1)[0].strip().lower() + != "application/x-protobuf" + ): + return Reply(body=b"", content_type="application/x-protobuf") + try: + return sink(request) + except (AssertionError, IndexError, KeyError, TypeError, ValueError) as error: + raise AssertionError(f"{request.method} {request.target}: {error!r}") from error + + return handle + + +def _provider_handler(upstream: Callable[[Request], Reply]) -> Callable[[Request], Reply]: + def handle(request: Request) -> Reply: + if ( + request.method != "POST" + or not request.body + or request.headers.get("content-type", "").split(";", maxsplit=1)[0].strip().lower() != "application/json" + ): + return Reply(body=b"{}", content_type="application/json") + try: + return upstream(request) + except (AssertionError, IndexError, KeyError, TypeError, ValueError) as error: + raise AssertionError(f"{request.method} {request.target}: {error!r}") from error + + return handle + + +@contextmanager +def _rig( + gateway: Gateway, + directory: Path, + upstream: Callable[[Request], Reply], + *, + callbacks: Sequence[str] = ("arize",), + callback_settings: Mapping[str, JsonValue] | None = None, + litellm_settings: Mapping[str, JsonValue] | None = None, + general_settings: Mapping[str, JsonValue] | None = None, + environment: Mapping[str, str] | None = None, + remove_environment: tuple[str, ...] = (), + disabled_environment: tuple[str, ...] = (), + model_name: str = "openai/gpt-4o-mini", + api_base_suffix: str = "/v1", + destination_handler: Callable[[Request], Reply] | None = None, + destination_wire: Wire | None = None, + fresh_client_connections: bool = False, + workers: int = 2, +) -> Iterator[Rig]: + destination_context: Final[AbstractContextManager[Wire]] = ( + nullcontext(destination_wire) + if destination_wire is not None + else wire_server(_owned_sink_handler(destination_handler or _collector)) + ) + with wire_server(_provider_handler(upstream)) as provider, destination_context as destination: + resolved_settings: Final = { + key: ( + {**value, "endpoint": destination.url + "/v1/traces"} + if key == "otel" and isinstance(value, dict) and value.get("endpoint") == "unused" + else value + ) + for key, value in (callback_settings or {}).items() + } + config: Final = _write_config( + directory, + callbacks=callbacks, + callback_settings=resolved_settings, + litellm_settings=litellm_settings, + general_settings=general_settings, + ) + preset_environment: Final = { + **({"PHOENIX_COLLECTOR_ENDPOINT": destination.url + "/v1/traces"} if "arize_phoenix" in callbacks else {}), + **( + { + "WANDB_HOST": destination.url, + "WANDB_API_KEY": "integration-weave-key", + "WANDB_PROJECT_ID": "integration/project", + } + if "weave_otel" in callbacks + else {} + ), + **( + { + "LANGFUSE_OTEL_HOST": destination.url, + "LANGFUSE_PUBLIC_KEY": "integration-public", + "LANGFUSE_SECRET_KEY": "integration-secret", + } + if "langfuse_otel" in callbacks + else {} + ), + **( + { + "LEVOAI_API_KEY": "integration-levo-key", + "LEVOAI_ORG_ID": "integration-org", + "LEVOAI_WORKSPACE_ID": "integration-workspace", + "LEVOAI_COLLECTOR_URL": destination.url + "/v1/traces", + } + if "levo" in callbacks + else {} + ), + **( + { + "SIGNOZ_INGESTION_ENDPOINT": destination.url + "/v1/traces", + "SIGNOZ_INGESTION_KEY": "integration-signoz-key", + } + if "signoz" in callbacks + else {} + ), + **( + {"OTEL_EXPORTER_OTLP_ENDPOINT": destination.url} + if any(callback_name in callbacks for callback_name in ("langtrace", "newrelic", "agentops")) + else {} + ), + } + overrides: Final = { + key: value + for key, value in _environment(destination, {**preset_environment, **(environment or {})}).items() + if key not in disabled_environment + } + with ( + owned_proxy_process( + gateway, + directory, + overrides, + config=config, + remove_environment=remove_environment, + workers=workers, + ) as owned, + ExitStack() as resources, + ): + proxy_client: Final = ( + resources.enter_context( + httpx.Client( + base_url=str(owned.gateway.client.base_url), + timeout=15, + trust_env=False, + limits=httpx.Limits(max_keepalive_connections=0), + ) + ) + if fresh_client_connections + else owned.gateway.client + ) + proxy: Final = Gateway( + client=proxy_client, + key=owned.gateway.key, + upstream_url=owned.gateway.upstream_url, + ) + scenario: Final = resources.enter_context(proxy.scenario()) + model: Final = scenario.model( + model=model_name, + api_base=provider.url + api_base_suffix, + ) + yield Rig(proxy, owned, model, provider, destination) + + +def _chat_tool_call(identity: str, city: str = "Paris") -> dict[str, JsonValue]: + return { + "id": "call_" + identity, + "type": "function", + "function": { + "name": "lookup_weather", + "arguments": json.dumps({"city": city}), + }, + } + + +def _chat_response(identity: str, calls: Sequence[dict[str, JsonValue]] | None = None) -> Reply: + tool_calls: Final = list(calls if calls is not None else (_chat_tool_call(identity),)) + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "finish_reason": "tool_calls", + "message": {"role": "assistant", "content": None, "tool_calls": tool_calls}, + } + ], + "usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}, + } + ).encode() + ) + + +def _chat_plain_response(identity: str, content: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": content}, + } + ], + "usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9}, + } + ).encode() + ) + + +def _chat_stream_response(identity: str, calls: Sequence[dict[str, JsonValue]], *, include_usage: bool = True) -> Reply: + def first_chunk(index: int, call: dict[str, JsonValue]) -> dict[str, JsonValue]: + fields: Final = object_value(JSON_VALUE.validate_python(call)) + return { + "index": index, + "id": string_value(fields["id"]), + "type": "function", + "function": {"name": "lookup_weather"}, + } + + def arguments_chunk(index: int, call: dict[str, JsonValue]) -> dict[str, JsonValue]: + fields: Final = object_value(JSON_VALUE.validate_python(call)) + function: Final = object_value(fields["function"]) + return { + "id": identity, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "delta": { + "tool_calls": [ + { + "index": index, + "function": {"arguments": string_value(function["arguments"])}, + } + ] + }, + "finish_reason": None, + } + ], + } + + chunks: Final = [ + { + "id": identity, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "delta": { + "role": "assistant", + "tool_calls": [first_chunk(index, call) for index, call in enumerate(calls)], + }, + "finish_reason": None, + } + ], + }, + *(arguments_chunk(index, call) for index, call in enumerate(calls)), + { + "id": identity, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "delta": {}, "finish_reason": "tool_calls"}], + **({"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}} if include_usage else {}), + }, + ] + return Reply( + content_type="text/event-stream", + chunks=tuple(b"data: " + json.dumps(chunk).encode() + b"\n\n" for chunk in chunks) + (b"data: [DONE]\n\n",), + ) + + +def _responses_response(identity: str, calls: Sequence[dict[str, JsonValue]] | None = None) -> Reply: + def response_item(call: dict[str, JsonValue]) -> dict[str, JsonValue]: + fields: Final = object_value(JSON_VALUE.validate_python(call)) + function: Final = object_value(fields["function"]) + call_id: Final = string_value(fields["id"]) + return { + "type": "function_call", + "id": "fc_" + call_id, + "call_id": call_id, + "name": string_value(function["name"]), + "arguments": string_value(function["arguments"]), + "status": "completed", + } + + response: Final = { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [response_item(call) for call in calls if calls is not None] + if calls is not None + else [response_item(_chat_tool_call(identity))], + "usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15}, + } + return Reply(body=json.dumps(response).encode()) + + +def _responses_stream_response(identity: str, calls: Sequence[dict[str, JsonValue]]) -> Reply: + response: Final = _json_object(_responses_response(identity, calls).body) + output: Final = response["output"] + items: Final = tuple(object_value(item) for item in output) if isinstance(output, list) else () + events: Final = ( + { + "type": "response.created", + "sequence_number": 0, + "response": {**response, "status": "in_progress", "output": []}, + }, + *( + { + "type": "response.output_item.added", + "sequence_number": index + 1, + "output_index": index, + "item": item, + } + for index, item in enumerate(items) + ), + *( + { + "type": "response.function_call_arguments.delta", + "sequence_number": index + len(items) + 1, + "item_id": str(item["id"]), + "output_index": index, + "delta": str(item["arguments"]), + } + for index, item in enumerate(items) + ), + { + "type": "response.completed", + "sequence_number": len(items) * 2 + 1, + "response": response, + }, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _anthropic_response(identity: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": identity, + "type": "message", + "role": "assistant", + "model": "claude-opus-5-5", + "content": [ + { + "type": "tool_use", + "id": "call_" + identity, + "name": "lookup_weather", + "input": {"city": "Paris"}, + } + ], + "stop_reason": "tool_use", + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 4}, + } + ).encode() + ) + + +def _anthropic_stream_response(identity: str) -> Reply: + message: Final = { + "id": identity, + "type": "message", + "role": "assistant", + "model": "claude-opus-5-5", + "content": [ + { + "type": "tool_use", + "id": "call_" + identity, + "name": "lookup_weather", + "input": {"city": "Paris"}, + } + ], + "stop_reason": "tool_use", + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 4}, + } + events: Final = ( + {"type": "message_start", "message": {**message, "content": [], "stop_reason": None}}, + { + "type": "content_block_start", + "index": 0, + "content_block": { + "type": "tool_use", + "id": "call_" + identity, + "name": "lookup_weather", + "input": {}, + }, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "input_json_delta", "partial_json": '{"city": "Paris"}'}, + }, + {"type": "content_block_stop", "index": 0}, + { + "type": "message_delta", + "delta": {"stop_reason": "tool_use", "stop_sequence": None}, + "usage": {"output_tokens": 4}, + }, + {"type": "message_stop"}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _tool_call_attributes( + attributes: Mapping[str, str], + *, + message_index: int = 0, + tool_index: int = 0, +) -> tuple[str, str, dict[str, JsonValue]]: + prefix: Final = _TOOL_PREFIX.format(message=message_index, tool=tool_index) + return ( + attributes[prefix + "id"], + attributes[prefix + "function.name"], + JSON_OBJECT.validate_json(attributes[prefix + "function.arguments"]), + ) + + +def _assert_tool_span( + attributes: Mapping[str, str], + *, + marker: str, + output: JsonValue, + calls: Sequence[tuple[str, str, Mapping[str, JsonValue]]], + metadata: Mapping[str, JsonValue] | None = None, + baggage: Mapping[str, str] | None = None, +) -> None: + for index, (call_id, name, arguments) in enumerate(calls): + observed: Final = _tool_call_attributes(attributes, tool_index=index) + assert observed == (call_id, name, dict(arguments)), f"tool call {index} for {marker}: {observed!r}" + observed_output: Final = JSON_MESSAGES.validate_json(attributes["output.value"]) + assert observed_output == output, (marker, observed_output) + if metadata is None: + assert "metadata" not in attributes, attributes + else: + assert json.loads(attributes["metadata"]) == dict(metadata), attributes + for key, value in (baggage or {}).items(): + assert attributes.get("litellm.metadata." + key) == value, attributes + + +def _response_tool_calls(identity: str, cities: Sequence[str] = ("Paris",)) -> list[dict[str, JsonValue]]: + return [_chat_tool_call(identity + "-" + city, city) for city in cities] diff --git a/tests/integration/observability/test_akto_logging_only.py b/tests/integration/observability/test_akto_logging_only.py new file mode 100644 index 00000000000..31adacdec75 --- /dev/null +++ b/tests/integration/observability/test_akto_logging_only.py @@ -0,0 +1,547 @@ +"""Akto guardrail in `mode: logging_only`, driven through a real proxy. + +The Akto service is the only guardrail double: an owned wire peer that answers the `/api/http-proxy` +verdict protocol. The provider is a second owned peer. The proxy, its guardrail registry, Postgres and +Redis run for real with two workers. Every test waits for the spend row, which is written after the +logging-only scans finish, before it reads what Akto received. +""" + +from __future__ import annotations + +import asyncio +import json +import uuid +from collections.abc import Iterator +from dataclasses import dataclass +from pathlib import Path +from typing import Final +from urllib.parse import parse_qs, urlsplit + +import anthropic +import openai +import pytest +import yaml +from integration._support.client import Gateway, eventually, gateway_from_environment, object_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server + +LOG_KEY: Final = "synthetic-akto-log-key" +INPUT_KEY: Final = "synthetic-akto-input-key" +OUTPUT_KEY: Final = "synthetic-akto-output-key" +DOWN_KEY: Final = "synthetic-akto-down-key" +DOWN_OPEN_KEY: Final = "synthetic-akto-down-open-key" +MIXED_KEY: Final = "synthetic-akto-mixed-key" +BLOCK_MARK: Final = "SYNTHETIC-AKTO-BLOCK" +BLOCK_REASON: Final = "Synthetic Akto policy block" +AKTO_DROP_MARK: Final = "SYNTHETIC-AKTO-DROP" +AKTO_ERROR_MARK: Final = "SYNTHETIC-AKTO-500" +AKTO_GARBAGE_MARK: Final = "SYNTHETIC-AKTO-GARBAGE" +PROVIDER_FAIL_MARK: Final = "SYNTHETIC-PROVIDER-FAIL" +REQUEST_CHECK: Final = {"akto_connector": "litellm", "guardrails": "true", "ingest_data": "true"} +RESPONSE_CHECK: Final = {"akto_connector": "litellm", "response_guardrails": "true", "ingest_data": "true"} +FILE_CHECK: Final = {"akto_connector": "litellm", "file_guardrails": "true"} + + +@dataclass(frozen=True, slots=True) +class AktoCall: + path: str + flags: dict[str, str] + authorization: str + payload: dict[str, object] + + def request_text(self) -> str: + return str(self.payload.get("requestPayload", "")) + + def response_text(self) -> str: + return str(self.payload.get("responsePayload", "")) + + +def _akto_call(request: Request) -> AktoCall: + target: Final = urlsplit(request.target) + return AktoCall( + path=target.path, + flags={name: values[0] for name, values in parse_qs(target.query).items()}, + authorization=request.headers.get("authorization", ""), + payload=json.loads(request.body), + ) + + +def _akto_verdict(request: Request) -> Reply: + if AKTO_DROP_MARK.encode() in request.body: + return Reply(drop_connection=True) + if AKTO_ERROR_MARK.encode() in request.body: + return Reply(status=500, body=b'{"error": "synthetic akto failure"}') + if AKTO_GARBAGE_MARK.encode() in request.body: + return Reply(body=b"synthetic akto garbage", content_type="text/plain") + verdict: Final = ( + {"Allowed": False, "Behaviour": "block", "Reason": BLOCK_REASON} + if BLOCK_MARK.encode() in request.body + else {"Allowed": True} + ) + return Reply(body=json.dumps({"data": {"guardrailsResult": verdict}}).encode()) + + +def _answer(marker: str) -> str: + return "synthetic answer " + marker + + +def _marker_in(body: bytes) -> str: + text: Final = body.decode() + start: Final = text.find("mark-") + assert start >= 0, text + return text[start : start + 37] + + +def _sse(events: tuple[dict[str, object], ...]) -> tuple[bytes, ...]: + return tuple(("data: " + json.dumps(event) + "\n\n").encode() for event in events) + + +def _chat_reply(marker: str, streaming: bool) -> Reply: + answer: Final = _answer(marker) + if not streaming: + return Reply( + body=json.dumps( + { + "id": "chatcmpl-" + marker, + "object": "chat.completion", + "created": 1, + "model": "gpt-5.4-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": answer}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + ).encode() + ) + + def chunk(delta: dict[str, object], finish: str | None) -> dict[str, object]: + return { + "id": "chatcmpl-" + marker, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-5.4-mini", + "choices": [{"index": 0, "delta": delta, "finish_reason": finish}], + } + + events: Final = (chunk({"role": "assistant", "content": answer[:9]}, None), chunk({"content": answer[9:]}, "stop")) + return Reply(chunks=(*_sse(events), b"data: [DONE]\n\n"), content_type="text/event-stream") + + +def _messages_reply(marker: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": "msg_" + marker, + "type": "message", + "role": "assistant", + "model": "claude-haiku-5-5", + "content": [{"type": "text", "text": _answer(marker)}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 10, "output_tokens": 5}, + } + ).encode() + ) + + +def _responses_reply(marker: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": "resp_" + marker, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-5.4-mini", + "output": [ + { + "type": "message", + "id": "msgo_" + marker, + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": _answer(marker), "annotations": []}], + } + ], + "usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15}, + } + ).encode() + ) + + +def _provider(request: Request) -> Reply: + if PROVIDER_FAIL_MARK.encode() in request.body: + return Reply(status=500, body=b'{"error": {"message": "synthetic provider failure", "type": "server_error"}}') + marker: Final = _marker_in(request.body) + if request.target.endswith("/v1/messages"): + return _messages_reply(marker) + if request.target.endswith("/v1/responses"): + return _responses_reply(marker) + assert request.target.endswith("/v1/chat/completions"), request.target + return _chat_reply(marker, bool(json.loads(request.body).get("stream"))) + + +def _guardrail(name: str, key: str, url: str, **params: object) -> dict[str, object]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "akto", + "mode": "logging_only", + "default_on": True, + "akto_base_url": url, + "akto_api_key": key, + **params, + }, + } + + +def _rig_config(akto_url: str, down_url: str, root: Path) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["cache"] = False + config["guardrails"] = [ + _guardrail("akto-log", LOG_KEY, akto_url), + _guardrail("akto-log-input", INPUT_KEY, akto_url, logging_only_scope="input"), + _guardrail("akto-log-output", OUTPUT_KEY, akto_url, logging_only_scope="output"), + _guardrail("akto-log-down", DOWN_KEY, down_url), + _guardrail("akto-log-down-open", DOWN_OPEN_KEY, down_url, unreachable_fallback="fail_open"), + _guardrail("akto-mixed", MIXED_KEY, akto_url, mode=["pre_call", "logging_only"], default_on=False), + ] + path: Final = root / "akto-logging-only.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@dataclass(frozen=True, slots=True) +class Rig: + proxy: Gateway + akto: Wire + akto_down: Wire + provider: Wire + chat_model: str + claude_model: str + responses_model: str + + def base(self) -> str: + return str(self.proxy.client.base_url).rstrip("/") + + def auth(self) -> dict[str, str]: + return {"Authorization": f"Bearer {self.proxy.key}"} + + def provider_calls(self, marker: str) -> tuple[Request, ...]: + return tuple(request for request in self.provider.drain() if marker.encode() in request.body) + + def guardrail_entries(self, response_id: str, name: str) -> tuple[dict[str, object], ...]: + rows: Final = spend_rows(response_id) + assert len(rows) == 1, rows + entries: Final = object_value(rows[0]["metadata"]).get("guardrail_information") + assert isinstance(entries, list), rows[0] + return tuple(entry for entry in (object_value(item) for item in entries) if entry.get("guardrail_name") == name) + + def settled_akto_calls(self, response_id: str, marker: str, key: str) -> tuple[AktoCall, ...]: + self.guardrail_entries(response_id, "akto-log") + calls: Final = tuple(_akto_call(request) for request in self.akto.drain() if marker.encode() in request.body) + return tuple(call for call in calls if call.authorization == key) + + +def spend_rows(request_id: str) -> tuple[dict[str, object], ...]: + return tuple( + eventually( + lambda: read_rows( + 'SELECT request_id, status, metadata FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (request_id,) + ), + lambda values: len(values) >= 1, + seconds=70, + ) + ) + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]: + root: Final = tmp_path_factory.mktemp("akto-logging-only") + with ( + gateway_from_environment() as gateway, + wire_server(_provider) as provider, + wire_server(_akto_verdict) as akto, + wire_server(lambda _: Reply(drop_connection=True)) as akto_down, + owned_proxy_process(gateway, root, {}, config=_rig_config(akto.url, akto_down.url, root), workers=2) as owned, + owned.gateway.scenario() as scenario, + ): + chat: Final = scenario.model( + model="openai/gpt-5.4-mini", api_base=provider.url + "/v1", api_key="synthetic-openai-key" + ) + claude: Final = scenario.model( + model="anthropic/claude-haiku-5-5", api_base=provider.url, api_key="synthetic-anthropic-key" + ) + responses: Final = scenario.model( + model="openai/gpt-5.4-mini", api_base=provider.url + "/v1", api_key="synthetic-openai-key" + ) + yield Rig(owned.gateway, akto, akto_down, provider, chat, claude, responses) + + +def _marker() -> str: + return "mark-" + uuid.uuid4().hex + + +def _assert_checked_both_ways(calls: tuple[AktoCall, ...], marker: str) -> None: + assert [call.flags for call in calls] == [REQUEST_CHECK, RESPONSE_CHECK], calls + assert all(call.path == "/api/http-proxy" for call in calls), calls + assert marker in calls[0].request_text(), calls[0].payload + assert _answer(marker) in calls[1].response_text(), calls[1].payload + + +def _chat(rig: Rig, content: object, **extra: object) -> dict[str, object]: + response: Final = rig.proxy.client.post( + "/v1/chat/completions", + json={"model": rig.chat_model, "messages": [{"role": "user", "content": content}], **extra}, + headers=rig.auth(), + ) + assert response.status_code == 200, response.text + return response.json() + + +def _content(body: dict[str, object]) -> object: + choices: Final = body["choices"] + assert isinstance(choices, list), body + return object_value(object_value(choices[0])["message"])["content"] + + +def test_logging_only_checks_a_sync_openai_chat_request_and_response(rig: Rig) -> None: + marker: Final = _marker() + client: Final = openai.OpenAI(base_url=rig.base() + "/v1", api_key=rig.proxy.key, max_retries=0) + completion: Final = client.chat.completions.create( + model=rig.chat_model, messages=[{"role": "user", "content": "hello " + marker}] + ) + assert completion.choices[0].message.content == _answer(marker) + assert len(rig.provider_calls(marker)) == 1 + + _assert_checked_both_ways(rig.settled_akto_calls(completion.id, marker, LOG_KEY), marker) + + +@pytest.mark.asyncio +async def test_logging_only_checks_an_async_openai_chat_request_and_response(rig: Rig) -> None: + marker: Final = _marker() + client: Final = openai.AsyncOpenAI(base_url=rig.base() + "/v1", api_key=rig.proxy.key, max_retries=0) + completion: Final = await client.chat.completions.create( + model=rig.chat_model, messages=[{"role": "user", "content": "hello " + marker}] + ) + assert completion.choices[0].message.content == _answer(marker) + + _assert_checked_both_ways(rig.settled_akto_calls(completion.id, marker, LOG_KEY), marker) + + +def test_logging_only_block_verdict_never_blocks_the_caller(rig: Rig) -> None: + marker: Final = _marker() + body: Final = _chat(rig, f"{BLOCK_MARK} {marker}") + assert _content(body) == _answer(marker) + assert len(rig.provider_calls(marker)) == 1 + + calls: Final = rig.settled_akto_calls(str(body["id"]), marker, LOG_KEY) + assert [call.flags for call in calls] == [REQUEST_CHECK], calls + entries: Final = rig.guardrail_entries(str(body["id"]), "akto-log") + assert [(entry["guardrail_mode"], entry["guardrail_status"]) for entry in entries] == [ + ("logging_only", "guardrail_intervened") + ], entries + + +def test_logging_only_scope_input_and_output_each_check_one_direction(rig: Rig) -> None: + marker: Final = _marker() + body: Final = _chat(rig, "scoped " + marker) + + input_calls: Final = rig.settled_akto_calls(str(body["id"]), marker, INPUT_KEY) + assert [call.flags for call in input_calls] == [REQUEST_CHECK], input_calls + assert marker in input_calls[0].request_text(), input_calls[0].payload + entries: Final = rig.guardrail_entries(str(body["id"]), "akto-log-output") + assert [entry["guardrail_status"] for entry in entries] == ["success"], entries + + +def test_logging_only_scope_output_checks_only_the_response(rig: Rig) -> None: + marker: Final = _marker() + body: Final = _chat(rig, "scoped output " + marker) + + output_calls: Final = rig.settled_akto_calls(str(body["id"]), marker, OUTPUT_KEY) + assert [call.flags for call in output_calls] == [RESPONSE_CHECK], output_calls + assert _answer(marker) in output_calls[0].response_text(), output_calls[0].payload + + +def test_unreachable_akto_under_logging_only_never_fails_the_caller(rig: Rig) -> None: + marker: Final = _marker() + body: Final = _chat(rig, "outage " + marker) + assert _content(body) == _answer(marker) + + entries: Final = rig.guardrail_entries(str(body["id"]), "akto-log-down") + assert [entry["guardrail_mode"] for entry in entries] == ["logging_only"], entries + dropped: Final = [_akto_call(request) for request in rig.akto_down.drain() if marker.encode() in request.body] + assert [call.flags for call in dropped if call.authorization == DOWN_KEY] == [REQUEST_CHECK], dropped + + +def test_unreachable_akto_with_fail_open_still_attempts_the_response_check(rig: Rig) -> None: + marker: Final = _marker() + body: Final = _chat(rig, "outage open " + marker) + assert _content(body) == _answer(marker) + + entries: Final = rig.guardrail_entries(str(body["id"]), "akto-log-down-open") + assert [entry["guardrail_mode"] for entry in entries] == ["logging_only", "logging_only"], entries + dropped: Final = [_akto_call(request) for request in rig.akto_down.drain() if marker.encode() in request.body] + assert [call.flags for call in dropped if call.authorization == DOWN_OPEN_KEY] == [REQUEST_CHECK, RESPONSE_CHECK] + + +def test_logging_only_streaming_chat_sends_the_assembled_answer(rig: Rig) -> None: + marker: Final = _marker() + client: Final = openai.OpenAI(base_url=rig.base() + "/v1", api_key=rig.proxy.key, max_retries=0) + chunks: Final = tuple( + client.chat.completions.create( + model=rig.chat_model, messages=[{"role": "user", "content": "stream " + marker}], stream=True + ) + ) + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == _answer(marker) + + _assert_checked_both_ways(rig.settled_akto_calls(chunks[0].id, marker, LOG_KEY), marker) + + +def test_logging_only_anthropic_messages_checks_both_directions(rig: Rig) -> None: + marker: Final = _marker() + client: Final = anthropic.Anthropic(base_url=rig.base(), api_key=rig.proxy.key, max_retries=0) + message: Final = client.messages.create( + model=rig.claude_model, max_tokens=64, messages=[{"role": "user", "content": "claude " + marker}] + ) + assert message.content[0].type == "text" and message.content[0].text == _answer(marker) + + _assert_checked_both_ways(rig.settled_akto_calls(message.id, marker, LOG_KEY), marker) + + +def test_logging_only_responses_api_checks_both_directions(rig: Rig) -> None: + marker: Final = _marker() + client: Final = openai.OpenAI(base_url=rig.base() + "/v1", api_key=rig.proxy.key, max_retries=0) + result: Final = client.responses.create(model=rig.responses_model, input="responses " + marker) + assert result.output_text == _answer(marker) + + _assert_checked_both_ways(rig.settled_akto_calls(result.id, marker, LOG_KEY), marker) + + +def test_logging_only_sends_each_chat_attachment_to_akto_once(rig: Rig) -> None: + marker: Final = _marker() + image_url: Final = f"https://example.com/{marker}.png" + body: Final = _chat( + rig, [{"type": "text", "text": "describe " + marker}, {"type": "image_url", "image_url": {"url": image_url}}] + ) + + calls: Final = rig.settled_akto_calls(str(body["id"]), marker, LOG_KEY) + file_checks: Final = [call.payload["files"] for call in calls if call.flags == FILE_CHECK] + assert file_checks == [[{"filename": f"{marker}.png", "type": "image", "url": image_url}]], calls + + +def test_logging_only_sends_each_responses_attachment_to_akto_once(rig: Rig) -> None: + marker: Final = _marker() + image_url: Final = f"https://example.com/{marker}.png" + client: Final = openai.OpenAI(base_url=rig.base() + "/v1", api_key=rig.proxy.key, max_retries=0) + result: Final = client.responses.create( + model=rig.responses_model, + input=[ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "describe " + marker}, + {"type": "input_image", "image_url": image_url, "detail": "auto"}, + ], + } + ], + ) + + calls: Final = rig.settled_akto_calls(result.id, marker, LOG_KEY) + file_checks: Final = [call.payload["files"] for call in calls if call.flags == FILE_CHECK] + assert file_checks == [[{"filename": f"{marker}.png", "type": "image", "url": image_url}]], calls + + +def test_pre_call_plus_logging_only_checks_the_request_inline_and_both_directions_after(rig: Rig) -> None: + marker: Final = _marker() + body: Final = _chat(rig, "mixed " + marker, guardrails=["akto-mixed"]) + + flags: Final = [call.flags for call in rig.settled_akto_calls(str(body["id"]), marker, MIXED_KEY)] + assert (flags.count(REQUEST_CHECK), flags.count(RESPONSE_CHECK), len(flags)) == (2, 1, 3), flags + + +def test_dashboard_offers_logging_only_for_akto(rig: Rig) -> None: + response: Final = rig.proxy.client.get("/guardrails/ui/add_guardrail_settings", headers=rig.auth()) + assert response.status_code == 200, response.text + settings: Final = response.json() + assert "logging_only" in settings["supported_modes_by_provider"]["akto"], settings["supported_modes_by_provider"] + assert "akto" not in settings["providers_without_directional_logging_only_scope"] + + +@pytest.mark.parametrize("failure", [AKTO_ERROR_MARK, AKTO_GARBAGE_MARK, AKTO_DROP_MARK]) +def test_failing_akto_under_logging_only_never_fails_the_caller(rig: Rig, failure: str) -> None: + marker: Final = _marker() + body: Final = _chat(rig, f"{failure} {marker}") + assert _content(body) == _answer(marker) + + calls: Final = rig.settled_akto_calls(str(body["id"]), marker, LOG_KEY) + assert [call.flags for call in calls] == [REQUEST_CHECK], calls + entries: Final = rig.guardrail_entries(str(body["id"]), "akto-log") + assert [entry["guardrail_mode"] for entry in entries] == ["logging_only"], entries + + +def test_provider_failure_under_logging_only_reaches_the_caller_and_checks_nothing(rig: Rig) -> None: + marker: Final = _marker() + response: Final = rig.proxy.client.post( + "/v1/chat/completions", + json={"model": rig.chat_model, "messages": [{"role": "user", "content": f"{PROVIDER_FAIL_MARK} {marker}"}]}, + headers=rig.auth(), + ) + assert response.status_code == 500, response.text + assert "synthetic provider failure" in response.text + + rows: Final = spend_rows(response.headers["x-litellm-call-id"]) + assert [row["status"] for row in rows] == ["failure"], rows + assert [request for request in rig.akto.drain() if marker.encode() in request.body] == [] + + +async def _burst_call(rig: Rig, kind: str, content: str) -> tuple[str, str]: + if kind == "messages": + claude: Final = anthropic.AsyncAnthropic(base_url=rig.base(), api_key=rig.proxy.key, max_retries=0) + message: Final = await claude.messages.create( + model=rig.claude_model, max_tokens=64, messages=[{"role": "user", "content": content}] + ) + assert message.content[0].type == "text" + return message.id, message.content[0].text + client: Final = openai.AsyncOpenAI(base_url=rig.base() + "/v1", api_key=rig.proxy.key, max_retries=0) + if kind == "responses": + result: Final = await client.responses.create(model=rig.responses_model, input=content) + return result.id, result.output_text + if kind == "stream": + stream: Final = await client.chat.completions.create( + model=rig.chat_model, messages=[{"role": "user", "content": content}], stream=True + ) + chunks: Final = [chunk async for chunk in stream] + return chunks[0].id, "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) + completion: Final = await client.chat.completions.create( + model=rig.chat_model, messages=[{"role": "user", "content": content}] + ) + return completion.id, completion.choices[0].message.content or "" + + +@pytest.mark.asyncio +async def test_burst_with_akto_failing_for_half_the_calls_answers_and_logs_every_call_once(rig: Rig) -> None: + kinds: Final = ("chat", "stream", "messages", "responses") * 6 + markers: Final = tuple(_marker() for _ in kinds) + failing: Final = frozenset(markers[::2]) + results: Final = await asyncio.gather( + *( + _burst_call(rig, kind, (AKTO_DROP_MARK + " " if marker in failing else "") + "burst " + marker) + for kind, marker in zip(kinds, markers, strict=True) + ) + ) + assert [answer for _, answer in results] == [_answer(marker) for marker in markers] + + for response_id, _ in results: + assert [row["status"] for row in spend_rows(response_id)] == ["success"], response_id + calls: Final = tuple( + call for call in (_akto_call(request) for request in rig.akto.drain()) if call.authorization == LOG_KEY + ) + observed: Final = { + marker: [call.flags for call in calls if marker in json.dumps(call.payload)] for marker in markers + } + expected: Final = { + marker: [REQUEST_CHECK] if marker in failing else [REQUEST_CHECK, RESPONSE_CHECK] for marker in markers + } + assert observed == expected + assert rig.proxy.client.get("/health/liveliness").status_code == 200 diff --git a/tests/integration/observability/test_arize_otel_v2_openinference_chaos.py b/tests/integration/observability/test_arize_otel_v2_openinference_chaos.py new file mode 100644 index 00000000000..9a86bb28577 --- /dev/null +++ b/tests/integration/observability/test_arize_otel_v2_openinference_chaos.py @@ -0,0 +1,403 @@ +from __future__ import annotations + +import threading +import uuid +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack +from pathlib import Path +from typing import Final +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +from _openinference_support import ( + CHAT_TOOLS, + RESPONSES_TOOLS, + _anthropic_response, + _anthropic_stream_response, + _assert_chat_request, + _assert_messages_request, + _assert_responses_request, + _chat_caller_response, + _chat_caller_stream, + _chat_request_marker, + _chat_response, + _chat_stream_response, + _chat_tool_call, + _collect_marker_spans, + _json_object, + _llm_spans_through_markers, + _matching_marker_span, + _messages_caller_raw_stream, + _messages_caller_response, + _normalize_chat_caller_stream, + _normalize_responses_caller_body, + _normalize_responses_caller_stream, + _owned_sink_handler, + _responses_caller_response, + _responses_caller_stream, + _responses_response, + _responses_stream_response, + _rig, + _spans, + _sse_json_values, +) +from integration._support.client import Gateway, eventually +from integration._support.wire import Reply, Request, wire_server + + +def _call( + proxy: Gateway, + model: str, + marker: str, + *, + surface: str = "chat", + stream: bool = False, + prompt: str | None = None, +) -> httpx.Response: + match surface: + case "chat": + return proxy.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt or "weather in Paris?"}], + "tools": CHAT_TOOLS, + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "metadata": {"trace_marker": marker}, + **({"stream": True} if stream else {}), + **({"stream_options": {"include_usage": True}} if stream and surface == "chat" else {}), + "cache": {"no-cache": True}, + }, + ) + case "responses": + return proxy.request( + "POST", + "/v1/responses", + { + "model": model, + "input": prompt or "weather in Paris?", + "tools": RESPONSES_TOOLS, + "tool_choice": {"type": "function", "name": "lookup_weather"}, + "metadata": {"trace_marker": marker}, + **({"stream": True} if stream else {}), + "cache": {"no-cache": True}, + }, + ) + case "messages": + return proxy.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 64, + "messages": [{"role": "user", "content": prompt or "weather in Paris?"}], + "tools": [ + { + "name": "lookup_weather", + "description": "Get weather", + "input_schema": { + "type": "object", + "properties": {"city": {"type": "string"}}, + }, + } + ], + "tool_choice": {"type": "auto"}, + "metadata": {"trace_marker": marker}, + **({"stream": True} if stream else {}), + "cache": {"no-cache": True}, + }, + ) + case _: + raise AssertionError(f"Unknown endpoint: {surface}") + + +def _assert_response(response: httpx.Response, marker: str, surface: str, stream: bool, model: str) -> None: + assert response.status_code == 200, response.text + if stream: + observed_stream: Final = _sse_json_values(response.content) + if surface == "chat": + chat_reply: Final = _chat_stream_response(marker, (_chat_tool_call(marker),)) + assert _normalize_chat_caller_stream(observed_stream) == _chat_caller_stream(chat_reply, model), ( + response.text + ) + return + if surface == "responses": + responses_reply: Final = _responses_stream_response(marker, (_chat_tool_call(marker),)) + assert _normalize_responses_caller_stream(observed_stream) == _responses_caller_stream( + responses_reply, model + ), response.text + return + if surface == "messages": + messages_reply: Final = _anthropic_stream_response(marker) + assert observed_stream == _messages_caller_raw_stream(messages_reply, model), response.text + return + raise AssertionError(f"Unknown endpoint: {surface}") + + observed: Final = _json_object(response.content) + if surface == "chat": + chat_reply: Final = _chat_response(marker) + assert observed == _chat_caller_response(chat_reply, model), response.text + return + if surface == "responses": + responses_reply: Final = _responses_response(marker) + assert _normalize_responses_caller_body(observed) == _responses_caller_response(responses_reply, model), ( + response.text + ) + return + if surface == "messages": + messages_reply: Final = _anthropic_response(marker) + assert observed == _messages_caller_response(messages_reply, model), response.text + return + raise AssertionError(f"Unknown endpoint: {surface}") + + +def _call_without_worker_error( + proxy: Gateway, model: str, marker: str, *, prompt: str | None = None +) -> httpx.Response | None: + try: + return _call(proxy, model, marker, prompt=prompt) + except httpx.HTTPError: + return None + + +def test_arize_otel_v2_f1_sink_outage_and_recovery(gateway: Gateway, tmp_path: Path) -> None: + outage_markers: Final = tuple("f1-outage-" + uuid.uuid4().hex for _ in range(30)) + recovery_markers: Final = tuple("f1-recovery-" + uuid.uuid4().hex for _ in range(30)) + surfaces: Final = ("chat", "responses", "messages") + surface_streams: Final = tuple( + (surfaces[index % len(surfaces)], index % 2 == 0) for index in range(len(outage_markers)) + ) + calls: Final = tuple( + (marker, surface, stream) for marker, (surface, stream) in zip(outage_markers, surface_streams, strict=True) + ) + recovery_calls: Final = tuple( + (marker, surface, stream) for marker, (surface, stream) in zip(recovery_markers, surface_streams, strict=True) + ) + + def upstream(request: Request) -> Reply: + body: Final = _json_object(request.body) + if request.target.endswith("/messages"): + messages: Final = body.get("messages") + assert isinstance(messages, list) and isinstance(messages[0], dict), body + marker: Final = messages[0].get("content") + assert isinstance(marker, str), body + _assert_messages_request( + request, + marker=marker, + prompt=marker, + stream=True if body.get("stream") is True else False, + ) + return _anthropic_stream_response(marker) if body.get("stream") is True else _anthropic_response(marker) + if request.target.endswith("/responses"): + marker: Final = body.get("input") + assert isinstance(marker, str), body + _assert_responses_request( + request, + marker=marker, + input_value=marker, + stream=body.get("stream") is True, + ) + return ( + _responses_stream_response(marker, (_chat_tool_call(marker),)) + if body.get("stream") is True + else _responses_response(marker) + ) + marker: Final = _chat_request_marker(request) + _assert_chat_request( + request, + messages=[{"role": "user", "content": marker}], + stream=True if body.get("stream") is True else None, + stream_options={"include_usage": True} if body.get("stream") is True else None, + ) + return ( + _chat_stream_response(marker, (_chat_tool_call(marker),)) + if body.get("stream") is True + else _chat_response(marker) + ) + + def sink(_request: Request) -> Reply: + return Reply(body=b"", content_type="application/x-protobuf") + + def run_burst( + burst: tuple[tuple[str, str, bool], ...], + proxy: Gateway, + chat_model: str, + messages_model: str, + ) -> None: + with ThreadPoolExecutor(max_workers=len(burst)) as executor: + futures: Final = tuple( + executor.submit( + _call, + proxy, + messages_model if surface == "messages" else chat_model, + marker, + surface=surface, + stream=stream, + prompt=marker, + ) + for marker, surface, stream in burst + ) + responses: Final = tuple( + (marker, surface, stream, future.result(timeout=60)) + for (marker, surface, stream), future in zip(burst, futures, strict=True) + ) + for marker, surface, stream, response in responses: + model_name: Final = messages_model if surface == "messages" else chat_model + _assert_response(response, marker, surface, stream, model_name) + + with ExitStack() as servers: + initial_stack: Final = servers.enter_context(ExitStack()) + stopped_destination: Final = initial_stack.enter_context(wire_server(_owned_sink_handler(sink))) + sink_port: Final = urlsplit(stopped_destination.url).port + assert sink_port is not None, stopped_destination.url + initial_stack.close() + with _rig( + gateway, + tmp_path, + upstream, + destination_wire=stopped_destination, + workers=1, + ) as rig: + with httpx.Client(trust_env=False) as client, pytest.raises(httpx.ConnectError): + client.get(stopped_destination.url + "/health", timeout=2) + with rig.proxy.scenario() as scenario: + messages_model: Final = scenario.model( + model="anthropic/claude-opus-5-5", + api_base=rig.provider.url, + ) + run_burst(calls, rig.proxy, rig.model, messages_model) + with httpx.Client(trust_env=False) as client, pytest.raises(httpx.ConnectError): + client.get(stopped_destination.url + "/health", timeout=2) + recovered_stack: Final = servers.enter_context(ExitStack()) + recovered_destination: Final = recovered_stack.enter_context( + wire_server(_owned_sink_handler(sink), port=sink_port) + ) + run_burst(recovery_calls, rig.proxy, rig.model, messages_model) + sentinel: Final = f"f1-sentinel-{uuid.uuid4().hex}" + sentinel_response: Final = _call( + rig.proxy, + rig.model, + sentinel, + surface="chat", + stream=False, + prompt=sentinel, + ) + _assert_response(sentinel_response, sentinel, "chat", False, rig.model) + spans: Final = _llm_spans_through_markers(recovered_destination, (*recovery_markers, sentinel)) + assert all( + sum(span.get("litellm.metadata.trace_marker") == marker for span in spans) == 1 + for marker in recovery_markers + ), spans + assert all( + sum(span.get("litellm.metadata.trace_marker") == marker for span in spans) <= 1 + for marker in outage_markers + ), spans + assert sum(span.get("litellm.metadata.trace_marker") == sentinel for span in spans) == 1, spans + + +def test_arize_otel_v2_f2_slow_sink_does_not_deadlock(gateway: Gateway, tmp_path: Path) -> None: + release: Final = threading.Event() + blocked: Final = threading.Event() + completed: Final = threading.Event() + marker: Final = "f2-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return _chat_response(marker) + + def sink(request: Request) -> Reply: + if any(attributes.get("litellm.metadata.trace_marker") == marker for attributes in _spans((request,))): + blocked.set() + assert release.wait(timeout=5), "slow sink was not released" + completed.set() + return Reply(body=b"", content_type="application/x-protobuf") + + with _rig(gateway, tmp_path, upstream, destination_handler=sink) as rig: + timer: Final = threading.Timer(2, release.set) + try: + response: Final = _call(rig.proxy, rig.model, marker) + _assert_response(response, marker, "chat", False, rig.model) + assert eventually(lambda: blocked.is_set(), bool, seconds=10) + timer.start() + assert eventually(lambda: completed.is_set(), bool, seconds=10) + timer.join(timeout=5) + assert not timer.is_alive(), "Slow sink timer did not finish" + finally: + release.set() + timer.cancel() + if timer.ident is not None: + timer.join(timeout=5) + _matching_marker_span(rig.destination, marker) + + +def test_arize_otel_v2_f3_one_proxy_worker_can_die(gateway: Gateway, tmp_path: Path) -> None: + markers: Final = tuple("f3-" + uuid.uuid4().hex for _ in range(8)) + release: Final = threading.Event() + + def upstream(request: Request) -> Reply: + request_marker: Final = _chat_request_marker(request) + _assert_chat_request( + request, + messages=[{"role": "user", "content": request_marker}], + ) + assert release.wait(timeout=20), "F3 upstream barrier was not released" + return _chat_response(request_marker) + + def sink(_request: Request) -> Reply: + return Reply(body=b"", content_type="application/x-protobuf") + + with _rig( + gateway, + tmp_path, + upstream, + destination_handler=sink, + fresh_client_connections=True, + ) as rig: + children: Final = psutil.Process(rig.owned.process.pid).children(recursive=True) + workers: Final = tuple( + child for child in children if child.is_running() and "resource_tracker" not in " ".join(child.cmdline()) + ) + assert len(workers) >= 2, tuple((worker.pid, worker.name()) for worker in workers) + with ThreadPoolExecutor(max_workers=len(markers)) as executor: + try: + futures: Final = tuple( + executor.submit( + _call_without_worker_error, + rig.proxy, + rig.model, + marker, + prompt=marker, + ) + for marker in markers + ) + observed: Final = eventually( + lambda: rig.provider.received.qsize(), + lambda count: count >= 2, + seconds=10, + ) + assert observed >= 2, observed + workers[0].kill() + assert eventually(lambda: not workers[0].is_running(), bool, seconds=10), workers[0] + release.set() + in_flight: Final = tuple( + (marker, future.result(timeout=60)) for marker, future in zip(markers, futures, strict=True) + ) + finally: + release.set() + survivor: Final = "f3-survivor-" + uuid.uuid4().hex + survivor_response: Final = _call(rig.proxy, rig.model, survivor, prompt=survivor) + _assert_response(survivor_response, survivor, "chat", False, rig.model) + served_responses: Final = tuple( + (marker, response) for marker, response in in_flight if response is not None and response.status_code == 200 + ) + for marker, response in served_responses: + _assert_response(response, marker, "chat", False, rig.model) + served: Final = tuple(marker for marker, _response in served_responses) + (survivor,) + collected: Final = _collect_marker_spans(rig.destination, served) + assert len(collected) == len(served), collected + exported: Final = tuple(span["litellm.metadata.trace_marker"] for span in collected) + assert len(exported) == len(served), exported + assert frozenset(exported) == frozenset(served), exported diff --git a/tests/integration/observability/test_arize_otel_v2_openinference_family.py b/tests/integration/observability/test_arize_otel_v2_openinference_family.py new file mode 100644 index 00000000000..242b7bfe8c1 --- /dev/null +++ b/tests/integration/observability/test_arize_otel_v2_openinference_family.py @@ -0,0 +1,368 @@ +from __future__ import annotations + +import uuid +from collections.abc import Mapping +from pathlib import Path +from types import MappingProxyType +from typing import Final + +import httpx +import pytest +from _openinference_support import ( + CHAT_TOOLS, + _assert_chat_request, + _chat_caller_response, + _chat_response, + _json_messages, + _json_object, + _matching_genai_marker_span, + _matching_marker_span, + _rig, +) +from integration._support.client import Gateway +from integration._support.wire import Reply, Request +from pydantic import JsonValue + +_GENAI_B3_KEYS_WITHOUT_BAGGAGE: Final = frozenset( + { + "gen_ai.input.messages", + "gen_ai.operation.name", + "gen_ai.output.messages", + "gen_ai.provider.name", + "gen_ai.request.model", + "gen_ai.response.finish_reasons", + "gen_ai.response.id", + "gen_ai.response.model", + "gen_ai.system", + "gen_ai.tool.0.description", + "gen_ai.tool.0.name", + "gen_ai.tool.0.parameters", + "gen_ai.usage.completion_tokens", + "gen_ai.usage.input_tokens", + "gen_ai.usage.output_tokens", + "gen_ai.usage.prompt_tokens", + "gen_ai.usage.total_tokens", + "litellm.api_key.hash", + "litellm.call_id", + "litellm.call_type", + "litellm.cost.discount_amount", + "litellm.cost.discount_percent", + "litellm.cost.input", + "litellm.cost.margin_fixed_amount", + "litellm.cost.margin_percent", + "litellm.cost.margin_total_amount", + "litellm.cost.original", + "litellm.cost.output", + "litellm.cost.tool_usage", + "litellm.cost.total", + "litellm.provider.model", + "litellm.request.route", + "litellm.request.tools.declared", + "llm.request.functions.0.description", + "llm.request.functions.0.name", + "llm.request.functions.0.parameters", + "server.address", + "server.port", + } +) +_GENAI_B3_KEYS_WITH_BAGGAGE: Final = _GENAI_B3_KEYS_WITHOUT_BAGGAGE | frozenset({"litellm.metadata.trace_marker"}) +_LANGFUSE_B3_KEYS: Final = frozenset( + { + "gen_ai.input.messages", + "gen_ai.operation.name", + "gen_ai.output.messages", + "gen_ai.provider.name", + "gen_ai.request.model", + "gen_ai.response.finish_reasons", + "gen_ai.response.id", + "gen_ai.response.model", + "gen_ai.system", + "gen_ai.tool.0.description", + "gen_ai.tool.0.name", + "gen_ai.tool.0.parameters", + "gen_ai.usage.completion_tokens", + "gen_ai.usage.input_tokens", + "gen_ai.usage.output_tokens", + "gen_ai.usage.prompt_tokens", + "gen_ai.usage.total_tokens", + "langfuse.observation.cost_details", + "langfuse.observation.id", + "langfuse.observation.input", + "langfuse.observation.metadata.provider", + "langfuse.observation.model.name", + "langfuse.observation.output", + "langfuse.observation.type", + "langfuse.observation.usage_details", + "litellm.api_key.hash", + "litellm.call_id", + "litellm.call_type", + "litellm.cost.discount_amount", + "litellm.cost.discount_percent", + "litellm.cost.input", + "litellm.cost.margin_fixed_amount", + "litellm.cost.margin_percent", + "litellm.cost.margin_total_amount", + "litellm.cost.original", + "litellm.cost.output", + "litellm.cost.tool_usage", + "litellm.cost.total", + "litellm.provider.model", + "litellm.request.route", + "litellm.request.tools.declared", + "llm.request.functions.0.description", + "llm.request.functions.0.name", + "llm.request.functions.0.parameters", + "server.address", + "server.port", + } +) +_B3_ATTRIBUTE_KEYS: Final[Mapping[str, frozenset[str]]] = MappingProxyType( + { + "langfuse_otel": _LANGFUSE_B3_KEYS, + "langtrace": _GENAI_B3_KEYS_WITHOUT_BAGGAGE, + "signoz": _GENAI_B3_KEYS_WITH_BAGGAGE, + "newrelic": _GENAI_B3_KEYS_WITH_BAGGAGE, + "levo": _GENAI_B3_KEYS_WITH_BAGGAGE, + "agentops": _GENAI_B3_KEYS_WITH_BAGGAGE, + "otel": _GENAI_B3_KEYS_WITH_BAGGAGE, + } +) +_B3_BAGGAGE_CALLBACKS: Final = frozenset({"signoz", "newrelic", "levo", "agentops", "otel"}) +_B4_ATTRIBUTE_KEYS: Final = frozenset( + { + "input.value", + "litellm.trace_id", + "llm.cost.total", + "llm.input_messages.0.message.content", + "llm.input_messages.0.message.role", + "llm.invocation_parameters", + "llm.is_streaming", + "llm.model_name", + "llm.output_messages.0.message.content", + "llm.output_messages.0.message.role", + "llm.output_messages.0.message.tool_calls.0.tool_call.function.arguments", + "llm.output_messages.0.message.tool_calls.0.tool_call.function.name", + "llm.output_messages.0.message.tool_calls.0.tool_call.id", + "llm.provider", + "llm.request.type", + "llm.response.cost", + "llm.response.id", + "llm.response.model", + "llm.token_count.completion", + "llm.token_count.prompt", + "llm.token_count.total", + "llm.tools.0.description", + "llm.tools.0.name", + "llm.tools.0.parameters", + "metadata", + "openinference.span.kind", + "output.value", + "user.id", + } +) + + +def _request(proxy: Gateway, model: str, marker: str) -> httpx.Response: + return proxy.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "weather in Paris?"}], + "tools": CHAT_TOOLS, + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "metadata": {"trace_marker": marker}, + "cache": {"no-cache": True}, + }, + ) + + +def _assert_openinference(attributes: dict[str, str], marker: str) -> None: + assert attributes["llm.output_messages.0.message.tool_calls.0.tool_call.id"] == f"call_{marker}", attributes + assert attributes["llm.output_messages.0.message.tool_calls.0.tool_call.function.name"] == "lookup_weather", ( + attributes + ) + assert ( + attributes["llm.output_messages.0.message.tool_calls.0.tool_call.function.arguments"] == '{"city": "Paris"}' + ), attributes + assert _json_object(attributes["metadata"].encode()) == {"trace_marker": marker}, attributes + assert attributes["litellm.metadata.trace_marker"] == marker, attributes + assert _json_messages(attributes["output.value"]) == [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": f"call_{marker}", + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + } + ], + } + ], attributes + + +def test_arize_otel_v2_b1_phoenix_openinference(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "b1-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return _chat_response(marker) + + with _rig( + gateway, + tmp_path, + upstream, + callbacks=("arize_phoenix",), + callback_settings={"otel": {"exporter": "http/protobuf", "endpoint": "unused"}}, + environment={"PHOENIX_PROJECT_NAME": "integration"}, + ) as rig: + response: Final = _request(rig.proxy, rig.model, marker) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_chat_response(marker), rig.model), response.text + _assert_openinference(_matching_marker_span(rig.destination, marker), marker) + + +def test_arize_otel_v2_b2_weave_openinference(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "b2-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return _chat_response(marker) + + with _rig( + gateway, + tmp_path, + upstream, + callbacks=("weave_otel",), + callback_settings={"otel": {"exporter": "http/protobuf", "endpoint": "unused"}}, + environment={"WANDB_BASE_URL": "http://127.0.0.1"}, + ) as rig: + response: Final = _request(rig.proxy, rig.model, marker) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_chat_response(marker), rig.model), response.text + _assert_openinference(_matching_marker_span(rig.destination, marker), marker) + + +@pytest.mark.parametrize( + "callback", + ("langfuse_otel", "langtrace", "signoz", "newrelic", "levo", "agentops", "otel"), +) +def test_arize_otel_v2_b3_non_openinference_callback_family(callback: str, gateway: Gateway, tmp_path: Path) -> None: + marker: Final = f"b3-{callback}-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + body: Final = _json_object(request.body) + assert body == { + "messages": [{"role": "user", "content": "weather in Paris?"}], + "model": "gpt-4o-mini", + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "tools": CHAT_TOOLS, + }, body + return _chat_response(marker) + + callback_settings: Final[dict[str, JsonValue]] = { + "otel": {"exporter": "http/protobuf", "endpoint": "unused", "mapper_names": ["genai"]} + } + with _rig( + gateway, + tmp_path, + upstream, + callbacks=(callback,), + callback_settings=callback_settings, + environment={ + "LANGFUSE_HOST": "http://127.0.0.1", + **( + { + "HTTPS_PROXY": "http://127.0.0.1:0", + "https_proxy": "http://127.0.0.1:0", + "NO_PROXY": "", + "no_proxy": "", + } + if callback == "agentops" + else {} + ), + }, + remove_environment=( + ("NEW_RELIC_LICENSE_KEY",) + if callback == "newrelic" + else ("AGENTOPS_API_KEY",) + if callback == "agentops" + else () + ), + ) as rig: + response: Final = _request(rig.proxy, rig.model, marker) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_chat_response(marker), rig.model), response.text + attributes: Final = _matching_genai_marker_span(rig.destination, marker) + assert frozenset(attributes) == _B3_ATTRIBUTE_KEYS[callback], attributes + assert "metadata" not in attributes, attributes + assert not any(".tool_calls." in key for key in attributes), attributes + baggage: Final = tuple( + sorted((key, value) for key, value in attributes.items() if key.startswith("litellm.metadata.")) + ) + expected_baggage: Final = ( + (("litellm.metadata.trace_marker", marker),) if callback in _B3_BAGGAGE_CALLBACKS else () + ) + assert baggage == expected_baggage, attributes + + +def test_arize_otel_v2_b3_langfuse_carries_metadata_baggage(gateway: Gateway, tmp_path: Path) -> None: + pytest.skip("BUG: LIT-9128 Langfuse OTel v2 preset omits request-metadata baggage") + marker: Final = "b3-langfuse-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + body: Final = _json_object(request.body) + assert body == { + "messages": [{"role": "user", "content": "weather in Paris?"}], + "model": "gpt-4o-mini", + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "tools": CHAT_TOOLS, + }, body + return _chat_response(marker) + + callback_settings: Final[dict[str, JsonValue]] = { + "otel": {"exporter": "http/protobuf", "endpoint": "unused", "mapper_names": ["genai"]} + } + with _rig( + gateway, + tmp_path, + upstream, + callbacks=("langfuse_otel",), + callback_settings=callback_settings, + environment={"LANGFUSE_HOST": "http://127.0.0.1"}, + ) as rig: + response: Final = _request(rig.proxy, rig.model, marker) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_chat_response(marker), rig.model), response.text + attributes: Final = _matching_genai_marker_span(rig.destination, marker) + assert attributes["litellm.metadata.trace_marker"] == marker, attributes + + +def test_arize_otel_v2_b4_legacy_otel_is_unchanged(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "b4-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + body: Final = _json_object(request.body) + assert body == { + "messages": [{"role": "user", "content": "weather in Paris?"}], + "model": "gpt-4o-mini", + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "tools": CHAT_TOOLS, + }, body + return _chat_response(marker) + + with _rig( + gateway, + tmp_path, + upstream, + callbacks=("arize",), + callback_settings={"otel": {"exporter": "http/protobuf", "endpoint": "unused"}}, + remove_environment=("LITELLM_OTEL_V2",), + disabled_environment=("LITELLM_OTEL_V2",), + ) as rig: + response: Final = _request(rig.proxy, rig.model, marker) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_chat_response(marker), rig.model), response.text + attributes: Final = _matching_marker_span(rig.destination, marker) + assert frozenset(attributes) == _B4_ATTRIBUTE_KEYS, attributes diff --git a/tests/integration/observability/test_arize_otel_v2_openinference_modes.py b/tests/integration/observability/test_arize_otel_v2_openinference_modes.py new file mode 100644 index 00000000000..7d0f3f31787 --- /dev/null +++ b/tests/integration/observability/test_arize_otel_v2_openinference_modes.py @@ -0,0 +1,484 @@ +from __future__ import annotations + +import uuid +from collections.abc import Callable +from pathlib import Path +from typing import Final + +import httpx +import pytest +from _openinference_support import ( + CHAT_TOOLS, + _assert_chat_request, + _chat_caller_response, + _chat_plain_response, + _chat_request_marker, + _chat_response, + _json_messages, + _json_object, + _llm_spans_through_markers, + _matching_marker_span, + _matching_output_value_span, + _matching_span, + _response_tool_calls, + _rig, +) +from integration._support.client import Gateway +from integration._support.wire import Reply, Request + +_DEFAULT_METADATA: Final = { + "requester_ip_address": "127.0.0.1", + "user_api_key_user_id": "default_user_id", +} +_DEFAULT_METADATA_BAGGAGE: Final = frozenset({("litellm.metadata.user_api_key_user_id", "default_user_id")}) + + +def _request( + proxy: Gateway, + model: str, + marker: str, + *, + prompt: str = "weather in Paris?", + headers: dict[str, str] | None = None, + key: str | None = None, +) -> httpx.Response: + return proxy.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt}], + "tools": CHAT_TOOLS, + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "metadata": {"trace_marker": marker}, + "cache": {"no-cache": True}, + }, + key=key, + headers=headers, + ) + + +def _assert_success_body(response: httpx.Response, marker: str, model: str) -> None: + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_chat_response(marker), model), response.text + + +def _upstream(marker: str) -> Callable[[Request], Reply]: + def reply(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return _chat_response(marker) + + return reply + + +def _assert_output_tool_call(attributes: dict[str, str], marker: str) -> None: + assert attributes["llm.output_messages.0.message.tool_calls.0.tool_call.id"] == f"call_{marker}", attributes + assert attributes["llm.output_messages.0.message.tool_calls.0.tool_call.function.name"] == "lookup_weather", ( + attributes + ) + assert ( + attributes["llm.output_messages.0.message.tool_calls.0.tool_call.function.arguments"] == '{"city": "Paris"}' + ), attributes + assert _json_messages(attributes["output.value"]) == [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": f"call_{marker}", + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + } + ], + } + ], attributes + + +def _assert_default_allowlist_attributes(attributes: dict[str, str], marker: str) -> None: + observed_baggage: Final = frozenset( + (key, value) for key, value in attributes.items() if key.startswith("litellm.metadata.") + ) + assert observed_baggage == _DEFAULT_METADATA_BAGGAGE, attributes + metadata: Final = _json_object(attributes["metadata"].encode()) + assert metadata == _DEFAULT_METADATA, attributes + assert "trace_marker" not in metadata, attributes + assert "litellm.metadata.trace_marker" not in attributes, attributes + _assert_output_tool_call(attributes, marker) + + +def test_arize_otel_v2_c1_absent_allowlist(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "c1-" + uuid.uuid4().hex + with _rig( + gateway, + tmp_path, + _upstream(marker), + remove_environment=("LITELLM_OTEL_BAGGAGE_METADATA_KEYS",), + disabled_environment=("LITELLM_OTEL_BAGGAGE_METADATA_KEYS",), + ) as rig: + response: Final = _request(rig.proxy, rig.model, marker) + _assert_success_body(response, marker, rig.model) + attributes: Final = _matching_span(rig.destination, marker) + _assert_default_allowlist_attributes(attributes, marker) + + +def test_arize_otel_v2_c2_empty_allowlist(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "c2-" + uuid.uuid4().hex + with _rig(gateway, tmp_path, _upstream(marker), environment={"LITELLM_OTEL_BAGGAGE_METADATA_KEYS": ""}) as rig: + response: Final = _request(rig.proxy, rig.model, marker) + _assert_success_body(response, marker, rig.model) + attributes: Final = _matching_span(rig.destination, marker) + assert "metadata" not in attributes, attributes + assert not any(key.startswith("litellm.metadata.") for key in attributes), attributes + _assert_output_tool_call(attributes, marker) + + +def test_arize_otel_v2_c3_absent_allowlisted_key(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "c3-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return _chat_response(marker) + + with _rig(gateway, tmp_path, upstream) as rig: + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + { + "model": rig.model, + "messages": [{"role": "user", "content": "weather in Paris?"}], + "tools": CHAT_TOOLS, + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "metadata": {"other": "value"}, + "cache": {"no-cache": True}, + }, + ) + _assert_success_body(response, marker, rig.model) + attributes: Final = _matching_span(rig.destination, marker) + assert "metadata" not in attributes, attributes + assert "litellm.metadata.trace_marker" not in attributes, attributes + _assert_output_tool_call(attributes, marker) + + +def test_arize_otel_v2_c4_promotes_marker_and_alias(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "c4-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return _chat_response(marker) + + with ( + _rig( + gateway, + tmp_path, + upstream, + environment={"LITELLM_OTEL_BAGGAGE_METADATA_KEYS": "requester_metadata.trace_marker,user_api_key_alias"}, + ) as rig, + rig.proxy.scenario() as scenario, + ): + key: Final = scenario.key(key_alias="alias-c4") + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + { + "model": rig.model, + "messages": [{"role": "user", "content": "weather in Paris?"}], + "tools": CHAT_TOOLS, + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "metadata": {"trace_marker": marker}, + "cache": {"no-cache": True}, + }, + key=key, + ) + _assert_success_body(response, marker, rig.model) + attributes: Final = _matching_marker_span(rig.destination, marker) + assert _json_object(attributes["metadata"].encode()) == { + "trace_marker": marker, + "user_api_key_alias": "alias-c4", + }, attributes + assert attributes["litellm.metadata.trace_marker"] == marker, attributes + assert attributes["litellm.metadata.user_api_key_alias"] == "alias-c4", attributes + + +def test_arize_otel_v2_c5_yaml_allowlist_does_not_reach_preset_so_default_applies( + gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = "c5-" + uuid.uuid4().hex + with _rig( + gateway, + tmp_path, + _upstream(marker), + callback_settings={"otel": {"baggage_metadata_keys": ["requester_metadata.trace_marker"]}}, + remove_environment=("LITELLM_OTEL_BAGGAGE_METADATA_KEYS",), + disabled_environment=("LITELLM_OTEL_BAGGAGE_METADATA_KEYS",), + ) as rig: + response: Final = _request(rig.proxy, rig.model, marker) + _assert_success_body(response, marker, rig.model) + attributes: Final = _matching_marker_span(rig.destination, marker) + _assert_default_allowlist_attributes(attributes, marker) + + +def test_arize_otel_v2_c5_yaml_allowlist_reaches_preset(gateway: Gateway, tmp_path: Path) -> None: + pytest.skip("BUG: LIT-9124 arize preset ignores callback_settings.otel.baggage_metadata_keys from config.yaml") + marker: Final = "c5-allowlist-" + uuid.uuid4().hex + with _rig( + gateway, + tmp_path, + _upstream(marker), + callback_settings={"otel": {"baggage_metadata_keys": ["requester_metadata.trace_marker"]}}, + remove_environment=("LITELLM_OTEL_BAGGAGE_METADATA_KEYS",), + disabled_environment=("LITELLM_OTEL_BAGGAGE_METADATA_KEYS",), + ) as rig: + response: Final = _request(rig.proxy, rig.model, marker) + _assert_success_body(response, marker, rig.model) + attributes: Final = _matching_marker_span(rig.destination, marker) + assert _json_object(attributes["metadata"].encode()) == {"trace_marker": marker}, attributes + assert attributes["litellm.metadata.trace_marker"] == marker, attributes + + +def test_arize_otel_v2_c6_content_capture_disabled(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "c6-" + uuid.uuid4().hex + with _rig( + gateway, + tmp_path, + _upstream(marker), + environment={ + "LITELLM_OTEL_BAGGAGE_METADATA_KEYS": "requester_metadata.trace_marker", + }, + remove_environment=("OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT",), + disabled_environment=("OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT",), + ) as rig: + response: Final = _request(rig.proxy, rig.model, marker) + _assert_success_body(response, marker, rig.model) + attributes: Final = _matching_marker_span(rig.destination, marker) + assert _json_object(attributes["metadata"].encode()) == {"trace_marker": marker}, attributes + assert attributes["litellm.metadata.trace_marker"] == marker, attributes + assert "gen_ai.input.messages" not in attributes, attributes + assert "gen_ai.output.messages" not in attributes, attributes + assert not any( + key.startswith("llm.input_messages.") or key.startswith("llm.output_messages.") for key in attributes + ), attributes + assert "input.value" not in attributes, attributes + assert "output.value" not in attributes, attributes + assert not any(".tool_calls." in key for key in attributes), attributes + + +def test_arize_otel_v2_c7_key_and_team_logging_callbacks(gateway: Gateway, tmp_path: Path) -> None: + key_marker: Final = "c7-key-" + uuid.uuid4().hex + team_marker: Final = "c7-team-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + marker: Final = _chat_request_marker(request) + assert marker in (key_marker, team_marker), request + _assert_chat_request(request, messages=[{"role": "user", "content": marker}]) + return _chat_response(marker) + + logging_metadata: Final = {"logging": [{"callback_name": "arize", "callback_type": "success"}]} + with _rig(gateway, tmp_path, upstream) as rig, rig.proxy.scenario() as scenario: + key: Final = scenario.key(key_alias="key-c7", metadata=logging_metadata) + team: Final = scenario.team(metadata=logging_metadata) + team_key: Final = scenario.key(team_id=team) + key_response: Final = _request(rig.proxy, rig.model, key_marker, prompt=key_marker, key=key) + _assert_success_body(key_response, key_marker, rig.model) + key_attributes: Final = _matching_marker_span(rig.destination, key_marker) + team_response: Final = _request(rig.proxy, rig.model, team_marker, prompt=team_marker, key=team_key) + _assert_success_body(team_response, team_marker, rig.model) + team_attributes: Final = _matching_marker_span(rig.destination, team_marker) + assert _json_object(key_attributes["metadata"].encode()) == {"trace_marker": key_marker}, key_attributes + assert _json_object(team_attributes["metadata"].encode()) == {"trace_marker": team_marker}, team_attributes + assert key_attributes["litellm.metadata.trace_marker"] == key_marker, key_attributes + assert team_attributes["litellm.metadata.trace_marker"] == team_marker, team_attributes + _assert_output_tool_call(key_attributes, key_marker) + _assert_output_tool_call(team_attributes, team_marker) + + +def test_arize_otel_v2_c8_request_callback_disable(gateway: Gateway, tmp_path: Path) -> None: + control_marker: Final = "c8-control-" + uuid.uuid4().hex + disabled_marker: Final = "c8-disabled-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + marker: Final = _chat_request_marker(request) + assert marker in (control_marker, disabled_marker), marker + return _chat_response(marker) + + with _rig( + gateway, + tmp_path, + upstream, + litellm_settings={"allow_dynamic_callback_disabling": True}, + ) as rig: + control_response: Final = _request(rig.proxy, rig.model, control_marker, prompt=control_marker) + _assert_success_body(control_response, control_marker, rig.model) + control_attributes: Final = _matching_marker_span(rig.destination, control_marker) + _assert_output_tool_call(control_attributes, control_marker) + + disabled_response: Final = _request( + rig.proxy, + rig.model, + disabled_marker, + prompt=disabled_marker, + headers={"x-litellm-disable-callbacks": "arize"}, + ) + _assert_success_body(disabled_response, disabled_marker, rig.model) + + +def test_arize_otel_v2_c8_disabled_callback_exports_no_span(gateway: Gateway, tmp_path: Path) -> None: + pytest.skip("BUG: LIT-9049 x-litellm-disable-callbacks: arize still exports the OTel v2 span") + disabled_marker: Final = "c8-disabled-" + uuid.uuid4().hex + sentinel: Final = "c8-sentinel-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + marker: Final = _chat_request_marker(request) + assert marker in (disabled_marker, sentinel), request + _assert_chat_request(request, messages=[{"role": "user", "content": marker}]) + return _chat_response(marker) + + with _rig( + gateway, + tmp_path, + upstream, + litellm_settings={"allow_dynamic_callback_disabling": True}, + workers=1, + ) as rig: + disabled_response: Final = _request( + rig.proxy, + rig.model, + disabled_marker, + prompt=disabled_marker, + headers={"x-litellm-disable-callbacks": "arize"}, + ) + _assert_success_body(disabled_response, disabled_marker, rig.model) + sentinel_response: Final = _request(rig.proxy, rig.model, sentinel, prompt=sentinel) + _assert_success_body(sentinel_response, sentinel, rig.model) + spans: Final = _llm_spans_through_markers(rig.destination, (sentinel,)) + assert not any(attributes.get("litellm.metadata.trace_marker") == disabled_marker for attributes in spans), ( + spans + ) + + +@pytest.mark.parametrize("failure_status", (401, 500)) +def test_arize_otel_v2_c9_upstream_failures_are_recorded(failure_status: int, gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "c9-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.body, f"{request.method} {request.target}" + body: Final = _json_object(request.body) + assert body == { + "messages": [{"role": "user", "content": "weather in Paris?"}], + "model": "gpt-4o-mini", + }, body + return Reply( + status=failure_status, + body=b'{"error":{"message":"upstream failure"}}', + content_type="application/json", + ) + + with _rig(gateway, tmp_path, upstream) as rig: + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + { + "model": rig.model, + "messages": [{"role": "user", "content": "weather in Paris?"}], + "metadata": {"trace_marker": marker, "failure_status": failure_status}, + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == failure_status, response.text + error_body: Final = _json_object(response.content) + error_name: Final = {401: "AuthenticationError", 500: "InternalServerError"}[failure_status] + error_type: Final = {401: "authentication_error", 500: "internal_server_error"}[failure_status] + provider_message: Final = f"litellm.{error_name}: {error_name}: OpenAIException - upstream failure" + caller_message: Final = ( + f"{provider_message}\n\nLiteLLM: model group '{rig.model}' failed with the error above. " + "No fallback was attempted." + ) + assert error_body == { + "error": { + "message": caller_message, + "type": error_type, + "param": None, + "code": str(failure_status), + } + }, error_body + attributes: Final = _matching_marker_span(rig.destination, marker) + metadata: Final = attributes.get("metadata") + assert metadata is not None, attributes + assert _json_object(metadata.encode()) == {"trace_marker": marker}, attributes + assert attributes["litellm.metadata.trace_marker"] == marker, attributes + assert attributes["error.message"] == provider_message, attributes + assert attributes["error.type"] == error_name, attributes + assert not any(".tool_calls." in key for key in attributes), attributes + + +def test_arize_otel_v2_c10_attribute_limit_keeps_tool_prefix(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "c10-" + uuid.uuid4().hex + calls: Final = _response_tool_calls(marker, ("Paris", "Berlin", "Rome", "Tokyo", "Oslo", "Lima", "Accra", "Delhi")) + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return _chat_response(marker, calls) + + with _rig(gateway, tmp_path, upstream, environment={"OTEL_SPAN_ATTRIBUTE_COUNT_LIMIT": "57"}) as rig: + response: Final = _request(rig.proxy, rig.model, marker) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_chat_response(marker, calls), rig.model), ( + response.text + ) + attributes: Final = _matching_output_value_span(rig.destination, marker) + indexes: Final = tuple( + int(key.split(".tool_calls.")[1].split(".")[0]) + for key in attributes + if ".tool_calls." in key and key.endswith(".tool_call.id") + ) + assert indexes == (0,), attributes + assert attributes["llm.output_messages.0.message.role"] == "assistant", attributes + fields: Final = ("id", "function.name", "function.arguments") + expected_tool_call_keys: Final = frozenset().union( + *( + frozenset(f"llm.output_messages.0.message.tool_calls.{index}.tool_call.{field}" for field in fields) + for index in indexes + ) + ) + observed_tool_call_keys: Final = frozenset(key for key in attributes if ".tool_calls." in key) + assert observed_tool_call_keys == expected_tool_call_keys, attributes + assert _json_messages(attributes["output.value"])[0]["tool_calls"] == calls, attributes + assert _json_object(attributes["metadata"].encode()) == {"trace_marker": marker}, attributes + assert attributes["litellm.metadata.trace_marker"] == marker, attributes + + history: Final = [{"role": "user", "content": f"history-{index}"} for index in range(40)] + history_marker: Final = marker + "-history" + + def history_upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=history, include_tools=False) + return _chat_plain_response(history_marker, "history retained") + + with _rig(gateway, tmp_path, history_upstream) as history_rig: + history_response: Final = history_rig.proxy.request( + "POST", + "/v1/chat/completions", + { + "model": history_rig.model, + "messages": history, + "metadata": {"trace_marker": history_marker}, + "cache": {"no-cache": True}, + }, + ) + assert history_response.status_code == 200, history_response.text + assert _json_object(history_response.content) == _chat_caller_response( + _chat_plain_response(history_marker, "history retained"), history_rig.model + ), history_response.text + history_attributes: Final = _matching_marker_span(history_rig.destination, history_marker) + assert _json_object(history_attributes["metadata"].encode()) == {"trace_marker": history_marker}, ( + history_attributes + ) + assert history_attributes["litellm.metadata.trace_marker"] == history_marker, history_attributes + assert tuple(history_attributes[f"llm.input_messages.{index}.message.role"] for index in range(40)) == tuple( + str(message["role"]) for message in history + ) + assert tuple(history_attributes[f"llm.input_messages.{index}.message.content"] for index in range(40)) == tuple( + str(message["content"]) for message in history + ) + assert _json_messages(history_attributes["output.value"]) == [ + {"role": "assistant", "content": "history retained"} + ], history_attributes diff --git a/tests/integration/observability/test_arize_otel_v2_openinference_sad_edge.py b/tests/integration/observability/test_arize_otel_v2_openinference_sad_edge.py new file mode 100644 index 00000000000..a05b3f0f045 --- /dev/null +++ b/tests/integration/observability/test_arize_otel_v2_openinference_sad_edge.py @@ -0,0 +1,547 @@ +from __future__ import annotations + +import json +import uuid +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Final + +import httpx +import pytest +from _openinference_support import ( + CHAT_TOOLS, + _assert_chat_request, + _chat_caller_response, + _chat_output_value, + _chat_request_marker, + _collect_marker_spans, + _json_messages, + _json_object, + _json_object_value, + _matching_marker_span, + _rig, + _span_attributes, + _spans, +) +from integration._support.client import Gateway +from integration._support.wire import Reply, Request +from pydantic import JsonValue + + +def _call( + proxy: Gateway, + model: str, + marker: str, + metadata: JsonValue | None = None, + key: str | None = None, + *, + prompt: str = "weather in Paris?", +) -> httpx.Response: + return proxy.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": prompt}], + "tools": CHAT_TOOLS, + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "metadata": metadata if metadata is not None else {"trace_marker": marker}, + "cache": {"no-cache": True}, + }, + key=key, + ) + + +def _success(marker: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": marker, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_" + marker, + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + } + ], + }, + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12}, + } + ).encode() + ) + + +def _chat_message(body: bytes) -> dict[str, JsonValue]: + response: Final = _json_object(body) + choices: Final = response["choices"] + assert isinstance(choices, list) and len(choices) == 1 and isinstance(choices[0], dict), response + message: Final = choices[0]["message"] + assert isinstance(message, dict), response + return message + + +@pytest.mark.parametrize( + ("value", "expected_trace"), + ( + (7, "7"), + (["one", 2], None), + ("", None), + ("x" * 5000, "x" * 5000), + ({"enabled": True}, None), + ), +) +def test_arize_otel_v2_d1_metadata_value_shapes( + value: JsonValue, expected_trace: str | None, gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = "d1-" + uuid.uuid4().hex + metadata: Final = {"trace_marker": value} + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return _success(marker) + + with _rig(gateway, tmp_path, upstream) as rig: + response: Final = _call(rig.proxy, rig.model, marker, metadata) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_success(marker), rig.model), response.text + attributes: Final = _matching_marker_span(rig.destination, marker) + if expected_trace is None: + assert "metadata" not in attributes, attributes + assert "litellm.metadata.trace_marker" not in attributes, attributes + else: + assert json.loads(attributes["metadata"]) == {"trace_marker": expected_trace}, attributes + assert attributes["litellm.metadata.trace_marker"] == expected_trace, attributes + + +def test_arize_otel_v2_d2_duplicate_json_metadata_keys(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "d2-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return _success(marker) + + with _rig(gateway, tmp_path, upstream) as rig: + request_body: Final = ( + '{"model":"' + + rig.model + + '","messages":' + + json.dumps([{"role": "user", "content": "weather in Paris?"}]) + + ',"tools":' + + json.dumps(CHAT_TOOLS) + + "," + + '"tool_choice":{"type":"function","function":{"name":"lookup_weather"}},' + + '"metadata":{"trace_marker":"' + + marker + + '","trace_marker":"' + + marker + + '"},"cache":{"no-cache":true}}' + ) + response: Final = rig.proxy.client.post( + "/v1/chat/completions", + content=request_body, + headers={ + "authorization": f"Bearer {rig.proxy.key}", + "content-type": "application/json", + }, + ) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_success(marker), rig.model), response.text + attributes: Final = _matching_marker_span(rig.destination, marker) + assert _json_object(attributes["metadata"].encode()) == {"trace_marker": marker}, attributes + assert attributes["litellm.metadata.trace_marker"] == marker, attributes + + +def test_arize_otel_v2_d4_unauthenticated_request_has_no_span(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "d4-" + uuid.uuid4().hex + with _rig(gateway, tmp_path, lambda _request: _success(marker)) as rig: + with httpx.Client(base_url=str(rig.proxy.client.base_url), trust_env=False) as client: + response: Final = client.post( + "/v1/chat/completions", + json={"model": rig.model, "messages": [{"role": "user", "content": marker}]}, + ) + assert response.status_code == 401, response.text + body: Final = _json_object(response.content) + assert body == { + "error": { + "message": "Authentication Error, No api key passed in.", + "type": "auth_error", + "param": "None", + "code": "401", + } + }, body + assert rig.provider.received.qsize() == 0 + spans: Final = tuple( + attributes + for attributes in _spans(rig.destination.drain()) + if attributes.get("openinference.span.kind") == "LLM" + ) + assert spans == (), "unauthenticated request exported an LLM span" + + +def test_arize_otel_v2_d5_unknown_model_leaves_proxy_ready(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "d5-" + uuid.uuid4().hex + unknown_model: Final = "unknown-model-" + uuid.uuid4().hex + with _rig(gateway, tmp_path, lambda _request: _success(marker)) as rig: + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": unknown_model, "messages": [{"role": "user", "content": marker}]}, + ) + assert response.status_code == 400, response.text + body: Final = _json_object(response.content) + error_message: Final = ( + f"/chat/completions: Invalid model name passed in model={unknown_model}. " + "Call `/v1/models` to view available models for your key." + ) + assert body == { + "error": { + "message": error_message, + "type": "invalid_request_error", + "param": None, + "code": "400", + "provider_specific_fields": {"error": error_message}, + } + }, body + assert rig.provider.received.qsize() == 0 + readiness: Final = rig.proxy.client.get("/health/readiness") + assert readiness.status_code == 200, readiness.text + assert _json_object(readiness.content) == {"status": "healthy", "db": "connected"}, readiness.text + spans: Final = tuple( + attributes + for attributes in _spans(rig.destination.drain()) + if attributes.get("openinference.span.kind") == "LLM" + ) + assert spans == (), "unknown model exported an LLM span" + + +def test_arize_otel_v2_d6_sink_rejections_do_not_change_caller_response(gateway: Gateway, tmp_path: Path) -> None: + markers: Final = tuple(f"d6-{status}-" + uuid.uuid4().hex for status in (403, 404)) + unrelated_marker: Final = "d6-unrelated-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + response_marker: Final = _chat_request_marker(request) + _assert_chat_request(request, messages=[{"role": "user", "content": response_marker}]) + return _success(response_marker) + + def sink(request: Request) -> Reply: + exported_markers: Final = tuple( + attributes.get("litellm.metadata.trace_marker") + for attributes in _span_attributes(request) + if "litellm.metadata.trace_marker" in attributes + ) + if markers[0] in exported_markers: + return Reply(status=403, body=b"rejected") + if markers[1] in exported_markers: + return Reply(status=404, body=b"rejected") + return Reply(body=b"", content_type="application/x-protobuf") + + with ( + _rig( + gateway, + tmp_path, + upstream, + destination_handler=sink, + ) as rig, + rig.proxy.scenario() as scenario, + ): + unrelated_key: Final = scenario.key(key_alias="unrelated-d6") + first: Final = _call(rig.proxy, rig.model, markers[0], prompt=markers[0]) + assert first.status_code == 200, first.text + assert _json_object(first.content) == _chat_caller_response(_success(markers[0]), rig.model), first.text + assert _matching_marker_span(rig.destination, markers[0])["litellm.metadata.trace_marker"] == markers[0] + second: Final = _call(rig.proxy, rig.model, markers[1], prompt=markers[1]) + assert second.status_code == 200, second.text + assert _json_object(second.content) == _chat_caller_response(_success(markers[1]), rig.model), second.text + assert _matching_marker_span(rig.destination, markers[1])["litellm.metadata.trace_marker"] == markers[1] + unrelated: Final = _call( + rig.proxy, + rig.model, + unrelated_marker, + key=unrelated_key, + prompt=unrelated_marker, + ) + assert unrelated.status_code == 200, unrelated.text + assert _json_object(unrelated.content) == _chat_caller_response(_success(unrelated_marker), rig.model), ( + unrelated.text + ) + assert ( + _matching_marker_span(rig.destination, unrelated_marker)["litellm.metadata.trace_marker"] + == unrelated_marker + ) + + +def test_arize_otel_v2_d7_missing_space_id_is_stable(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "d7-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return _success(marker) + + with _rig( + gateway, + tmp_path, + upstream, + remove_environment=("ARIZE_SPACE_ID",), + disabled_environment=("ARIZE_SPACE_ID",), + ) as rig: + response: Final = _call(rig.proxy, rig.model, marker) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_success(marker), rig.model), response.text + assert _matching_marker_span(rig.destination, marker)["litellm.metadata.trace_marker"] == marker + + +def test_arize_otel_v2_e1_uncached_request_exports_one_llm_span(gateway: Gateway, tmp_path: Path) -> None: + markers: Final = tuple("e1-" + uuid.uuid4().hex for _ in range(3)) + + def upstream(request: Request) -> Reply: + marker: Final = _chat_request_marker(request) + assert marker in markers, request + _assert_chat_request(request, messages=[{"role": "user", "content": marker}]) + return _success(marker) + + with _rig(gateway, tmp_path, upstream) as rig: + responses: Final = tuple(_call(rig.proxy, rig.model, marker, prompt=marker) for marker in markers) + assert all(response.status_code == 200 for response in responses), tuple( + response.text for response in responses + ) + assert tuple(_json_object(response.content) for response in responses) == tuple( + _chat_caller_response(_success(marker), rig.model) for marker in markers + ), responses + spans: Final = _collect_marker_spans(rig.destination, markers) + assert len(spans) == len(markers), spans + spans_by_id: Final = {span["gen_ai.response.id"]: span for span in spans} + assert frozenset(spans_by_id) == frozenset(markers), spans + assert all( + spans_by_id[marker]["gen_ai.response.id"] in response.text + for marker, response in zip(markers, responses, strict=True) + ), spans + + +def test_arize_otel_v2_e2_concurrent_unique_markers(gateway: Gateway, tmp_path: Path) -> None: + markers: Final = tuple("e2-" + uuid.uuid4().hex for _ in range(20)) + + def upstream(request: Request) -> Reply: + marker: Final = _chat_request_marker(request) + assert marker.startswith("e2-"), request + _assert_chat_request(request, messages=[{"role": "user", "content": marker}]) + return _success(marker) + + with _rig(gateway, tmp_path, upstream) as rig: + with ThreadPoolExecutor(max_workers=20) as executor: + responses: Final = tuple( + executor.map( + lambda marker: _call(rig.proxy, rig.model, marker, prompt=marker), + markers, + ) + ) + assert all(response.status_code == 200 for response in responses), tuple( + response.text for response in responses + ) + assert tuple(_json_object(response.content) for response in responses) == tuple( + _chat_caller_response(_success(marker), rig.model) for marker in markers + ), responses + spans: Final = _collect_marker_spans(rig.destination, markers) + assert len(spans) == len(markers), spans + assert tuple( + _json_object(next(span for span in spans if marker in span.values())["metadata"].encode()) + for marker in markers + ) == tuple({"trace_marker": marker} for marker in markers), spans + assert ( + tuple( + next(span for span in spans if marker in span.values())["litellm.metadata.trace_marker"] + for marker in markers + ) + == markers + ), spans + + +@pytest.mark.parametrize( + "shape", + ("object-arguments", "missing-name", "non-dict-call", "null-tool-calls", "integer-id"), +) +def test_arize_otel_v2_d3_malformed_tool_calls_are_normalized(shape: str, gateway: Gateway, tmp_path: Path) -> None: + marker: Final = f"d3-{shape}-" + uuid.uuid4().hex + + def malformed_reply() -> Reply: + call: Final = { + "id": 17 if shape == "integer-id" else f"call_{marker}", + "type": "function", + "function": { + **({} if shape == "missing-name" else {"name": "lookup_weather"}), + "arguments": {"city": "Paris"} if shape == "object-arguments" else '{"city": "Paris"}', + }, + } + tool_calls: Final[JsonValue] = ( + None if shape == "null-tool-calls" else ["not-a-call"] if shape == "non-dict-call" else [call] + ) + message: Final = { + "role": "assistant", + "content": None, + "tool_calls": tool_calls, + } + return Reply( + body=json.dumps( + { + "id": marker, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "finish_reason": "tool_calls", "message": message}], + "usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12}, + } + ).encode() + ) + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return malformed_reply() + + with _rig(gateway, tmp_path, upstream) as rig: + proxy_log: Final = rig.owned.log + response: Final = _call(rig.proxy, rig.model, marker) + if shape == "non-dict-call": + assert response.status_code == 400, response.text + error_body: Final = _json_object(response.content) + assert frozenset(error_body) == frozenset({"error"}), error_body + error: Final = _json_object_value(error_body["error"]) + assert frozenset(error) == frozenset({"type", "code", "param", "message"}), error + assert error["type"] == "invalid_request_error", error + assert error["code"] == "400", error + assert error["param"] is None, error + message: Final = error["message"] + assert isinstance(message, str), error + assert "AttributeError: 'str' object has no attribute 'get'" in message, error + spans: Final = tuple(_spans(rig.destination.drain())) + assert all(not any(".tool_calls." in key for key in attributes) for attributes in spans), spans + readiness: Final = rig.proxy.client.get("/health/readiness") + assert readiness.status_code == 200, readiness.text + assert _json_object(readiness.content) == {"status": "healthy", "db": "connected"}, readiness.text + else: + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(malformed_reply(), rig.model), response.text + attributes: Final = _matching_marker_span(rig.destination, marker) + indexed_prefix: Final = "llm.output_messages.0.message.tool_calls.0.tool_call." + expected_fields: Final = ( + ("id", "function.name", "function.arguments"), + ("id", "function.arguments"), + (), + ("function.name", "function.arguments"), + )[("object-arguments", "missing-name", "null-tool-calls", "integer-id").index(shape)] + expected_keys: Final = frozenset(indexed_prefix + field for field in expected_fields) + observed_keys: Final = frozenset(key for key in attributes if ".tool_calls." in key) + assert observed_keys == expected_keys, attributes + if shape == "object-arguments": + assert attributes["llm.output_messages.0.message.tool_calls.0.tool_call.id"] == f"call_{marker}", ( + attributes + ) + assert ( + attributes["llm.output_messages.0.message.tool_calls.0.tool_call.function.name"] == "lookup_weather" + ), attributes + assert ( + attributes["llm.output_messages.0.message.tool_calls.0.tool_call.function.arguments"] + == '{"city": "Paris"}' + ), attributes + elif shape == "missing-name": + assert attributes["llm.output_messages.0.message.tool_calls.0.tool_call.id"] == f"call_{marker}", ( + attributes + ) + assert ( + attributes["llm.output_messages.0.message.tool_calls.0.tool_call.function.arguments"] + == '{"city": "Paris"}' + ), attributes + elif shape == "integer-id": + assert ( + attributes["llm.output_messages.0.message.tool_calls.0.tool_call.function.name"] == "lookup_weather" + ), attributes + assert ( + attributes["llm.output_messages.0.message.tool_calls.0.tool_call.function.arguments"] + == '{"city": "Paris"}' + ), attributes + assert _json_messages(attributes["output.value"]) == _json_messages( + _chat_output_value(malformed_reply()) + ), attributes + assert "Exception while exporting Span batch" not in proxy_log.read_text(), proxy_log.read_text() + + +def test_arize_otel_v2_d3_non_dict_tool_call_is_not_a_caller_error(gateway: Gateway, tmp_path: Path) -> None: + pytest.skip("BUG: LIT-9125 non-dict tool_calls entry returns HTTP 400 with a server traceback") + marker: Final = "d3-non-dict-call-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return Reply( + body=json.dumps( + { + "id": marker, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "finish_reason": "tool_calls", + "message": {"role": "assistant", "content": None, "tool_calls": ["not-a-call"]}, + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12}, + } + ).encode() + ) + + with _rig(gateway, tmp_path, upstream) as rig: + response: Final = _call(rig.proxy, rig.model, marker) + assert response.status_code not in range(400, 500), response.text + assert "Traceback" not in response.text, response.text + assert "AttributeError" not in response.text, response.text + readiness: Final = rig.proxy.client.get("/health/readiness") + assert readiness.status_code == 200, readiness.text + assert _json_object(readiness.content) == {"status": "healthy", "db": "connected"}, readiness.text + + +@pytest.mark.parametrize("shape", ("empty", "null", "missing")) +def test_arize_otel_v2_e3_empty_or_missing_tool_calls_never_indexed( + shape: str, gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = f"e3-{shape}-" + uuid.uuid4().hex + message: Final = ( + {"role": "assistant", "content": None, "tool_calls": []} + if shape == "empty" + else {"role": "assistant", "content": None, "tool_calls": None} + if shape == "null" + else {"role": "assistant", "content": None} + ) + expected_response: Final[dict[str, JsonValue]] = { + "id": marker, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "finish_reason": "stop", "message": message}], + "usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12}, + } + + expected_reply: Final = Reply(body=json.dumps(expected_response).encode()) + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return expected_reply + + with _rig(gateway, tmp_path, upstream) as rig: + response: Final = _call(rig.proxy, rig.model, marker) + assert response.status_code == 200, response.text + response_body: Final = _json_object(response.content) + assert response_body == _chat_caller_response(expected_reply, rig.model), response.text + attributes: Final = _matching_marker_span(rig.destination, marker) + assert not any(".tool_calls." in key for key in attributes), attributes + assert "tool_calls" not in _json_messages(attributes["output.value"])[0], attributes diff --git a/tests/integration/observability/test_arize_otel_v2_openinference_spans.py b/tests/integration/observability/test_arize_otel_v2_openinference_spans.py new file mode 100644 index 00000000000..d5b4dbeacea --- /dev/null +++ b/tests/integration/observability/test_arize_otel_v2_openinference_spans.py @@ -0,0 +1,966 @@ +from __future__ import annotations + +import asyncio +import json +import uuid +from itertools import chain +from pathlib import Path +from typing import Final + +import anthropic +import httpx +import openai +import pytest +from _openinference_support import ( + CHAT_TOOLS, + RESPONSES_TOOLS, + Rig, + _anthropic_stream_response, + _assert_chat_request, + _assert_messages_request, + _assert_responses_request, + _assert_tool_span, + _chat_cache_hit_caller_stream, + _chat_caller_response, + _chat_caller_stream, + _chat_plain_response, + _chat_request_marker, + _chat_response, + _chat_stream_response, + _chat_tool_call, + _json_messages, + _json_object, + _llm_spans_through_markers, + _matching_marker_span, + _messages_caller_response, + _messages_caller_stream, + _messages_caller_stream_response, + _normalize_chat_caller_stream, + _normalize_responses_caller_body, + _normalize_responses_caller_stream, + _response_tool_calls, + _responses_caller_response, + _responses_caller_stream, + _responses_response, + _responses_stream_response, + _rig, +) +from integration._support.client import Gateway, object_value, string_value +from integration._support.wire import Reply, Request +from pydantic import JsonValue + + +def _chat_upstream(request: Request) -> None: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + + +def _responses_upstream(request: Request, marker: str, *, stream: bool = False) -> None: + _assert_responses_request(request, marker=marker, stream=stream) + + +def _messages_upstream(request: Request, marker: str, *, stream: bool = False) -> None: + _assert_messages_request(request, marker=marker, stream=stream) + + +def _messages_response(identity: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": identity, + "type": "message", + "role": "assistant", + "model": "claude-opus-5-5", + "content": [ + { + "type": "tool_use", + "id": "call_" + identity, + "name": "lookup_weather", + "input": {"city": "Paris"}, + } + ], + "stop_reason": "tool_use", + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 4}, + } + ).encode() + ) + + +def _assert_tool_span_for_marker(attributes: dict[str, str], marker: str, *, content: str | None = None) -> None: + _assert_tool_span( + attributes, + marker=marker, + output=[ + { + "role": "assistant", + "content": content, + "tool_calls": [ + { + "id": "call_" + marker, + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + } + ], + } + ], + calls=[("call_" + marker, "lookup_weather", {"city": "Paris"})], + metadata={"trace_marker": marker}, + baggage={"trace_marker": marker}, + ) + + +def _chat_request( + proxy: Gateway, model: str, marker: str, *, stream: bool = False, no_cache: bool = True +) -> httpx.Response: + return proxy.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "weather in Paris?"}], + "tools": CHAT_TOOLS, + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "metadata": {"trace_marker": marker}, + **({"stream": True} if stream else {}), + **({"cache": {"no-cache": True}} if no_cache else {}), + }, + ) + + +def _responses_request(proxy: Gateway, model: str, marker: str) -> httpx.Response: + return proxy.request( + "POST", + "/v1/responses", + { + "model": model, + "input": "weather in Paris?", + "tools": RESPONSES_TOOLS, + "tool_choice": {"type": "function", "name": "lookup_weather"}, + "metadata": {"trace_marker": marker}, + "cache": {"no-cache": True}, + }, + ) + + +def _messages_request(proxy: Gateway, model: str, marker: str) -> httpx.Response: + return proxy.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 64, + "messages": [{"role": "user", "content": "weather in Paris?"}], + "tools": [ + { + "name": "lookup_weather", + "description": "Get weather", + "input_schema": { + "type": "object", + "properties": {"city": {"type": "string"}}, + }, + } + ], + "tool_choice": {"type": "auto"}, + "metadata": {"trace_marker": marker}, + "cache": {"no-cache": True}, + }, + ) + + +def _openai_client(proxy: Gateway) -> openai.OpenAI: + return openai.OpenAI(base_url=str(proxy.client.base_url) + "/v1", api_key=proxy.key, max_retries=0) + + +def _async_openai_client(proxy: Gateway) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI(base_url=str(proxy.client.base_url) + "/v1", api_key=proxy.key, max_retries=0) + + +def test_arize_otel_v2_a1_chat_sync_sdk(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a1-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _chat_upstream(request) + body: Final = _json_object(request.body) + assert body == { + "messages": [{"role": "user", "content": "weather in Paris?"}], + "model": "gpt-4o-mini", + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "tools": CHAT_TOOLS, + }, body + return _chat_response(marker) + + with _rig(gateway, tmp_path, upstream) as rig: + client: Final = _openai_client(rig.proxy) + response: Final = client.chat.completions.create( + model=rig.model, + messages=[{"role": "user", "content": "weather in Paris?"}], + tools=CHAT_TOOLS, + tool_choice={"type": "function", "function": {"name": "lookup_weather"}}, + extra_body={"metadata": {"trace_marker": marker}, "cache": {"no-cache": True}}, + ) + assert response.id == marker, response + assert response.model_dump(mode="json", exclude_unset=True) == _chat_caller_response( + _chat_response(marker), rig.model + ), response + calls: Final = response.choices[0].message.tool_calls + assert calls is not None and len(calls) == 1, response + assert (calls[0].id, calls[0].function.name, calls[0].function.arguments) == ( + f"call_{marker}", + "lookup_weather", + '{"city": "Paris"}', + ), response + _assert_tool_span_for_marker(_matching_marker_span(rig.destination, marker), marker) + + +def test_arize_otel_v2_a2_chat_async_sdk(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a2-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _chat_upstream(request) + return _chat_response(marker) + + async def call() -> None: + with _rig(gateway, tmp_path, upstream) as rig: + client: Final = _async_openai_client(rig.proxy) + response: Final = await client.chat.completions.create( + model=rig.model, + messages=[{"role": "user", "content": "weather in Paris?"}], + tools=CHAT_TOOLS, + tool_choice={"type": "function", "function": {"name": "lookup_weather"}}, + extra_body={"metadata": {"trace_marker": marker}, "cache": {"no-cache": True}}, + ) + assert response.id == marker, response + assert response.model_dump(mode="json", exclude_unset=True) == _chat_caller_response( + _chat_response(marker), rig.model + ), response + calls: Final = response.choices[0].message.tool_calls + assert calls is not None and len(calls) == 1, response + assert (calls[0].id, calls[0].function.name, calls[0].function.arguments) == ( + f"call_{marker}", + "lookup_weather", + '{"city": "Paris"}', + ), response + _assert_tool_span_for_marker(_matching_marker_span(rig.destination, marker), marker) + + asyncio.run(call()) + + +def test_arize_otel_v2_a5_responses_sync_sdk(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a5-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _responses_upstream(request, marker) + return _responses_response(marker) + + with _rig(gateway, tmp_path, upstream) as rig: + client: Final = _openai_client(rig.proxy) + response: Final = client.responses.create( + model=rig.model, + input="weather in Paris?", + tools=RESPONSES_TOOLS, + tool_choice={"type": "function", "name": "lookup_weather"}, + extra_body={"metadata": {"trace_marker": marker}, "cache": {"no-cache": True}}, + ) + assert response.id.startswith("resp_"), response + assert _normalize_responses_caller_body(response.model_dump(mode="json", exclude_unset=True)) == ( + _responses_caller_response(_responses_response(marker), rig.model) + ), response + assert (response.status, response.model) == ("completed", rig.model), response + assert ( + response.output[0].type, + response.output[0].call_id, + response.output[0].name, + response.output[0].arguments, + ) == ("function_call", f"call_{marker}", "lookup_weather", '{"city": "Paris"}'), response + _assert_tool_span_for_marker(_matching_marker_span(rig.destination, marker), marker) + + +def test_arize_otel_v2_a6_responses_async_streaming_sdk(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a6-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _responses_upstream(request, marker, stream=True) + return _responses_stream_response(marker, (_chat_tool_call(marker),)) + + async def call() -> None: + with _rig(gateway, tmp_path, upstream) as rig: + client: Final = _async_openai_client(rig.proxy) + stream: Final = await client.responses.create( + model=rig.model, + input="weather in Paris?", + tools=RESPONSES_TOOLS, + tool_choice={"type": "function", "name": "lookup_weather"}, + stream=True, + extra_body={"metadata": {"trace_marker": marker}, "cache": {"no-cache": True}}, + ) + events: Final = tuple([event async for event in stream]) + assert events[-1].type == "response.completed", events + assert events[-1].response.id.startswith("resp_"), events[-1] + assert _normalize_responses_caller_stream( + tuple(event.model_dump(mode="json", exclude_unset=True) for event in events) + ) == _responses_caller_stream(_responses_stream_response(marker, (_chat_tool_call(marker),)), rig.model), ( + events + ) + call: Final = events[-1].response.output[0] + assert (call.type, call.call_id, call.name, call.arguments) == ( + "function_call", + f"call_{marker}", + "lookup_weather", + '{"city": "Paris"}', + ), events[-1] + _assert_tool_span_for_marker(_matching_marker_span(rig.destination, marker), marker) + + asyncio.run(call()) + + +def test_arize_otel_v2_a7_messages_sync_sdk(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a7-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _messages_upstream(request, marker) + return _messages_response(marker) + + with _rig(gateway, tmp_path, upstream, model_name="anthropic/claude-opus-5-5", api_base_suffix="") as rig: + client: Final = anthropic.Anthropic( + base_url=str(rig.proxy.client.base_url), api_key=rig.proxy.key, max_retries=0 + ) + response: Final = client.messages.create( + model=rig.model, + max_tokens=64, + messages=[{"role": "user", "content": "weather in Paris?"}], + tools=[ + { + "name": "lookup_weather", + "description": "Get weather", + "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}}, + } + ], + tool_choice={"type": "auto"}, + metadata={"trace_marker": marker}, + ) + assert response.id == marker, response + assert response.model_dump(mode="json", exclude_unset=True) == _messages_caller_response( + _messages_response(marker), rig.model + ), response + call: Final = response.content[0] + assert (call.type, call.id, call.name, call.input) == ( + "tool_use", + f"call_{marker}", + "lookup_weather", + {"city": "Paris"}, + ), response + _assert_tool_span_for_marker(_matching_marker_span(rig.destination, marker), marker) + + +def test_arize_otel_v2_a3_chat_streaming(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a3-" + uuid.uuid4().hex + stream_reply: Final = _chat_stream_response(marker, (_chat_tool_call(marker),), include_usage=False) + + def upstream(request: Request) -> Reply: + _assert_chat_request( + request, + messages=[{"role": "user", "content": "weather in Paris?"}], + stream=True, + stream_options={"include_usage": False}, + ) + return stream_reply + + with _rig(gateway, tmp_path, upstream, general_settings={"always_include_stream_usage": False}) as rig: + client: Final = _openai_client(rig.proxy) + stream: Final = client.chat.completions.create( + model=rig.model, + messages=[{"role": "user", "content": "weather in Paris?"}], + tools=CHAT_TOOLS, + tool_choice={"type": "function", "function": {"name": "lookup_weather"}}, + stream=True, + stream_options={"include_usage": False}, + extra_body={"metadata": {"trace_marker": marker}, "cache": {"no-cache": True}}, + ) + chunks: Final = tuple(stream) + assert chunks[0].id == marker and chunks[-1].id == marker, chunks + assert _normalize_chat_caller_stream( + tuple(chunk.model_dump(mode="json", exclude_unset=True) for chunk in chunks) + ) == _chat_caller_stream(stream_reply, rig.model), chunks + assert chunks[-1].choices[0].finish_reason == "tool_calls", chunks + tool_call_deltas: Final = tuple( + chain.from_iterable(chunk.choices[0].delta.tool_calls or () for chunk in chunks) + ) + assert len(tool_call_deltas) == 2, chunks + assert ( + tool_call_deltas[0].id, + tool_call_deltas[0].function.name, + tool_call_deltas[1].function.arguments, + ) == (f"call_{marker}", "lookup_weather", '{"city": "Paris"}'), chunks + _assert_tool_span_for_marker(_matching_marker_span(rig.destination, marker), marker) + + +def test_arize_otel_v2_a4_chat_async_streaming(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a4-" + uuid.uuid4().hex + stream_reply: Final = _chat_stream_response(marker, (_chat_tool_call(marker),), include_usage=False) + + def upstream(request: Request) -> Reply: + _assert_chat_request( + request, + messages=[{"role": "user", "content": "weather in Paris?"}], + stream=True, + stream_options={"include_usage": False}, + ) + return stream_reply + + async def call() -> None: + with _rig(gateway, tmp_path, upstream, general_settings={"always_include_stream_usage": False}) as rig: + client: Final = _async_openai_client(rig.proxy) + stream: Final = await client.chat.completions.create( + model=rig.model, + messages=[{"role": "user", "content": "weather in Paris?"}], + tools=CHAT_TOOLS, + tool_choice={"type": "function", "function": {"name": "lookup_weather"}}, + stream=True, + stream_options={"include_usage": False}, + extra_body={"metadata": {"trace_marker": marker}, "cache": {"no-cache": True}}, + ) + chunks: Final = tuple([chunk async for chunk in stream]) + assert chunks[0].id == marker and chunks[-1].id == marker, chunks + assert _normalize_chat_caller_stream( + tuple(chunk.model_dump(mode="json", exclude_unset=True) for chunk in chunks) + ) == _chat_caller_stream(stream_reply, rig.model), chunks + assert chunks[-1].choices[0].finish_reason == "tool_calls", chunks + tool_call_deltas: Final = tuple( + chain.from_iterable(chunk.choices[0].delta.tool_calls or () for chunk in chunks) + ) + assert len(tool_call_deltas) == 2, chunks + assert ( + tool_call_deltas[0].id, + tool_call_deltas[0].function.name, + tool_call_deltas[1].function.arguments, + ) == (f"call_{marker}", "lookup_weather", '{"city": "Paris"}'), chunks + _assert_tool_span_for_marker(_matching_marker_span(rig.destination, marker), marker) + + asyncio.run(call()) + + +def test_arize_otel_v2_a8_messages_streaming(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a8-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _messages_upstream(request, marker, stream=True) + return _anthropic_stream_response(marker) + + async def call() -> None: + with _rig(gateway, tmp_path, upstream, model_name="anthropic/claude-opus-5-5", api_base_suffix="") as rig: + client: Final = anthropic.AsyncAnthropic( + base_url=str(rig.proxy.client.base_url), api_key=rig.proxy.key, max_retries=0 + ) + async with client.messages.stream( + model=rig.model, + max_tokens=64, + messages=[{"role": "user", "content": "weather in Paris?"}], + tools=[ + { + "name": "lookup_weather", + "description": "Get weather", + "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}}, + } + ], + tool_choice={"type": "auto"}, + metadata={"trace_marker": marker}, + ) as stream: + events: Final = tuple([event async for event in stream]) + response: Final = await stream.get_final_message() + assert events[-1].type == "message_stop", events + assert response.id == marker, response + assert response.model_dump(mode="json", exclude_unset=True) == _messages_caller_stream_response( + _messages_response(marker), rig.model + ), response + assert tuple(event.model_dump(mode="json", exclude_unset=True) for event in events) == ( + _messages_caller_stream( + _anthropic_stream_response(marker), + rig.model, + final_message=_messages_caller_stream_response(_messages_response(marker), rig.model), + ) + ), events + call: Final = response.content[0] + assert (call.type, call.id, call.name, call.input) == ( + "tool_use", + f"call_{marker}", + "lookup_weather", + {"city": "Paris"}, + ), response + _assert_tool_span_for_marker(_matching_marker_span(rig.destination, marker), marker, content="") + + asyncio.run(call()) + + +def test_arize_otel_v2_llm_span_carries_openinference_tool_calls_and_metadata(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a9-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _chat_upstream(request) + return _chat_response(marker) + + with _rig(gateway, tmp_path, upstream) as rig: + response: Final = _chat_request(rig.proxy, rig.model, marker) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_chat_response(marker), rig.model), response.text + attributes: Final = _matching_marker_span(rig.destination, marker) + _assert_tool_span_for_marker(attributes, marker) + + +def test_arize_otel_v2_responses_span_carries_openinference_tool_calls_and_metadata( + gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = "a10-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _responses_upstream(request, marker) + return _responses_response(marker) + + with _rig(gateway, tmp_path, upstream) as rig: + response: Final = _responses_request(rig.proxy, rig.model, marker) + assert response.status_code == 200, response.text + response_body: Final = _json_object(response.content) + assert _normalize_responses_caller_body(response_body) == _responses_caller_response( + _responses_response(marker), rig.model + ), response.text + _assert_tool_span_for_marker(_matching_marker_span(rig.destination, marker), marker) + + +def test_arize_otel_v2_a11_parallel_output_tool_calls(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a11-" + uuid.uuid4().hex + calls: Final = _response_tool_calls(marker, ("Paris", "Berlin")) + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=[{"role": "user", "content": "weather in Paris?"}]) + return _chat_response(marker, calls) + + with _rig(gateway, tmp_path, upstream) as rig: + response: Final = _chat_request(rig.proxy, rig.model, marker) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_chat_response(marker, calls), rig.model), ( + response.text + ) + attributes: Final = _matching_marker_span(rig.destination, marker) + expected_calls: Final = ( + ("call_" + marker + "-Paris", "lookup_weather", {"city": "Paris"}), + ("call_" + marker + "-Berlin", "lookup_weather", {"city": "Berlin"}), + ) + _assert_tool_span( + attributes, + marker=marker, + output=[{"role": "assistant", "content": None, "tool_calls": calls}], + calls=expected_calls, + metadata={"trace_marker": marker}, + baggage={"trace_marker": marker}, + ) + + +def test_arize_otel_v2_a12_plain_text_has_metadata_without_tool_calls(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a12-" + uuid.uuid4().hex + + def upstream(request: Request) -> Reply: + _assert_chat_request( + request, + messages=[{"role": "user", "content": "weather in Paris?"}], + include_tools=False, + ) + return _chat_plain_response(marker, "The weather is clear") + + with _rig(gateway, tmp_path, upstream) as rig: + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + { + "model": rig.model, + "messages": [{"role": "user", "content": "weather in Paris?"}], + "metadata": {"trace_marker": marker}, + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response( + _chat_plain_response(marker, "The weather is clear"), rig.model + ), response.text + attributes: Final = _matching_marker_span(rig.destination, marker) + assert _json_object(attributes["metadata"].encode()) == {"trace_marker": marker}, attributes + assert attributes["litellm.metadata.trace_marker"] == marker, attributes + assert not any(".tool_calls." in key for key in attributes), attributes + assert _json_messages(attributes["output.value"]) == [ + {"role": "assistant", "content": "The weather is clear"} + ], attributes + + +def test_arize_otel_v2_a13_multiturn_input_tool_calls(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a13-" + uuid.uuid4().hex + messages: Final = [ + {"role": "user", "content": "weather in Paris?"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call-prior", + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + } + ], + }, + {"role": "tool", "tool_call_id": "call-prior", "content": '{"temperature": 20}'}, + ] + upstream_messages: Final = [ + {key: value for key, value in message.items() if value is not None} for message in messages + ] + + def upstream(request: Request) -> Reply: + _assert_chat_request(request, messages=upstream_messages) + return _chat_response(marker) + + with _rig(gateway, tmp_path, upstream) as rig: + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + { + "model": rig.model, + "messages": messages, + "tools": CHAT_TOOLS, + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "metadata": {"trace_marker": marker}, + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(_chat_response(marker), rig.model), response.text + attributes: Final = _matching_marker_span(rig.destination, marker) + _assert_tool_span_for_marker(attributes, marker) + assert not any(key.startswith("llm.input_messages.") and ".tool_calls." in key for key in attributes), ( + attributes + ) + assert tuple(attributes[f"llm.input_messages.{index}.message.role"] for index in range(3)) == ( + "user", + "assistant", + "tool", + ), attributes + assert tuple(attributes[f"llm.input_messages.{index}.message.content"] for index in (0, 2)) == ( + "weather in Paris?", + '{"temperature": 20}', + ), attributes + expected_input_value: Final = [ + messages[0], + messages[1], + {"role": "tool", "content": '{"temperature": 20}'}, + ] + assert _json_messages(attributes["input.value"]) == expected_input_value, attributes + + +def test_arize_otel_v2_a14_two_choices_each_with_tool_calls(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "a14-" + uuid.uuid4().hex + first_call: Final = _response_tool_calls(marker, ("Paris",))[0] + second_call: Final = _response_tool_calls(marker, ("Berlin",))[0] + + def choice(index: int, call: dict[str, JsonValue]) -> dict[str, JsonValue]: + return { + "index": index, + "finish_reason": "tool_calls", + "message": {"role": "assistant", "content": None, "tool_calls": [call]}, + } + + expected_response: Final = Reply( + body=json.dumps( + { + "id": marker, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [choice(0, first_call), choice(1, second_call)], + "usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}, + } + ).encode() + ) + + def upstream(request: Request) -> Reply: + _assert_chat_request( + request, + messages=[{"role": "user", "content": "weather in Paris and Berlin?"}], + n=2, + ) + return expected_response + + with _rig(gateway, tmp_path, upstream) as rig: + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + { + "model": rig.model, + "messages": [{"role": "user", "content": "weather in Paris and Berlin?"}], + "tools": CHAT_TOOLS, + "tool_choice": {"type": "function", "function": {"name": "lookup_weather"}}, + "n": 2, + "metadata": {"trace_marker": marker}, + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == 200, response.text + assert _json_object(response.content) == _chat_caller_response(expected_response, rig.model), response.text + attributes: Final = _matching_marker_span(rig.destination, marker) + assert _json_messages(attributes["output.value"]) == [ + {"role": "assistant", "content": None, "tool_calls": [first_call]}, + {"role": "assistant", "content": None, "tool_calls": [second_call]}, + ], attributes + assert attributes["llm.output_messages.0.message.tool_calls.0.tool_call.id"] == str(first_call["id"]), ( + attributes + ) + assert attributes["llm.output_messages.1.message.tool_calls.0.tool_call.id"] == str(second_call["id"]), ( + attributes + ) + assert attributes["llm.output_messages.0.message.tool_calls.0.tool_call.function.name"] == "lookup_weather", ( + attributes + ) + assert ( + attributes["llm.output_messages.0.message.tool_calls.0.tool_call.function.arguments"] == '{"city": "Paris"}' + ), attributes + assert attributes["llm.output_messages.1.message.tool_calls.0.tool_call.function.name"] == "lookup_weather", ( + attributes + ) + assert ( + attributes["llm.output_messages.1.message.tool_calls.0.tool_call.function.arguments"] + == '{"city": "Berlin"}' + ), attributes + assert _json_object(attributes["metadata"].encode()) == {"trace_marker": marker}, attributes + assert attributes["litellm.metadata.trace_marker"] == marker, attributes + + +def _cache_call( + rig: Rig, surface: str, marker: str, *, cache_hit: bool = False +) -> tuple[tuple[str, str, str], str, httpx.Headers]: + match surface: + case "chat": + client: Final = _openai_client(rig.proxy) + raw: Final = client.chat.completions.with_raw_response.create( + model=rig.model, + messages=[{"role": "user", "content": marker}], + tools=CHAT_TOOLS, + tool_choice={"type": "function", "function": {"name": "lookup_weather"}}, + extra_body={"metadata": {"trace_marker": marker}, "cache": {"no-cache": False}}, + ) + response: Final = raw.parse() + assert response.model_dump(mode="json", exclude_unset=True) == _chat_caller_response( + _chat_response(marker), rig.model + ), response + assert response.choices[0].message.tool_calls is not None, response + call: Final = response.choices[0].message.tool_calls[0] + return (call.id, call.function.name, call.function.arguments), response.id, raw.headers + case "chat-stream": + client: Final = _openai_client(rig.proxy) + with client.chat.completions.with_streaming_response.create( + model=rig.model, + messages=[{"role": "user", "content": marker}], + tools=CHAT_TOOLS, + tool_choice={"type": "function", "function": {"name": "lookup_weather"}}, + stream=True, + stream_options={"include_usage": False}, + extra_body={"metadata": {"trace_marker": marker}, "cache": {"no-cache": False}}, + ) as raw: + chunks: Final = tuple(raw.parse()) + expected_chunks: Final = ( + _chat_cache_hit_caller_stream(marker, rig.model, (_chat_tool_call(marker),)) + if cache_hit + else _chat_caller_stream( + _chat_stream_response(marker, (_chat_tool_call(marker),), include_usage=False), rig.model + ) + ) + assert ( + _normalize_chat_caller_stream( + tuple(chunk.model_dump(mode="json", exclude_unset=True) for chunk in chunks) + ) + == expected_chunks + ), chunks + calls: Final = tuple(chain.from_iterable(chunk.choices[0].delta.tool_calls or () for chunk in chunks)) + if cache_hit: + assert len(calls) == 1, chunks + call: Final = calls[0] + assert ( + call.id is not None and call.function.name is not None and call.function.arguments is not None + ), chunks + return ( + (call.id, call.function.name, call.function.arguments), + chunks[0].id, + raw.headers, + ) + assert len(calls) == 2, chunks + call: Final = calls[0] + arguments_call: Final = calls[1] + assert ( + call.id is not None + and call.function.name is not None + and arguments_call.function.arguments is not None + ), chunks + return ( + (call.id, call.function.name, arguments_call.function.arguments), + chunks[0].id, + raw.headers, + ) + case "responses": + client: Final = _openai_client(rig.proxy) + raw: Final = client.responses.with_raw_response.create( + model=rig.model, + input=marker, + tools=RESPONSES_TOOLS, + tool_choice={"type": "function", "name": "lookup_weather"}, + extra_body={"metadata": {"trace_marker": marker}, "cache": {"no-cache": False}}, + ) + response: Final = raw.parse() + assert _normalize_responses_caller_body(response.model_dump(mode="json", exclude_unset=True)) == ( + _responses_caller_response(_responses_response(marker), rig.model) + ), response + call: Final = response.output[0] + assert call.type == "function_call", response + return (call.call_id, call.name, call.arguments), response.id, raw.headers + case "messages": + client: Final = anthropic.Anthropic( + base_url=str(rig.proxy.client.base_url), api_key=rig.proxy.key, max_retries=0 + ) + raw: Final = client.messages.with_raw_response.create( + model=rig.model, + max_tokens=64, + messages=[{"role": "user", "content": marker}], + tools=[ + { + "name": "lookup_weather", + "description": "Get weather", + "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}}, + } + ], + tool_choice={"type": "auto"}, + extra_body={"cache": {"no-cache": False}, "metadata": {"trace_marker": marker}}, + ) + response: Final = raw.parse() + assert response.model_dump(mode="json", exclude_unset=True) == _messages_caller_response( + _messages_response(marker), rig.model + ), response + call: Final = response.content[0] + assert call.type == "tool_use", response + return (call.id, call.name, json.dumps(call.input)), response.id, raw.headers + case _: + raise AssertionError(f"Unknown cache surface: {surface}") + + +@pytest.mark.parametrize("surface", ("chat", "chat-stream", "responses", "messages")) +def test_arize_otel_v2_a_cache(surface: str, gateway: Gateway, tmp_path: Path) -> None: + marker: Final = f"a-cache-{surface}-" + uuid.uuid4().hex + sentinel: Final = f"{marker}-sentinel" + + def upstream(request: Request) -> Reply: + body: Final = _json_object(request.body) + request_marker: Final = ( + _chat_request_marker(request) + if surface in ("chat", "chat-stream", "messages") + else string_value(object_value(body["metadata"])["trace_marker"]) + ) + assert request_marker in (marker, sentinel), request + if surface in ("chat", "chat-stream"): + _assert_chat_request( + request, + messages=[{"role": "user", "content": request_marker}], + stream=True if surface == "chat-stream" else None, + stream_options={"include_usage": False} if surface == "chat-stream" else None, + ) + return ( + _chat_stream_response(request_marker, (_chat_tool_call(request_marker),), include_usage=False) + if surface == "chat-stream" + else _chat_response(request_marker) + ) + if surface == "responses": + assert body.get("metadata") == {"trace_marker": request_marker}, body + _assert_responses_request(request, marker=request_marker, input_value=request_marker) + return _responses_response(request_marker) + _assert_messages_request(request, marker=request_marker, prompt=request_marker) + return _messages_response(request_marker) + + with _rig( + gateway, + tmp_path, + upstream, + general_settings={"always_include_stream_usage": False} if surface == "chat-stream" else None, + model_name="anthropic/claude-opus-5-5" if surface == "messages" else "gpt-4o-mini", + api_base_suffix="" if surface == "messages" else "/v1", + workers=1, + ) as rig: + expected_arguments: Final = '{"city": "Paris"}' + first: Final = _cache_call(rig, surface, marker) + assert first[0] == (f"call_{marker}", "lookup_weather", expected_arguments), first + assert first[1].startswith("resp_") if surface == "responses" else first[1] == marker, first + assert not first[2].get("x-litellm-cache-key"), first[2] + first_span: Final = _matching_marker_span(rig.destination, marker) + _assert_tool_span_for_marker(first_span, marker) + second: Final = _cache_call(rig, surface, marker, cache_hit=True) + assert second[0] == first[0], second + assert second[1].startswith("resp_") if surface == "responses" else second[1] == first[1], second + forwarded: Final = tuple( + request for request in rig.provider.drain() if request.method == "POST" and marker.encode() in request.body + ) + assert len(forwarded) == 1, forwarded + if surface == "messages": + assert not second[2].get("x-litellm-cache-key"), second[2] + else: + assert second[2].get("x-litellm-cache-key"), second[2] + forwarded_body: Final = _json_object(forwarded[0].body) + if surface in ("chat", "chat-stream"): + assert "metadata" not in forwarded_body, forwarded[0] + elif surface == "responses": + assert forwarded_body["metadata"] == {"trace_marker": marker}, forwarded[0] + else: + assert forwarded_body["metadata"] == {}, forwarded[0] + sentinel_response: Final = _cache_call(rig, surface, sentinel) + assert not sentinel_response[2].get("x-litellm-cache-key"), sentinel_response[2] + sentinel_forwarded: Final = tuple( + request + for request in rig.provider.drain() + if request.method == "POST" and sentinel.encode() in request.body + ) + assert len(sentinel_forwarded) == 1, sentinel_forwarded + spans: Final = _llm_spans_through_markers(rig.destination, (sentinel,)) + sentinel_spans: Final = tuple( + attributes for attributes in spans if attributes.get("litellm.metadata.trace_marker") == sentinel + ) + assert len(sentinel_spans) == 1, spans + _assert_tool_span_for_marker(sentinel_spans[0], sentinel) + assert not any(attributes.get("litellm.metadata.trace_marker") == marker for attributes in spans), spans + + +def test_arize_otel_v2_cache_hit_exports_llm_span(gateway: Gateway, tmp_path: Path) -> None: + pytest.skip("BUG: LIT-9127 response-cache hits emit no OTel v2 LLM span") + marker: Final = "a-cache-hit-" + uuid.uuid4().hex + sentinel: Final = f"{marker}-sentinel" + + def upstream(request: Request) -> Reply: + request_marker: Final = _chat_request_marker(request) + assert request_marker in (marker, sentinel), request + _assert_chat_request(request, messages=[{"role": "user", "content": request_marker}]) + return _chat_response(request_marker) + + with _rig(gateway, tmp_path, upstream, workers=1) as rig: + first: Final = _cache_call(rig, "chat", marker) + assert first[0] == (f"call_{marker}", "lookup_weather", '{"city": "Paris"}'), first + assert first[1] == marker, first + assert not first[2].get("x-litellm-cache-key"), first[2] + second: Final = _cache_call(rig, "chat", marker, cache_hit=True) + assert second[0] == first[0], second + assert second[1] == first[1], second + assert second[2].get("x-litellm-cache-key"), second[2] + forwarded: Final = tuple( + request for request in rig.provider.drain() if request.method == "POST" and marker.encode() in request.body + ) + assert len(forwarded) == 1, forwarded + sentinel_response: Final = _cache_call(rig, "chat", sentinel) + assert not sentinel_response[2].get("x-litellm-cache-key"), sentinel_response[2] + sentinel_forwarded: Final = tuple( + request + for request in rig.provider.drain() + if request.method == "POST" and sentinel.encode() in request.body + ) + assert len(sentinel_forwarded) == 1, sentinel_forwarded + spans: Final = _llm_spans_through_markers(rig.destination, (sentinel,)) + marker_spans: Final = tuple( + attributes for attributes in spans if attributes.get("litellm.metadata.trace_marker") == marker + ) + assert len(marker_spans) == 2, spans + _assert_tool_span_for_marker(marker_spans[1], marker) diff --git a/tests/integration/observability/test_guardrail_attachment.py b/tests/integration/observability/test_guardrail_attachment.py new file mode 100644 index 00000000000..f291b3877f3 --- /dev/null +++ b/tests/integration/observability/test_guardrail_attachment.py @@ -0,0 +1,153 @@ +from __future__ import annotations + +import json +import shutil +import uuid +from collections.abc import Iterator +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import pytest +import yaml + +from tests.integration._support.client import JSON_OBJECT, Gateway, gateway_from_environment, object_value, string_value +from tests.integration._support.process import owned_proxy +from tests.integration._support.wire import Reply, Request, Wire, wire_server + +_ATTACHABLE: Final = "attachable-during-guard" +_WORDS: Final = "custom-words-during-guard" +_HEADER: Final = "x-litellm-applied-guardrails" + + +@dataclass(frozen=True, slots=True) +class _Rig: + gateway: Gateway + policy: Wire + upstream: Wire + model: str + + +def _completion(request: Request) -> Reply: + assert request.target == "/v1/chat/completions", request.target + return Reply( + body=json.dumps( + { + "id": f"chatcmpl-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "guarded"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4}, + } + ).encode() + ) + + +def _allow(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + +def _config(directory: Path, policy: Wire) -> Path: + shutil.copy(Path("litellm/proxy/example_config_yaml/custom_guardrail.py"), directory / "custom_guardrail.py") + configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + configuration["guardrails"] = [ + { + "guardrail_name": _ATTACHABLE, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "during_call", + "api_base": policy.url, + "api_key": "synthetic-guardrail-key", + }, + }, + { + "guardrail_name": _WORDS, + "litellm_params": {"guardrail": "custom_guardrail.myCustomGuardrail", "mode": "during_call"}, + }, + ] + path: Final = directory / "guardrail-attachment.yaml" + path.write_text(yaml.safe_dump(configuration)) + return path + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]: + directory: Final = tmp_path_factory.mktemp("guardrail-attachment") + with ( + wire_server(_allow) as policy, + wire_server(_completion) as upstream, + gateway_from_environment() as shared, + owned_proxy(shared, directory, {}, config=_config(directory, policy)) as owned, + owned.scenario() as scenario, + ): + yield _Rig(owned, policy, upstream, scenario.model(api_base=f"{upstream.url}/v1", api_key="sk-fixture")) + + +def _prompt(text: str) -> str: + return f"{text} {uuid.uuid4().hex}" + + +def _ask(rig: _Rig, key: str | None, prompt: str, guardrails: list[str] | None = None) -> httpx.Response: + body: Final = {"model": rig.model, "messages": [{"role": "user", "content": prompt}]} + return rig.gateway.request( + "POST", "/v1/chat/completions", body if guardrails is None else {**body, "guardrails": guardrails}, key=key + ) + + +def _served_without_guardrail(rig: _Rig, response: httpx.Response, prompt: str) -> None: + assert response.status_code == 200, response.text + assert _HEADER not in response.headers, dict(response.headers) + assert rig.policy.drain() == () + forwarded: Final = rig.upstream.drain() + assert len(forwarded) == 1 and prompt in forwarded[0].body.decode(), forwarded + + +def _served_with_attachable(rig: _Rig, response: httpx.Response, prompt: str) -> None: + assert response.status_code == 200, response.text + assert response.headers[_HEADER] == _ATTACHABLE, dict(response.headers) + inspected: Final = rig.policy.drain() + assert len(inspected) == 1 and prompt in inspected[0].body.decode(), inspected + forwarded: Final = rig.upstream.drain() + assert len(forwarded) == 1 and prompt in forwarded[0].body.decode(), forwarded + + +def test_a_request_with_an_empty_guardrail_list_is_served_without_the_applied_header(rig: _Rig) -> None: + prompt: Final = _prompt("no guardrails") + _served_without_guardrail(rig, _ask(rig, None, prompt, []), prompt) + + +def test_a_key_carrying_a_guardrail_applies_it_and_a_plain_key_does_not(rig: _Rig) -> None: + with rig.gateway.scenario() as scenario: + plain: Final = scenario.key() + guarded: Final = scenario.key(guardrails=[_ATTACHABLE]) + plain_prompt: Final = _prompt("plain key") + _served_without_guardrail(rig, _ask(rig, plain, plain_prompt), plain_prompt) + guarded_prompt: Final = _prompt("guarded key") + _served_with_attachable(rig, _ask(rig, guarded, guarded_prompt), guarded_prompt) + + +def test_a_team_carrying_a_guardrail_applies_it_to_its_keys_only(rig: _Rig) -> None: + with rig.gateway.scenario() as scenario: + team: Final = scenario.team(guardrails=[_ATTACHABLE]) + outside: Final = scenario.key() + member: Final = scenario.key(team_id=team) + outside_prompt: Final = _prompt("outside team") + _served_without_guardrail(rig, _ask(rig, outside, outside_prompt), outside_prompt) + member_prompt: Final = _prompt("team key") + _served_with_attachable(rig, _ask(rig, member, member_prompt), member_prompt) + + +def test_a_during_call_custom_guardrail_rejects_a_request_naming_the_banned_word(rig: _Rig) -> None: + unguarded_prompt: Final = _prompt("what is litellm") + _served_without_guardrail(rig, _ask(rig, None, unguarded_prompt), unguarded_prompt) + refused: Final = _ask(rig, None, _prompt("what is litellm"), [_WORDS]) + rig.upstream.drain() + assert refused.status_code >= 400, refused.text + error: Final = object_value(JSON_OBJECT.validate_json(refused.content)["error"]) + assert "Guardrail failed words - `litellm` detected" in string_value(error["message"]), refused.text + assert rig.policy.drain() == () diff --git a/tests/integration/observability/test_guardrail_stream_scope.py b/tests/integration/observability/test_guardrail_stream_scope.py new file mode 100644 index 00000000000..8b832f07aa4 --- /dev/null +++ b/tests/integration/observability/test_guardrail_stream_scope.py @@ -0,0 +1,856 @@ +from __future__ import annotations + +import base64 +import json +import os +import uuid +from collections.abc import Iterator, Mapping, Sequence +from dataclasses import dataclass +from itertools import chain +from pathlib import Path +from types import MappingProxyType +from typing import Final, Literal, TypeAlias + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, JsonValue, Scenario, eventually +from integration._support.database import read_rows, write_rows +from integration._support.process import owned_proxy_process +from integration._support.upstream import ( + _aws_event_frame, # pyright: ignore[reportPrivateUsage] # project Bedrock event-stream encoder +) +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import TypeAdapter + +import litellm + +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +EndpointScope: TypeAlias = Literal["streaming", "non_streaming"] +ModelKind: TypeAlias = Literal["router", "direct"] +StreamingAction: TypeAlias = Literal["converse-stream", "invoke-with-response-stream"] +NonStreamingAction: TypeAlias = Literal["converse", "invoke"] +BedrockAction: TypeAlias = StreamingAction | NonStreamingAction +BEDROCK_MODEL_ID: Final = "anthropic.claude-sonnet-5-v1:0" +BEDROCK_FALSE_POSITIVE_MODEL_ID: Final = "anthropic.claude-converse-stream-test-v1:0" +BEDROCK_EVENT_STREAM: Final = "application/vnd.amazon.eventstream" +GUARDRAIL_PATH: Final = "/beta/litellm_basic_guardrail_api" +SERVER_STREAMING_CLASSIFICATION_KEY: Final = "litellm_server_streaming_classification" +STREAMING_ACTIONS: Final[tuple[StreamingAction, ...]] = ( + "converse-stream", + "invoke-with-response-stream", +) +NON_STREAMING_ACTIONS: Final[tuple[NonStreamingAction, ...]] = ("converse", "invoke") +SCOPES: Final[tuple[EndpointScope, ...]] = ("streaming", "non_streaming") +HOSTILE_CLASSIFICATION_CASES: Final = ( + pytest.param("is_streaming_request", True, id="boolean-marker"), + pytest.param("is_streaming_request", "litellm-server-streaming", id="server-marker-string"), + pytest.param("litellm_server_streaming_classification", True, id="classification-field"), +) +WORKTREE: Final = Path(__file__).resolve().parents[3] +LITELLM_PATH: Final = Path(litellm.__file__).resolve() +assert LITELLM_PATH.is_relative_to(WORKTREE), (LITELLM_PATH, WORKTREE) +print(f"stream_scope repro litellm import: {LITELLM_PATH}") # noqa: T201 # required worktree evidence + + +def _json(value: object) -> bytes: + return json.dumps(value, separators=(",", ":")).encode() + + +def _strings(value: JsonValue) -> tuple[str, ...]: + if isinstance(value, str): + return (value,) + if isinstance(value, list): + return tuple(chain.from_iterable(_strings(item) for item in value)) + if isinstance(value, dict): + return tuple(chain.from_iterable(_strings(item) for item in value.values())) + return () + + +def _key_names(value: JsonValue) -> tuple[str, ...]: + if isinstance(value, dict): + return tuple(value) + tuple(chain.from_iterable(_key_names(item) for item in value.values())) + if isinstance(value, list): + return tuple(chain.from_iterable(_key_names(item) for item in value)) + return () + + +def _marker(body: JsonValue) -> str: + return next(value for value in _strings(body) if value.startswith("scope-")) + + +def _bedrock_converse_stream(marker: str) -> bytes: + return b"".join( + _aws_event_frame(event_type, payload, marker, marker) + for event_type, payload in ( + ("messageStart", {"role": "assistant"}), + ( + "contentBlockDelta", + {"delta": {"text": f"scripted Bedrock reply {marker}"}, "contentBlockIndex": 0}, + ), + ("messageStop", {"stopReason": "end_turn"}), + ("metadata", {"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}}), + ) + ) + + +def _invoke_chunk(payload: Mapping[str, JsonValue], marker: str) -> bytes: + encoded: Final = base64.b64encode(_json(payload)).decode() + return _aws_event_frame("chunk", {"bytes": encoded}, marker, marker) + + +def _bedrock_invoke_stream(marker: str) -> bytes: + events: Final = ( + { + "type": "message_start", + "message": { + "id": f"msg-{marker}", + "type": "message", + "role": "assistant", + "model": BEDROCK_MODEL_ID, + "content": [], + "stop_reason": None, + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 0}, + }, + }, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": f"scripted Bedrock reply {marker}"}, + }, + {"type": "content_block_stop", "index": 0}, + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn"}, + "usage": {"input_tokens": 11, "output_tokens": 4}, + }, + {"type": "message_stop"}, + ) + return b"".join(_invoke_chunk(event, marker) for event in events) + + +def _provider(request: Request) -> Reply: + if not request.body: + return Reply(status=400, body=_json({"error": "empty request body"})) + body: Final = JSON_OBJECT.validate_json(request.body) + marker: Final = _marker(body) + target: Final = request.target.split("?", 1)[0] + if target.startswith("/passthrough"): + if body.get("stream") is True: + streamed_response: Final = _json({"received": body}) + return Reply( + content_type="text/event-stream", + chunks=(b"data: " + streamed_response + b"\n\n", b"data: [DONE]\n\n"), + ) + return Reply(body=_json({"received": body})) + if target.endswith("/converse-stream"): + return Reply(body=_bedrock_converse_stream(marker), content_type=BEDROCK_EVENT_STREAM) + if target.endswith("/invoke-with-response-stream"): + return Reply(body=_bedrock_invoke_stream(marker), content_type=BEDROCK_EVENT_STREAM) + if target.endswith("/converse") or target.endswith("/invoke"): + return Reply(body=_json({"output": f"scripted Bedrock reply {marker}"})) + if target == "/v1/chat/completions": + if body.get("stream") is True: + chunk: Final = { + "id": f"chatcmpl-{marker}", + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": f"scripted chat reply {marker}"}, + "finish_reason": None, + } + ], + } + final_chunk: Final = { + **chunk, + "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], + } + return Reply( + content_type="text/event-stream", + chunks=( + b"data: " + _json(chunk) + b"\n\n", + b"data: " + _json(final_chunk) + b"\n\n", + b"data: [DONE]\n\n", + ), + ) + return Reply( + body=_json( + { + "id": f"chatcmpl-{marker}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": f"scripted chat reply {marker}"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } + ) + ) + return Reply(status=404, body=_json({"error": f"unexpected upstream path: {target}"})) + + +def _sink(request: Request) -> Reply: + assert request.target.endswith(GUARDRAIL_PATH), request.target + assert b"scope-" in request.body, request.body.decode() + return Reply(body=_json({"action": "NONE"})) + + +def _rail(name: str, sink: Wire, scope: EndpointScope) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": False, + "stream_scope": scope, + "api_base": f"{sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + }, + } + + +def _chat_proxy_config(provider_url: str, guardrails: list[dict[str, JsonValue]]) -> dict[str, object]: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + return { + **config, + "guardrails": guardrails, + "model_list": [ + { + "model_name": "scope-invalid-config-chat", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider_url}/v1", + "api_key": "synthetic-provider-key", + }, + } + ], + } + + +@dataclass(frozen=True, slots=True) +class ReproRig: + candidate: Gateway + direct_candidate: Gateway + scenario: Scenario + models: Mapping[str, str] + rails: Mapping[str, str] + provider: Wire + sink: Wire + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[ReproRig]: + with httpx.Client( + base_url=os.environ["INTEGRATION_PROXY_URL"], + timeout=30, + trust_env=False, + ) as root_client: + root_gateway: Final = Gateway( + root_client, + os.environ.get("INTEGRATION_MASTER_KEY", "sk-integration-master"), + os.environ["INTEGRATION_UPSTREAM_URL"], + ) + directory: Final = tmp_path_factory.mktemp("guardrail-stream-scope-repro") + with wire_server(_provider) as provider, wire_server(_sink) as sink: + rails: Final = MappingProxyType( + { + "bedrock_streaming": "bedrock_streaming", + "bedrock_non_streaming": "bedrock_non_streaming", + "chat_streaming": "chat_streaming", + "passthrough_streaming": "passthrough_streaming", + "passthrough_non_streaming": "passthrough_non_streaming", + "passthrough_spoof_streaming": "passthrough_spoof_streaming", + } + ) + models: Final = MappingProxyType( + { + "chat": "scope-chat", + "bedrock_router": "scope-bedrock-router", + "bedrock_false_positive_router": "scope-bedrock-converse-stream-model", + } + ) + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + _rail(rails["bedrock_streaming"], sink, "streaming"), + _rail(rails["bedrock_non_streaming"], sink, "non_streaming"), + _rail(rails["chat_streaming"], sink, "streaming"), + _rail(rails["passthrough_streaming"], sink, "streaming"), + _rail(rails["passthrough_non_streaming"], sink, "non_streaming"), + _rail(rails["passthrough_spoof_streaming"], sink, "streaming"), + ] + config["model_list"] = [ + { + "model_name": models["chat"], + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider.url}/v1", + "api_key": "synthetic-provider-key", + }, + }, + { + "model_name": models["bedrock_router"], + "litellm_params": { + "model": f"bedrock/{BEDROCK_MODEL_ID}", + "api_base": provider.url, + "aws_access_key_id": "AKIASYNTHETICSTREAMSCOPE", + "aws_secret_access_key": "synthetic-bedrock-secret", + "aws_region_name": "us-east-1", + }, + }, + { + "model_name": models["bedrock_false_positive_router"], + "litellm_params": { + "model": f"bedrock/{BEDROCK_FALSE_POSITIVE_MODEL_ID}", + "api_base": provider.url, + "aws_access_key_id": "AKIASYNTHETICSTREAMSCOPE", + "aws_secret_access_key": "synthetic-bedrock-secret", + "aws_region_name": "us-east-1", + }, + }, + ] + config["environment_variables"] = { + "AWS_BEDROCK_RUNTIME_ENDPOINT": provider.url, + "AWS_ACCESS_KEY_ID": "AKIASYNTHETICSTREAMSCOPE", + "AWS_SECRET_ACCESS_KEY": "synthetic-bedrock-secret", + "AWS_REGION": "us-east-1", + "AWS_REGION_NAME": "us-east-1", + } + config["general_settings"]["pass_through_endpoints"] = [ + { + "path": "/pt-forward", + "target": f"{provider.url}/passthrough", + "include_subpath": True, + }, + { + "path": "/pt-spoof", + "target": f"{provider.url}/passthrough", + "include_subpath": True, + "guardrails": {rails["passthrough_spoof_streaming"]: None}, + }, + { + "path": "/pt-scope", + "target": f"{provider.url}/passthrough", + "include_subpath": True, + "guardrails": { + rails["passthrough_streaming"]: None, + rails["passthrough_non_streaming"]: None, + }, + }, + ] + config_path: Final = directory / "stream-scope-repro.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy_process(root_gateway, directory, {}, config=config_path, workers=1) as owned: + with owned.gateway.scenario() as scenario: + direct_config: Final = { + **config, + "model_list": [ + *config["model_list"], + { + "model_name": f"scope-unused-{uuid.uuid4().hex}*", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider.url}/v1", + "api_key": "synthetic-provider-key", + }, + }, + ], + } + direct_config_path: Final = directory / "stream-scope-direct.yaml" + direct_config_path.write_text(yaml.safe_dump(direct_config)) + with owned_proxy_process( + root_gateway, + directory, + {}, + config=direct_config_path, + workers=1, + ) as direct: + yield ReproRig( + owned.gateway, + direct.gateway, + scenario, + models, + rails, + provider, + sink, + ) + + +def _bedrock_request( + action: BedrockAction, + model_path: str, + marker: str, +) -> tuple[str, dict[str, JsonValue]]: + if action in ("converse", "converse-stream"): + body: Final = { + "messages": [{"role": "user", "content": [{"text": marker}]}], + "inferenceConfig": {"maxTokens": 16}, + } + else: + body = { + "anthropic_version": "bedrock-2023-05-31", + "max_tokens": 16, + "messages": [{"role": "user", "content": [{"type": "text", "text": marker}]}], + } + return f"/bedrock/model/{model_path}/{action}", body + + +def _matching_requests(wire: Wire, marker: str) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if marker.encode() in request.body) + + +def _rail_scans(rows: Sequence[Request], rail_name: str, marker: str) -> tuple[Request, ...]: + return tuple( + request for request in rows if request.target.startswith(f"/{rail_name}/") and marker.encode() in request.body + ) + + +def _chat_request_with_scans( + gateway: Gateway, + sink: Wire, + model: str, + marker: str, + streamed: bool, + rail_name: str, +) -> tuple[httpx.Response, tuple[Request, ...]]: + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": marker}], + "stream": streamed, + "guardrails": [rail_name], + }, + ) + return response, _rail_scans(sink.drain(), rail_name, marker) + + +def _spend_row_for_call(call_id: str, content: bytes) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + "SELECT request_id, litellm_call_id, spend, prompt_tokens, completion_tokens, metadata " + 'FROM "LiteLLM_SpendLogs" WHERE request_id=%s OR litellm_call_id=%s', + (call_id, call_id), + ), + lambda values: len(values) >= 1, + seconds=70, + ) + assert len(rows) == 1, (call_id, rows, content) + assert call_id in (rows[0]["request_id"], rows[0]["litellm_call_id"]), content + return rows[0] + + +@pytest.mark.parametrize("model_kind", ("router", "direct")) +@pytest.mark.parametrize("action", STREAMING_ACTIONS) +@pytest.mark.parametrize("scope", SCOPES) +def test_bedrock_streaming_actions_run_streaming_scoped_rails( + rig: ReproRig, + model_kind: ModelKind, + action: StreamingAction, + scope: EndpointScope, +) -> None: + marker: Final = f"scope-bedrock-stream-{uuid.uuid4().hex}" + model_path: Final = rig.models["bedrock_router"] if model_kind == "router" else BEDROCK_MODEL_ID + path, body = _bedrock_request(action, model_path, marker) + expected_body: Final = ( + _bedrock_converse_stream(marker) if action == "converse-stream" else _bedrock_invoke_stream(marker) + ) + call_id: Final = f"stream-scope-{uuid.uuid4().hex}" + key: Final = rig.scenario.key(guardrails=[rig.rails[f"bedrock_{scope}"]]) + candidate: Final = rig.candidate if model_kind == "router" else rig.direct_candidate + response: Final = candidate.request( + "POST", + path, + body, + key=key, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 200, response.content + assert response.headers.get("content-type") == BEDROCK_EVENT_STREAM, dict(response.headers) + assert response.content == expected_body, response.content + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows, response.content) + provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + key_names: Final = _key_names(provider_body) + assert "is_streaming_request" not in key_names, provider_body + assert not tuple(name for name in key_names if name.startswith("litellm_")), provider_body + _spend_row_for_call(call_id, response.content) + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails[f"bedrock_{scope}"], marker) + expected_scans: Final = int(scope == "streaming") + assert len(sink_rows) == expected_scans, (marker, model_kind, action, scope, sink_rows, response.content) + + +@pytest.mark.parametrize("model_kind", ("router", "direct")) +@pytest.mark.parametrize("action", NON_STREAMING_ACTIONS) +@pytest.mark.parametrize("scope", SCOPES) +def test_bedrock_non_streaming_actions_run_non_streaming_scoped_rails( + rig: ReproRig, + model_kind: ModelKind, + action: NonStreamingAction, + scope: EndpointScope, +) -> None: + marker: Final = f"scope-bedrock-nonstream-{uuid.uuid4().hex}" + model_path: Final = rig.models["bedrock_router"] if model_kind == "router" else BEDROCK_MODEL_ID + path, body = _bedrock_request(action, model_path, marker) + key: Final = rig.scenario.key(guardrails=[rig.rails[f"bedrock_{scope}"]]) + candidate: Final = rig.candidate if model_kind == "router" else rig.direct_candidate + response: Final = candidate.request("POST", path, body, key=key) + assert response.status_code == 200, response.text + assert marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows, response.text) + provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + key_names: Final = _key_names(provider_body) + assert "is_streaming_request" not in key_names, provider_body + assert not tuple(name for name in key_names if name.startswith("litellm_")), provider_body + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails[f"bedrock_{scope}"], marker) + expected_scans: Final = int(scope == "non_streaming") + assert len(sink_rows) == expected_scans, (marker, model_kind, action, scope, sink_rows, response.text) + + +@pytest.mark.parametrize("model_kind", ("router", "direct")) +def test_bedrock_model_id_streaming_action_text_on_converse_is_non_streaming( + rig: ReproRig, + model_kind: ModelKind, +) -> None: + marker: Final = f"scope-bedrock-converse-model-{uuid.uuid4().hex}" + model_path: Final = ( + rig.models["bedrock_false_positive_router"] if model_kind == "router" else BEDROCK_FALSE_POSITIVE_MODEL_ID + ) + path, body = _bedrock_request("converse", model_path, marker) + call_id: Final = f"stream-scope-bedrock-{uuid.uuid4().hex}" + key: Final = rig.scenario.key( + guardrails=[ + rig.rails["bedrock_streaming"], + rig.rails["bedrock_non_streaming"], + ] + ) + candidate: Final = rig.candidate if model_kind == "router" else rig.direct_candidate + response: Final = candidate.request( + "POST", + path, + body, + key=key, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 200, response.text + assert marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, model_kind, provider_rows, response.text) + sink_rows: Final = rig.sink.drain() + streaming_rows: Final = _rail_scans(sink_rows, rig.rails["bedrock_streaming"], marker) + non_streaming_rows: Final = _rail_scans(sink_rows, rig.rails["bedrock_non_streaming"], marker) + assert streaming_rows == (), (marker, model_kind, streaming_rows, response.text) + assert len(non_streaming_rows) == 1, (marker, model_kind, non_streaming_rows, response.text) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("stream-absent", "stream-true")) +def test_configured_passthrough_forwards_caller_is_streaming_request_field( + rig: ReproRig, + streamed: bool, +) -> None: + marker: Final = f"scope-passthrough-forward-{uuid.uuid4().hex}" + caller_value: Final = f"caller-{uuid.uuid4().hex}" + body: Final = { + "marker": marker, + "is_streaming_request": caller_value, + **({"stream": True} if streamed else {}), + } + response: Final = rig.candidate.request("POST", "/pt-forward", body) + assert response.status_code == 200, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows, response.text) + upstream_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + assert upstream_body.get("is_streaming_request") == caller_value, ( + marker, + caller_value, + upstream_body, + response.text, + ) + assert upstream_body == body, (marker, body, upstream_body, response.text) + if streamed: + assert response.headers.get("content-type", "").lower().startswith("text/event-stream"), dict(response.headers) + event_body: Final = response.text.removeprefix("data: ").split("\n", maxsplit=1)[0] + response_body: Final = JSON_OBJECT.validate_json(event_body) + else: + response_body = JSON_OBJECT.validate_json(response.content) + assert response_body == {"received": body}, response.text + + +@pytest.mark.parametrize(("hostile_field", "hostile_value"), HOSTILE_CLASSIFICATION_CASES) +def test_client_cannot_spoof_server_stream_classification( + rig: ReproRig, + hostile_field: str, + hostile_value: JsonValue, +) -> None: + chat_marker: Final = f"scope-chat-spoof-{uuid.uuid4().hex}" + chat_body: Final = { + "model": rig.models["chat"], + "messages": [{"role": "user", "content": chat_marker}], + "stream": False, + hostile_field: hostile_value, + } + chat_key: Final = rig.scenario.key(guardrails=[rig.rails["chat_streaming"]]) + chat_response: Final = rig.candidate.request( + "POST", + "/v1/chat/completions", + chat_body, + key=chat_key, + ) + chat_provider_rows: Final = _matching_requests(rig.provider, chat_marker) + chat_upstream_body: Final = chat_provider_rows[0].body.decode() if chat_provider_rows else "" + print( # noqa: T201 # required chat hostile-body observation + f"chat hostile {hostile_field}={hostile_value!r}: " + f"status={chat_response.status_code}, response={chat_response.text!r}, upstream={chat_upstream_body}" + ) + assert chat_response.status_code == 200, chat_response.text + assert len(chat_provider_rows) == 1, (chat_marker, chat_provider_rows, chat_response.text) + assert JSON_OBJECT.validate_json(chat_provider_rows[0].body) == { + "messages": [{"role": "user", "content": chat_marker}], + "model": "gpt-4o-mini", + hostile_field: hostile_value, + }, (hostile_field, hostile_value, chat_upstream_body) + assert _rail_scans(rig.sink.drain(), rig.rails["chat_streaming"], chat_marker) == (), ( + chat_marker, + hostile_field, + hostile_value, + chat_response.text, + ) + + +@pytest.mark.parametrize(("hostile_field", "hostile_value"), HOSTILE_CLASSIFICATION_CASES) +def test_configured_passthrough_cannot_spoof_server_stream_classification( + rig: ReproRig, + hostile_field: str, + hostile_value: JsonValue, +) -> None: + passthrough_marker: Final = f"scope-passthrough-spoof-{uuid.uuid4().hex}" + passthrough_body: Final = { + "marker": passthrough_marker, + "stream": False, + hostile_field: hostile_value, + } + passthrough_response: Final = rig.candidate.request("POST", "/pt-spoof", passthrough_body) + assert passthrough_response.status_code == 200, passthrough_response.text + passthrough_provider_rows: Final = _matching_requests(rig.provider, passthrough_marker) + assert len(passthrough_provider_rows) == 1, ( + passthrough_marker, + passthrough_provider_rows, + passthrough_response.text, + ) + assert _rail_scans(rig.sink.drain(), rig.rails["passthrough_spoof_streaming"], passthrough_marker) == (), ( + passthrough_marker, + hostile_field, + hostile_value, + passthrough_response.text, + ) + + +@pytest.mark.parametrize( + ("streamed", "request_fields"), + ((True, {"stream": True}), (False, {"stream": False}), (False, {})), + ids=("stream-true", "stream-false", "stream-absent"), +) +def test_passthrough_scope_follows_proxy_stream_decision( + rig: ReproRig, + streamed: bool, + request_fields: dict[str, JsonValue], +) -> None: + marker: Final = f"scope-passthrough-scope-{uuid.uuid4().hex}" + body: Final = {"marker": marker, **request_fields} + response: Final = rig.candidate.request("POST", "/pt-scope", body) + assert response.status_code == 200, response.text + if streamed: + expected_frame: Final = f"data: {_json({'received': body}).decode()}\n\ndata: [DONE]\n\n" + assert response.headers.get("content-type", "").lower().startswith("text/event-stream"), dict(response.headers) + assert response.text == expected_frame, response.text + assert response.headers.get("transfer-encoding", "").lower() == "chunked", dict(response.headers) + else: + response_body: Final = JSON_OBJECT.validate_json(response.content) + assert response_body == {"received": body}, response.text + assert "content-length" in response.headers, dict(response.headers) + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows, response.text) + upstream_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + assert upstream_body == body, (marker, body, upstream_body, response.text) + assert SERVER_STREAMING_CLASSIFICATION_KEY not in upstream_body, upstream_body + sink_rows: Final = rig.sink.drain() + streaming_rows: Final = _rail_scans(sink_rows, rig.rails["passthrough_streaming"], marker) + non_streaming_rows: Final = _rail_scans(sink_rows, rig.rails["passthrough_non_streaming"], marker) + assert len(streaming_rows) == int(streamed), (marker, streamed, streaming_rows, response.text) + assert len(non_streaming_rows) == int(not streamed), (marker, streamed, non_streaming_rows, response.text) + + +def test_invalid_yaml_stream_scope_keeps_rail_running_on_both_shapes(rig: ReproRig, tmp_path: Path) -> None: + name: Final = f"scope-invalid-yaml-{uuid.uuid4().hex}" + invalid_rail: Final = { + "guardrail_name": name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": False, + "stream_scope": "sometimes", + "api_base": f"{rig.sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + }, + } + config: Final = _chat_proxy_config(rig.provider.url, [invalid_rail]) + config_path: Final = tmp_path / "invalid-stream-scope.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy_process(rig.candidate, tmp_path, {}, config=config_path, workers=1) as owned: + markers: Final = tuple(f"scope-invalid-yaml-{int(streamed)}-{uuid.uuid4().hex}" for streamed in (False, True)) + observations: Final = tuple( + _chat_request_with_scans( + owned.gateway, + rig.sink, + "scope-invalid-config-chat", + marker, + streamed, + name, + ) + for streamed, marker in zip((False, True), markers) + ) + assert tuple(response.status_code for response, _ in observations) == (200, 200), tuple( + response.text for response, _ in observations + ) + assert tuple(len(sink_rows) for _, sink_rows in observations) == (1, 1), (markers, observations) + listed: Final = owned.gateway.request("GET", "/guardrails/list") + assert listed.status_code == 200, listed.text + list_payload: Final = JSON_OBJECT.validate_json(listed.content) + listed_guardrails: Final = list_payload.get("guardrails") + assert isinstance(listed_guardrails, list), list_payload + listed_rail: Final = next( + (row for row in listed_guardrails if isinstance(row, dict) and row.get("guardrail_name") == name), + None, + ) + assert isinstance(listed_rail, dict), list_payload + listed_params: Final = listed_rail.get("litellm_params") + assert isinstance(listed_params, dict), listed_rail + assert listed_params.get("stream_scope") is None, listed_rail + + +def test_persisted_invalid_stream_scope_row_stays_readable_and_enforced( + rig: ReproRig, + tmp_path: Path, +) -> None: + guardrail_id: Final = str(uuid.uuid4()) + guardrail_name: Final = f"scope-invalid-persisted-{uuid.uuid4().hex}" + params: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": False, + "api_base": f"{rig.sink.url}/{guardrail_name}", + "api_key": "synthetic-guardrail-key", + "stream_scope": "sometimes", + } + database_url: Final = os.environ.get("INTEGRATION_PROXY_DATABASE_URL") or os.environ["DATABASE_URL"] + write_rows( + 'INSERT INTO "LiteLLM_GuardrailsTable" ' + '("guardrail_id", "guardrail_name", "litellm_params", "created_at", "updated_at") ' + "VALUES (%s, %s, %s::jsonb, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)", + (guardrail_id, guardrail_name, json.dumps(params)), + database_url=database_url, + ) + try: + config: Final = _chat_proxy_config(rig.provider.url, []) + config_path: Final = tmp_path / "persisted-invalid-stream-scope.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy_process(rig.candidate, tmp_path, {}, config=config_path, workers=1) as owned: + info: Final = eventually( + lambda: owned.gateway.request("GET", f"/guardrails/{guardrail_id}/info"), + lambda response: response.status_code != 404, + seconds=20, + ) + listed: Final = owned.gateway.request("GET", "/v2/guardrails/list") + list_payload: Final = JSON_OBJECT.validate_json(listed.content) if listed.status_code == 200 else {} + listed_guardrails: Final = list_payload.get("guardrails") + includes_row: Final = isinstance(listed_guardrails, list) and any( + isinstance(row, dict) and row.get("guardrail_id") == guardrail_id for row in listed_guardrails + ) + markers: Final = ( + f"scope-invalid-persisted-0-{uuid.uuid4().hex}", + f"scope-invalid-persisted-1-{uuid.uuid4().hex}", + ) + observations: Final = tuple( + _chat_request_with_scans( + owned.gateway, + rig.sink, + "scope-invalid-config-chat", + marker, + streamed, + guardrail_name, + ) + for streamed, marker in zip((False, True), markers) + ) + assert ( + info.status_code == 200 + and listed.status_code == 200 + and includes_row + and tuple(response.status_code for response, _ in observations) == (200, 200) + and tuple(len(sink_rows) for _, sink_rows in observations) == (1, 1) + ), { + "info": (info.status_code, info.text), + "list": (listed.status_code, listed.text), + "includes_row": includes_row, + "responses": tuple((response.status_code, response.text) for response, _ in observations), + "scan_counts": tuple(len(sink_rows) for _, sink_rows in observations), + } + finally: + write_rows( + 'DELETE FROM "LiteLLM_GuardrailsTable" WHERE guardrail_id=%s', + (guardrail_id,), + database_url=database_url, + ) + + +def test_management_rejects_invalid_stream_scope(rig: ReproRig) -> None: + name: Final = f"scope-invalid-management-{uuid.uuid4().hex}" + invalid_params: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": False, + "api_base": f"{rig.sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + "stream_scope": "sometimes", + } + created_invalid: Final = rig.candidate.request( + "POST", + "/guardrails", + {"guardrail": {"guardrail_name": name, "litellm_params": invalid_params}}, + ) + assert created_invalid.status_code == 422, created_invalid.text + + valid_params: Final = {**invalid_params, "stream_scope": "both"} + created: Final = rig.candidate.request( + "POST", + "/guardrails", + {"guardrail": {"guardrail_name": name, "litellm_params": valid_params}}, + ) + assert created.status_code == 200, created.text + guardrail_id: Final = str(created.json()["guardrail_id"]) + try: + put_response: Final = rig.candidate.request( + "PUT", + f"/guardrails/{guardrail_id}", + {"guardrail": {"guardrail_name": name, "litellm_params": invalid_params}}, + ) + patch_response: Final = rig.candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"stream_scope": "sometimes"}}, + ) + assert put_response.status_code == 422, put_response.text + assert patch_response.status_code == 422, patch_response.text + finally: + deleted: Final = rig.candidate.request("DELETE", f"/guardrails/{guardrail_id}") + assert deleted.status_code == 200, deleted.text diff --git a/tests/integration/observability/test_guardrail_stream_scope_chaos.py b/tests/integration/observability/test_guardrail_stream_scope_chaos.py new file mode 100644 index 00000000000..f8a9d0a4bba --- /dev/null +++ b/tests/integration/observability/test_guardrail_stream_scope_chaos.py @@ -0,0 +1,1079 @@ +from __future__ import annotations + +import json +import os +import shutil +import signal +import socket +import subprocess +import uuid +from collections.abc import Iterator, Mapping, Sequence +from concurrent.futures import Future, ThreadPoolExecutor +from contextlib import ExitStack, contextmanager +from dataclasses import dataclass +from itertools import chain +from pathlib import Path +from threading import Barrier, Event +from types import MappingProxyType +from typing import Final, Literal, TypeAlias, cast + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import Gateway, JsonValue, eventually +from integration._support.database import read_rows +from integration._support.process import OwnedProxy, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import TypeAdapter + +ChaosEndpoint: TypeAlias = Literal["chat", "messages", "responses"] +CHAOS_ENDPOINTS: Final[tuple[ChaosEndpoint, ...]] = ("chat", "messages", "responses") +CHAOS_MODELS: Final = MappingProxyType( + {"chat": "chaos-chat", "messages": "chaos-messages", "responses": "chaos-responses"} +) +GUARDRAIL_PATH: Final = "/beta/litellm_basic_guardrail_api" +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +POSTGRES_IMAGE: Final = "postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5" + + +def _json(value: object) -> bytes: + return json.dumps(value, separators=(",", ":")).encode() + + +def _texts(value: JsonValue) -> tuple[str, ...]: + if isinstance(value, str): + return (value,) + if isinstance(value, list): + return tuple(chain.from_iterable(_texts(item) for item in value)) + if isinstance(value, dict): + return tuple(chain.from_iterable(_texts(item) for item in value.values())) + return () + + +def _marker(body: Mapping[str, JsonValue]) -> str: + return next((text for text in _texts(dict(body)) if text.startswith("audit-")), "audit-chaos") + + +def _sse(events: Sequence[Mapping[str, JsonValue]]) -> tuple[bytes, ...]: + return tuple(f"data: {json.dumps(event, separators=(',', ':'))}\n\n".encode() for event in events) + ( + b"data: [DONE]\n\n", + ) + + +def _messages_stream(message: Mapping[str, JsonValue]) -> tuple[bytes, ...]: + content: Final = cast(list[JsonValue], message["content"]) + text: Final = cast(dict[str, JsonValue], content[0])["text"] + assert isinstance(text, str) + return ( + f"event: message_start\ndata: {json.dumps({**message, 'content': [], 'stop_reason': None, 'usage': {'input_tokens': 2, 'output_tokens': 0}})}\n\n".encode(), + b'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}\n\n', + f"event: content_block_delta\ndata: {json.dumps({'type': 'content_block_delta', 'index': 0, 'delta': {'type': 'text_delta', 'text': text}})}\n\n".encode(), + b'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n', + f"event: message_delta\ndata: {json.dumps({'type': 'message_delta', 'delta': {'stop_reason': 'end_turn', 'stop_sequence': None}, 'usage': {'output_tokens': 2}})}\n\n".encode(), + b'event: message_stop\ndata: {"type":"message_stop"}\n\n', + ) + + +def _responses_stream( + response: Mapping[str, JsonValue], + output: Mapping[str, JsonValue], + marker: str, +) -> tuple[bytes, ...]: + events: Final[tuple[dict[str, JsonValue], ...]] = ( + {"type": "response.created", "response": {**response, "status": "in_progress", "output": []}}, + {"type": "response.in_progress", "response": {**response, "status": "in_progress", "output": []}}, + {"type": "response.output_item.added", "item": dict(output), "output_index": 0}, + { + "type": "response.content_part.added", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "part": {"type": "output_text", "text": "", "annotations": []}, + }, + { + "type": "response.output_text.delta", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "delta": marker, + }, + { + "type": "response.output_text.done", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "text": marker, + }, + { + "type": "response.content_part.done", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "part": cast(list[JsonValue], output["content"])[0], + }, + {"type": "response.output_item.done", "item": dict(output), "output_index": 0}, + {"type": "response.completed", "response": dict(response)}, + ) + return tuple( + f"event: {event['type']}\ndata: {json.dumps({**event, 'sequence_number': index}, separators=(',', ':'))}\n\n".encode() + for index, event in enumerate(events) + ) + + +def _provider(request: Request) -> Reply: + if request.method == "GET" and request.target.partition("?")[0] == "/v1/models": + return Reply(body=_json({"object": "list", "data": [{"id": "gpt-4o-mini", "object": "model"}]})) + if not request.body: + return Reply(status=400, body=_json({"error": "request body is required"})) + body: Final = JSON_OBJECT.validate_json(request.body) + marker: Final = _marker(body) + streamed: Final = bool(body.get("stream")) + if request.target == "/v1/chat/completions": + if streamed: + return Reply( + content_type="text/event-stream", + chunks=_sse( + ( + { + "id": f"chatcmpl-{marker}", + "object": "chat.completion.chunk", + "choices": [{"index": 0, "delta": {"content": marker}, "finish_reason": None}], + }, + ) + ), + ) + return Reply( + body=_json( + { + "id": f"chatcmpl-{marker}", + "object": "chat.completion", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": marker}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 2, "completion_tokens": 2, "total_tokens": 4}, + } + ) + ) + if request.target == "/v1/messages": + message: Final = { + "id": f"msg-{marker}", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20250929", + "content": [{"type": "text", "text": marker}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 2, "output_tokens": 2}, + } + return ( + Reply( + content_type="text/event-stream", + chunks=_messages_stream(message), + ) + if streamed + else Reply(body=_json(message)) + ) + if request.target == "/v1/responses": + response: Final = { + "id": f"resp-{marker}", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": f"msg-{marker}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": marker, "annotations": []}], + } + ], + "usage": {"input_tokens": 2, "output_tokens": 2, "total_tokens": 4}, + } + return ( + Reply( + content_type="text/event-stream", + chunks=_responses_stream(response, cast(dict[str, JsonValue], response["output"][0]), marker), + ) + if streamed + else Reply(body=_json(response)) + ) + return Reply(status=404, body=_json({"error": f"unexpected provider target {request.target}"})) + + +def _sink(request: Request) -> Reply: + assert request.target.endswith(GUARDRAIL_PATH), request.target + body: Final = JSON_OBJECT.validate_json(request.body) + assert body.get("litellm_call_id") is not None or any("audit-" in text for text in _texts(body)), ( + request.body.decode() + ) + return Reply(body=_json({"action": "NONE"})) + + +def _rail( + name: str, + sink_url: str, + scope: Literal["streaming", "non_streaming"], + *, + default_on: bool = False, +) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": default_on, + "stream_scope": scope, + "api_base": f"{sink_url}/{name}", + "api_key": "synthetic-chaos-key", + }, + } + + +def _config(provider_url: str, rails: Sequence[dict[str, JsonValue]]) -> dict[str, JsonValue]: + base: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + return cast( + dict[str, JsonValue], + { + **base, + "guardrails": list(rails), + "model_list": [ + { + "model_name": CHAOS_MODELS["chat"], + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider_url}/v1", + "api_key": "synthetic-provider-key", + "num_retries": 0, + }, + }, + { + "model_name": CHAOS_MODELS["messages"], + "litellm_params": { + "model": "anthropic/claude-sonnet-4-5-20250929", + "api_base": provider_url, + "api_key": "synthetic-provider-key", + "num_retries": 0, + }, + }, + { + "model_name": CHAOS_MODELS["responses"], + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider_url}/v1", + "api_key": "synthetic-provider-key", + "num_retries": 0, + }, + }, + ], + "environment_variables": { + **base.get("environment_variables", {}), + "OPENAI_API_BASE": provider_url, + "OPENAI_API_KEY": "synthetic-provider-key", + "ANTHROPIC_API_BASE": provider_url, + "ANTHROPIC_API_KEY": "synthetic-provider-key", + }, + }, + ) + + +@dataclass(frozen=True, slots=True) +class ChaosRig: + gateway: Gateway + provider: Wire + directory: Path + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[ChaosRig]: + with ( + httpx.Client( + base_url=os.environ["INTEGRATION_PROXY_URL"], + timeout=30, + trust_env=False, + ) as root_client, + wire_server(_provider) as provider, + ): + root_gateway: Final = Gateway( + root_client, + os.environ.get("INTEGRATION_MASTER_KEY", "sk-integration-master"), + os.environ["INTEGRATION_UPSTREAM_URL"], + ) + yield ChaosRig(root_gateway, provider, tmp_path_factory.mktemp("stream-scope-chaos")) + + +@dataclass(frozen=True, slots=True) +class CallPlan: + marker: str + call_id: str + endpoint: ChaosEndpoint + streamed: bool + + +def _plans(prefix: str, count: int) -> tuple[CallPlan, ...]: + return tuple( + CallPlan( + f"audit-{prefix}-{index}-{uuid.uuid4().hex}", + f"{prefix}-{uuid.uuid4().hex}", + CHAOS_ENDPOINTS[index % len(CHAOS_ENDPOINTS)], + index % 2 == 1, + ) + for index in range(count) + ) + + +def _request(gateway: Gateway, plan: CallPlan, rails: Sequence[str]) -> httpx.Response: + guardrail_field: Final[dict[str, JsonValue]] = {"guardrails": list(rails)} if rails else {} + match plan.endpoint: + case "chat": + return gateway.request( + "POST", + "/v1/chat/completions", + { + "model": CHAOS_MODELS["chat"], + "messages": [{"role": "user", "content": plan.marker}], + **guardrail_field, + **({"stream": True} if plan.streamed else {}), + }, + headers={"x-litellm-call-id": plan.call_id}, + ) + case "messages": + return gateway.request( + "POST", + "/v1/messages", + { + "model": CHAOS_MODELS["messages"], + "max_tokens": 32, + "messages": [{"role": "user", "content": plan.marker}], + **guardrail_field, + **({"stream": True} if plan.streamed else {}), + }, + headers={"x-litellm-call-id": plan.call_id}, + ) + case "responses": + return gateway.request( + "POST", + "/v1/responses", + { + "model": CHAOS_MODELS["responses"], + "input": plan.marker, + **guardrail_field, + **({"stream": True} if plan.streamed else {}), + }, + headers={"x-litellm-call-id": plan.call_id}, + ) + raise AssertionError(plan.endpoint) + + +def _rows_for_marker(rows: Sequence[Request], marker: str) -> tuple[Request, ...]: + return tuple(request for request in rows if marker.encode() in request.body) + + +def _rail_scans(rows: Sequence[Request], rail_name: str, marker: str) -> tuple[Request, ...]: + return tuple( + request for request in rows if request.target.startswith(f"/{rail_name}/") and marker.encode() in request.body + ) + + +def _spend_rows(call_id: str, database_url: str | None = None) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT request_id, litellm_call_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s OR litellm_call_id=%s', + (call_id, call_id), + database_url=database_url, + ) + + +def _one_spend_row( + call_id: str, + database_url: str | None = None, + *, + seconds: float = 70, +) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: _spend_rows(call_id, database_url), + lambda values: len(values) >= 1, + seconds=seconds, + ) + assert len(rows) == 1, (call_id, rows) + return rows[0] + + +def _expected_in_scope(plan: CallPlan, sink_a_name: str, sink_b_name: str) -> tuple[str, str]: + return (sink_a_name, sink_b_name) if plan.streamed else (sink_b_name, sink_a_name) + + +def _assert_successful_calls( + plans: Sequence[CallPlan], + responses: Sequence[httpx.Response], + provider_rows: Sequence[Request], + sink_a_rows: Sequence[Request], + sink_b_rows: Sequence[Request], + sink_a_name: str, + sink_b_name: str, +) -> None: + for plan, response in zip(plans, responses): + if response.status_code != 200: + continue + assert plan.marker in response.text, (plan, response.text) + provider_match: Final = _rows_for_marker(provider_rows, plan.marker) + assert len(provider_match) == 1, (plan, provider_match) + in_sink, out_sink = _expected_in_scope(plan, sink_a_name, sink_b_name) + in_rows: Final = _rail_scans( + sink_a_rows if in_sink == sink_a_name else sink_b_rows, + in_sink, + plan.marker, + ) + out_rows: Final = _rail_scans( + sink_a_rows if out_sink == sink_a_name else sink_b_rows, + out_sink, + plan.marker, + ) + assert len(in_rows) == 1 and len(out_rows) == 0, (plan, in_rows, out_rows) + spend: Final = _one_spend_row(plan.call_id) + assert plan.call_id in (spend.get("request_id"), spend.get("litellm_call_id")), (plan, spend) + + +def _call_wave(gateway: Gateway, plans: Sequence[CallPlan], rails: Sequence[str]) -> tuple[httpx.Response, ...]: + with ThreadPoolExecutor(max_workers=20) as pool: + futures: Final[tuple[Future[httpx.Response], ...]] = tuple( + pool.submit(_request, gateway, plan, rails) for plan in plans + ) + return tuple(future.result() for future in futures) + + +@contextmanager +def _owned_proxy( + rig: ChaosRig, + directory: Path, + rails: Sequence[dict[str, JsonValue]], + *, + workers: int = 1, +) -> Iterator[OwnedProxy]: + config_path: Final = directory / f"chaos-{uuid.uuid4().hex}.yaml" + config_path.write_text(yaml.safe_dump(_config(rig.provider.url, rails))) + with owned_proxy_process(rig.gateway, directory, {}, config=config_path, workers=workers) as owned: + yield owned + + +def _free_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return int(reserve.getsockname()[1]) + + +def _assert_down_wave( + plans: Sequence[CallPlan], + responses: Sequence[httpx.Response], + provider_rows: Sequence[Request], + sink_a_rows: Sequence[Request], + sink_b_rows: Sequence[Request], + sink_a_name: str, + sink_b_name: str, +) -> None: + for plan, response in zip(plans, responses): + if plan.streamed: + assert response.status_code >= 500 and response.content, (plan, response.status_code, response.text) + assert _rows_for_marker(provider_rows, plan.marker) == (), (plan, provider_rows) + assert _rail_scans(sink_b_rows, sink_b_name, plan.marker) == (), (plan, sink_b_name) + continue + assert response.status_code == 200 and plan.marker in response.text, (plan, response.status_code, response.text) + assert len(_rows_for_marker(provider_rows, plan.marker)) == 1, (plan, provider_rows) + assert _rail_scans(sink_a_rows, sink_a_name, plan.marker) == (), (plan, sink_a_name) + assert len(_rail_scans(sink_b_rows, sink_b_name, plan.marker)) == 1, (plan, sink_b_name) + spend: Final = _one_spend_row(plan.call_id) + assert plan.call_id in (spend.get("request_id"), spend.get("litellm_call_id")), (plan, spend) + + +def test_h1_sink_outage_keeps_scope_isolated_through_recovery(rig: ChaosRig, tmp_path: Path) -> None: + port_a: Final = _free_port() + started: Final = Event() + unavailable: Final = Event() + release: Final = Event() + + def gated_sink(request: Request) -> Reply: + started.set() + assert release.wait(timeout=45), "sink outage gate was not released" + if unavailable.is_set(): + return Reply(status=503, body=_json({"error": "synthetic sink outage"})) + return _sink(request) + + with wire_server(_sink) as sink_b, ExitStack() as sink_a_stack: + sink_a: Final = sink_a_stack.enter_context(wire_server(gated_sink, port=port_a)) + name_a: Final = f"h1-stream-{uuid.uuid4().hex}" + name_b: Final = f"h1-non-stream-{uuid.uuid4().hex}" + rails: Final = ( + _rail(name_a, sink_a.url, "streaming"), + _rail(name_b, sink_b.url, "non_streaming"), + ) + with _owned_proxy(rig, tmp_path, rails) as owned, ThreadPoolExecutor(max_workers=30) as pool: + outage_plans: Final = _plans("h1-burst", 30) + futures: Final[tuple[Future[httpx.Response], ...]] = tuple( + pool.submit(_request, owned.gateway, plan, (name_a, name_b)) for plan in outage_plans + ) + try: + assert started.wait(timeout=30), "streaming rail did not reach sink A" + unavailable.set() + finally: + release.set() + sink_a_stack.close() + outage_responses: Final = tuple(future.result(timeout=70) for future in futures) + outage_provider: Final = rig.provider.drain() + outage_a: Final = sink_a.drain() + outage_b: Final = sink_b.drain() + _assert_down_wave(outage_plans, outage_responses, outage_provider, outage_a, outage_b, name_a, name_b) + + with wire_server(_sink, port=port_a) as recovered_sink_a: + recovery_plans: Final = _plans("h1-recovery", 20) + recovery_responses: Final = _call_wave(owned.gateway, recovery_plans, (name_a, name_b)) + recovery_provider: Final = rig.provider.drain() + recovery_a: Final = recovered_sink_a.drain() + recovery_b: Final = sink_b.drain() + assert tuple(response.status_code for response in recovery_responses) == (200,) * 20, recovery_responses + _assert_successful_calls( + recovery_plans, + recovery_responses, + recovery_provider, + recovery_a, + recovery_b, + name_a, + name_b, + ) + + +def test_h2_stream_sink_gate_does_not_block_out_of_scope_calls(rig: ChaosRig, tmp_path: Path) -> None: + started: Final = Event() + release: Final = Event() + blocked_marker: Final = f"audit-h2-stream-{uuid.uuid4().hex}" + + def gated_sink(request: Request) -> Reply: + if blocked_marker.encode() in request.body: + started.set() + assert release.wait(timeout=45), "stream sink gate was not released" + return _sink(request) + + with wire_server(gated_sink) as sink_a, wire_server(_sink) as sink_b: + name_a: Final = f"h2-stream-{uuid.uuid4().hex}" + name_b: Final = f"h2-non-stream-{uuid.uuid4().hex}" + rails: Final = (_rail(name_a, sink_a.url, "streaming"), _rail(name_b, sink_b.url, "non_streaming")) + with _owned_proxy(rig, tmp_path, rails) as owned, ThreadPoolExecutor(max_workers=2) as pool: + streaming_plan: Final = CallPlan(blocked_marker, f"h2-stream-{uuid.uuid4().hex}", "chat", True) + non_streaming_plan: Final = CallPlan( + f"audit-h2-non-stream-{uuid.uuid4().hex}", + f"h2-non-stream-{uuid.uuid4().hex}", + "messages", + False, + ) + streaming_future: Final = pool.submit(_request, owned.gateway, streaming_plan, (name_a, name_b)) + assert started.wait(timeout=30), "stream request did not reach the gated sink" + try: + non_streaming_future: Final = pool.submit( + _request, + owned.gateway, + non_streaming_plan, + (name_a, name_b), + ) + non_streaming_response: Final = non_streaming_future.result(timeout=15) + assert non_streaming_response.status_code == 200, non_streaming_response.text + assert non_streaming_plan.marker in non_streaming_response.text, non_streaming_response.text + finally: + release.set() + streaming_response: Final = streaming_future.result(timeout=30) + assert streaming_response.status_code == 200, streaming_response.text + assert streaming_plan.marker in streaming_response.text, streaming_response.text + provider_rows: Final = rig.provider.drain() + sink_a_rows: Final = sink_a.drain() + sink_b_rows: Final = sink_b.drain() + _assert_successful_calls( + (streaming_plan, non_streaming_plan), + (streaming_response, non_streaming_response), + provider_rows, + sink_a_rows, + sink_b_rows, + name_a, + name_b, + ) + + +def _worker_processes(owned: OwnedProxy) -> tuple[psutil.Process, ...]: + return tuple( + child for child in psutil.Process(owned.process.pid).children() if "spawn_main" in _process_command(child) + ) + + +def _process_command(process: psutil.Process) -> str: + try: + return " ".join(process.cmdline()) + except psutil.Error: + return "" + + +def _safe_response(future: Future[httpx.Response]) -> httpx.Response | None: + try: + return future.result(timeout=70) + except (httpx.HTTPError, TimeoutError): + return None + + +def test_h3_worker_kill_mid_burst_keeps_remaining_worker_serving(rig: ChaosRig, tmp_path: Path) -> None: + plans: Final = _plans("h3-burst", 30) + gate_markers: Final = frozenset(plan.marker for plan in plans if plan.streamed) + started: Final = Event() + release: Final = Event() + + def gated_sink(request: Request) -> Reply: + if any(marker.encode() in request.body for marker in gate_markers): + started.set() + assert release.wait(timeout=60), "worker-kill sink gate was not released" + return _sink(request) + + with wire_server(gated_sink) as sink_a, wire_server(_sink) as sink_b: + name_a: Final = f"h3-stream-{uuid.uuid4().hex}" + name_b: Final = f"h3-non-stream-{uuid.uuid4().hex}" + rails: Final = ( + _rail(name_a, sink_a.url, "streaming", default_on=True), + _rail(name_b, sink_b.url, "non_streaming", default_on=True), + ) + with _owned_proxy(rig, tmp_path, rails, workers=2) as owned, ThreadPoolExecutor(max_workers=30) as pool: + workers: Final = _worker_processes(owned) + assert len(workers) == 2, tuple(worker.pid for worker in workers) + futures: Final[tuple[Future[httpx.Response], ...]] = tuple( + pool.submit(_request, owned.gateway, plan, ()) for plan in plans + ) + assert started.wait(timeout=30), "stream requests did not reach the owned sink" + victim: Final = workers[0] + survivor: Final = workers[1] + survivor_plan: Final = CallPlan( + f"audit-h3-survivor-{uuid.uuid4().hex}", + f"h3-survivor-{uuid.uuid4().hex}", + "messages", + False, + ) + try: + victim.send_signal(signal.SIGKILL) + survivor_response: Final = _request(owned.gateway, survivor_plan, ()) + assert survivor_response.status_code == 200 and survivor_plan.marker in survivor_response.text, ( + survivor_response.status_code, + survivor_response.text, + ) + assert survivor.is_running(), survivor.pid + finally: + release.set() + responses: Final = tuple(_safe_response(future) for future in futures) + assert owned.gateway.request("GET", "/health/liveliness").status_code == 200 + successful: Final = tuple( + (plan, response) + for plan, response in zip(plans, responses) + if response is not None and response.status_code == 200 + ) + provider_rows: Final = rig.provider.drain() + sink_a_rows: Final = sink_a.drain() + sink_b_rows: Final = sink_b.drain() + successful_plans: Final = (*tuple(plan for plan, _ in successful), survivor_plan) + successful_responses: Final = (*tuple(response for _, response in successful), survivor_response) + _assert_successful_calls( + successful_plans, + successful_responses, + provider_rows, + sink_a_rows, + sink_b_rows, + name_a, + name_b, + ) + + +def _create_stored_rail( + gateway: Gateway, + name: str, + sink_url: str, + stream_scope: Literal["streaming"] | None = "streaming", +) -> str: + response: Final = gateway.request( + "POST", + "/guardrails", + { + "guardrail": { + "guardrail_name": name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": False, + "api_base": f"{sink_url}/{name}", + "api_key": "synthetic-chaos-key", + **({"stream_scope": stream_scope} if stream_scope is not None else {}), + }, + } + }, + ) + assert response.status_code == 200, response.text + identity: Final = JSON_OBJECT.validate_json(response.content).get("guardrail_id") + assert isinstance(identity, str), response.text + return identity + + +def _assert_one_scope_wave( + plans: Sequence[CallPlan], + responses: Sequence[httpx.Response], + provider_rows: Sequence[Request], + sink_rows: Sequence[Request], + sink_name: str, + *, + database_url: str | None = None, +) -> None: + for plan, response in zip(plans, responses): + assert response.status_code == 200 and plan.marker in response.text, (plan, response.status_code, response.text) + provider_match: Final = _rows_for_marker(provider_rows, plan.marker) + sink_match: Final = _rail_scans(sink_rows, sink_name, plan.marker) + assert len(provider_match) == 1, (plan, provider_match) + assert len(sink_match) == int(plan.streamed), (plan, sink_match) + spend: Final = _one_spend_row(plan.call_id, database_url) + assert plan.call_id in (spend.get("request_id"), spend.get("litellm_call_id")), (plan, spend) + + +def test_h4_stored_scope_survives_owned_proxy_restart(rig: ChaosRig, tmp_path: Path) -> None: + with wire_server(_sink) as sink: + name: Final = f"h4-stored-{uuid.uuid4().hex}" + identity: Final = _create_stored_rail(rig.gateway, name, sink.url) + try: + with ExitStack() as first_stack: + first: Final = first_stack.enter_context(_owned_proxy(rig, tmp_path, ())) + first_plans: Final = _plans("h4-before", 20) + first_responses: Final = _call_wave(first.gateway, first_plans, (name,)) + first_provider: Final = rig.provider.drain() + first_sink: Final = sink.drain() + assert tuple(response.status_code for response in first_responses) == (200,) * 20, first_responses + _assert_one_scope_wave(first_plans, first_responses, first_provider, first_sink, name) + first_stack.close() + with _owned_proxy(rig, tmp_path, ()) as restarted: + recovery_plans: Final = _plans("h4-after", 20) + recovery_responses: Final = _call_wave(restarted.gateway, recovery_plans, (name,)) + recovery_provider: Final = rig.provider.drain() + recovery_sink: Final = sink.drain() + assert tuple(response.status_code for response in recovery_responses) == (200,) * 20, ( + recovery_responses, + ) + _assert_one_scope_wave( + recovery_plans, + recovery_responses, + recovery_provider, + recovery_sink, + name, + ) + finally: + deleted: Final = rig.gateway.request("DELETE", f"/guardrails/{identity}") + assert deleted.status_code == 200, deleted.text + + +@contextmanager +def _owned_postgres(directory: Path) -> Iterator[PostgresCluster]: + docker: Final = shutil.which("docker") + assert docker is not None, "Docker CLI is required for the H5 PostgreSQL outage test" + docker_info: Final = subprocess.run( + [docker, "info", "--format", "{{.ServerVersion}}"], + capture_output=True, + text=True, + check=False, + ) + assert docker_info.returncode == 0, docker_info.stderr + directory.mkdir(parents=True, exist_ok=True) + port: Final = _free_port() + container_name: Final = f"litellm-stream-scope-h5-{uuid.uuid4().hex}" + password: Final = uuid.uuid4().hex + created: Final = subprocess.run( + [ + docker, + "create", + "--name", + container_name, + "--env", + "POSTGRES_USER=postgres", + "--env", + "POSTGRES_PASSWORD", + "--env", + "POSTGRES_DB=postgres", + "--publish", + f"127.0.0.1:{port}:5432/tcp", + POSTGRES_IMAGE, + ], + env=os.environ | {"POSTGRES_PASSWORD": password}, + capture_output=True, + text=True, + check=False, + ) + assert created.returncode == 0, created.stderr + cluster: Final = PostgresCluster( + f"postgresql://postgres:{password}@127.0.0.1:{port}/postgres?sslmode=disable", + container_name, + docker, + directory / "postgres.log", + ) + try: + started: Final = _start_postgres(cluster) + assert started.returncode == 0, started.stderr + assert eventually(lambda: _postgres_is_ready(cluster), bool, seconds=70) + yield cluster + finally: + logs: Final = subprocess.run( + [docker, "logs", container_name], + capture_output=True, + text=True, + check=False, + ) + cluster.log_path.write_text(logs.stdout + logs.stderr) + removed: Final = subprocess.run( + [docker, "rm", "-f", container_name], + capture_output=True, + text=True, + check=False, + ) + assert removed.returncode == 0, removed.stderr + remaining: Final = subprocess.run( + [docker, "ps", "--all", "--quiet", "--filter", f"name={container_name}"], + capture_output=True, + text=True, + check=False, + ) + assert remaining.returncode == 0, remaining.stderr + assert not remaining.stdout.strip(), remaining.stdout + + +@dataclass(frozen=True, slots=True) +class PostgresCluster: + database_url: str + container_name: str + docker: str + log_path: Path + + +def _start_postgres(cluster: PostgresCluster) -> subprocess.CompletedProcess[str]: + return subprocess.run( + [cluster.docker, "start", cluster.container_name], + capture_output=True, + text=True, + check=False, + ) + + +def _postgres_is_ready(cluster: PostgresCluster) -> bool: + readiness: Final = subprocess.run( + [ + cluster.docker, + "exec", + cluster.container_name, + "pg_isready", + "-h", + "127.0.0.1", + "-p", + "5432", + "-d", + "postgres", + "-U", + "postgres", + ], + capture_output=True, + text=True, + check=False, + ) + return readiness.returncode == 0 + + +def _postgres_is_running(cluster: PostgresCluster) -> bool: + state: Final = subprocess.run( + [cluster.docker, "inspect", "--format", "{{.State.Running}}", cluster.container_name], + capture_output=True, + text=True, + check=False, + ) + return state.returncode == 0 and state.stdout.strip() == "true" + + +@pytest.mark.timeout(180) +def test_h5_stored_scope_survives_owned_postgres_outage_and_recovers_once( + rig: ChaosRig, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + postgres_directory: Final = Path(os.environ["INTEGRATION_RESULTS_DIR"]) / f"owned-postgres-{uuid.uuid4().hex}" + with _owned_postgres(postgres_directory) as database, wire_server(_sink) as sink: + monkeypatch.setenv("INTEGRATION_PROXY_DATABASE_URL", database.database_url) + name: Final = f"h5-stored-{uuid.uuid4().hex}" + with _owned_proxy(rig, tmp_path, ()) as registrar: + _create_stored_rail(registrar.gateway, name, sink.url) + with _owned_proxy(rig, tmp_path, ()) as owned: + try: + preflight_plans: Final = ( + CallPlan( + f"audit-h5-preflight-{uuid.uuid4().hex}", + f"h5-preflight-{uuid.uuid4().hex}", + "chat", + True, + ), + ) + preflight_responses: Final = _call_wave(owned.gateway, preflight_plans, (name,)) + preflight_provider: Final = rig.provider.drain() + preflight_sink: Final = sink.drain() + assert tuple(response.status_code for response in preflight_responses) == (200,), (preflight_responses,) + _assert_one_scope_wave( + preflight_plans, + preflight_responses, + preflight_provider, + preflight_sink, + name, + database_url=database.database_url, + ) + outage_result: Final = subprocess.run( + [database.docker, "stop", database.container_name], + capture_output=True, + text=True, + check=False, + ) + assert outage_result.returncode == 0, outage_result.stderr + outage_plans: Final = _plans("h5-outage", 20) + outage_responses: Final = _call_wave(owned.gateway, outage_plans, (name,)) + outage_provider: Final = rig.provider.drain() + outage_sink: Final = sink.drain() + for plan, response in zip(outage_plans, outage_responses): + assert response.status_code == 200 and plan.marker in response.text, ( + plan, + response.status_code, + response.text, + ) + assert len(_rows_for_marker(outage_provider, plan.marker)) == 1, (plan, outage_provider) + sink_match: Final = _rail_scans(outage_sink, name, plan.marker) + assert len(sink_match) == int(plan.streamed), (plan, sink_match) + + recovered_database: Final = _start_postgres(database) + assert recovered_database.returncode == 0, recovered_database.stderr + postgres_ready: Final = eventually(lambda: _postgres_is_ready(database), bool, seconds=70) + assert postgres_ready + readiness: Final = eventually( + lambda: owned.gateway.client.get("/health/readiness"), + lambda response: response.status_code == 200 and response.json().get("db") == "connected", + seconds=70, + ) + assert readiness.json().get("db") == "connected", readiness.text + recovery_plans: Final = _plans("h5-recovery", 20) + recovery_responses: Final = _call_wave(owned.gateway, recovery_plans, (name,)) + recovery_provider: Final = rig.provider.drain() + recovery_sink: Final = sink.drain() + assert tuple(response.status_code for response in recovery_responses) == (200,) * 20, recovery_responses + _assert_one_scope_wave( + recovery_plans, + recovery_responses, + recovery_provider, + recovery_sink, + name, + database_url=database.database_url, + ) + finally: + if not _postgres_is_running(database): + restarted: Final = _start_postgres(database) + assert restarted.returncode == 0, restarted.stderr + assert eventually(lambda: _postgres_is_ready(database), bool, seconds=70) + + +@pytest.mark.timeout(360) +def test_h6_spend_rows_for_requests_served_during_postgres_restart_land_once( + rig: ChaosRig, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + pytest.skip("BUG: LIT-9050 spend rows for requests served during a Postgres restart are dropped") + postgres_directory: Final = Path(os.environ["INTEGRATION_RESULTS_DIR"]) / f"owned-postgres-{uuid.uuid4().hex}" + with _owned_postgres(postgres_directory) as database, wire_server(_sink) as sink: + monkeypatch.setenv("INTEGRATION_PROXY_DATABASE_URL", database.database_url) + name: Final = f"h6-stored-{uuid.uuid4().hex}" + plans: Final = _plans("h6-postgres-restart", 20) + outage_markers: Final = frozenset(plan.marker for plan in plans[:10]) + arrivals: Final = Barrier(len(plans) + 1) + during_outage: Final = Event() + after_restart: Final = Event() + + def _gated_provider(request: Request) -> Reply: + if request.method == "GET" and request.target.partition("?")[0] == "/v1/models": + return _provider(request) + body: Final = JSON_OBJECT.validate_json(request.body) + marker: Final = _marker(body) + arrivals.wait(timeout=70) + gate: Final = during_outage if marker in outage_markers else after_restart + assert gate.wait(timeout=70), marker + return _provider(request) + + with wire_server(_gated_provider) as provider: + h6_rig: Final = ChaosRig(rig.gateway, provider, rig.directory) + with _owned_proxy(h6_rig, tmp_path, ()) as registrar: + _create_stored_rail(registrar.gateway, name, sink.url, stream_scope=None) + with _owned_proxy(h6_rig, tmp_path, ()) as owned, ThreadPoolExecutor(max_workers=len(plans)) as pool: + futures: Final = tuple( + pool.submit(_request, owned.gateway, plan, (name,)) for plan in plans + ) + try: + arrivals.wait(timeout=70) + stopped: Final = subprocess.run( + [database.docker, "stop", database.container_name], + capture_output=True, + text=True, + check=False, + ) + assert stopped.returncode == 0, stopped.stderr + during_outage.set() + outage_responses: Final = tuple(future.result(timeout=70) for future in futures[:10]) + for plan, response in zip(plans[:10], outage_responses): + assert response.status_code == 200 and plan.marker in response.text, ( + plan, + response.status_code, + response.text, + ) + + restarted: Final = _start_postgres(database) + assert restarted.returncode == 0, restarted.stderr + postgres_ready: Final = eventually(lambda: _postgres_is_ready(database), bool, seconds=70) + assert postgres_ready + readiness: Final = eventually( + lambda: owned.gateway.client.get("/health/readiness"), + lambda response: response.status_code == 200 + and JSON_OBJECT.validate_python(cast(object, response.json())).get("db") == "connected", + seconds=70, + ) + readiness_body: Final = JSON_OBJECT.validate_python(cast(object, readiness.json())) + assert readiness_body.get("db") == "connected", readiness.text + after_restart.set() + responses: Final = tuple(future.result(timeout=70) for future in futures) + finally: + during_outage.set() + after_restart.set() + if not _postgres_is_running(database): + recovered: Final = _start_postgres(database) + assert recovered.returncode == 0, recovered.stderr + assert eventually(lambda: _postgres_is_ready(database), bool, seconds=70) + + provider_rows: Final = provider.drain() + sink_rows: Final = sink.drain() + for plan, response in zip(plans, responses): + assert response.status_code == 200 and plan.marker in response.text, ( + plan, + response.status_code, + response.text, + ) + assert len(_rows_for_marker(provider_rows, plan.marker)) == 1, (plan, provider_rows) + assert len(_rail_scans(sink_rows, name, plan.marker)) == 1, (plan, sink_rows) + + spend_rows: Final = eventually( + lambda: tuple(_spend_rows(plan.call_id, database.database_url) for plan in plans), + lambda values: all(len(rows) == 1 for rows in values), + seconds=70, + return_last_on_timeout=True, + ) + counts: Final = tuple(len(rows) for rows in spend_rows) + missing_ids: Final = tuple( + plan.call_id for plan, rows in zip(plans, spend_rows) if not rows + ) + duplicate_ids: Final = tuple( + plan.call_id for plan, rows in zip(plans, spend_rows) if len(rows) > 1 + ) + assert counts == (1,) * len(plans), { + "missing_ids": missing_ids, + "duplicate_ids": duplicate_ids, + "counts": counts, + } diff --git a/tests/integration/observability/test_guardrail_stream_scope_matrix.py b/tests/integration/observability/test_guardrail_stream_scope_matrix.py new file mode 100644 index 00000000000..b5c66b202cc --- /dev/null +++ b/tests/integration/observability/test_guardrail_stream_scope_matrix.py @@ -0,0 +1,2076 @@ +from __future__ import annotations + +import asyncio +import contextlib +import json +import os +import socket +import socketserver +import ssl +import threading +import uuid +from asyncio import run, wait_for +from collections.abc import Generator, Iterable, Iterator, Mapping, Sequence +from contextlib import contextmanager +from dataclasses import dataclass +from itertools import chain +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal, TypeAlias, cast + +import anthropic +import httpx +import openai +import pytest +import websockets +import yaml +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway, JsonValue, Scenario, eventually +from integration._support.database import read_rows +from integration._support.mcp import McpCaller, echo_tool, register_mcp, scripted_peer +from integration._support.process import owned_proxy_process +from integration._support.tls import server_context, write_self_signed_cert +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.cost_calculation.cost_tracking_case import RealtimeResponse +from pydantic import TypeAdapter +from websockets.asyncio.server import ServerConnection, serve + +Endpoint: TypeAlias = Literal["chat", "messages", "responses"] +Mode: TypeAlias = Literal["pre_call", "during_call", "post_call", "logging_only"] +Scope: TypeAlias = Literal["streaming", "non_streaming"] +MatrixScope: TypeAlias = Scope | None +LoggingPhase: TypeAlias = Literal["request", "response"] +Endpoints: Final[tuple[Endpoint, ...]] = ("chat", "messages", "responses") +Modes: Final[tuple[Mode, ...]] = ("pre_call", "during_call", "post_call", "logging_only") +Scopes: Final[tuple[MatrixScope, ...]] = ("streaming", "non_streaming", None) +PROVIDER_TEXT: Final = "provider stream_scope control" +GUARDRAIL_PATH: Final = "/beta/litellm_basic_guardrail_api" +F1_BLOCKED_WORD: Final = "streamscope-mcp-blocked-word" +F3_STREAM_BLOCKED_WORD: Final = "matrix-f3-stream-block" +F3_NON_STREAM_BLOCKED_WORD: Final = "matrix-f3-non-stream-block" +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +LOGGING_PHASE: Final = TypeAdapter(LoggingPhase) + + +def _json(value: object) -> bytes: + return json.dumps(value, separators=(",", ":")).encode() + + +def _strings(value: JsonValue) -> tuple[str, ...]: + if isinstance(value, str): + return (value,) + if isinstance(value, list): + return tuple(chain.from_iterable(_strings(item) for item in value)) + if isinstance(value, dict): + return tuple(chain.from_iterable(_strings(item) for item in value.values())) + return () + + +def _marker(body: Mapping[str, JsonValue]) -> str: + return next((text for text in _strings(dict(body)) if text.startswith("audit-")), "audit-provider") + + +def _sse(events: Iterable[Mapping[str, JsonValue]]) -> tuple[bytes, ...]: + return tuple(f"data: {json.dumps(event, separators=(',', ':'))}\n\n".encode() for event in events) + ( + b"data: [DONE]\n\n", + ) + + +def _messages_stream(message: Mapping[str, JsonValue]) -> tuple[bytes, ...]: + content: Final = cast(list[JsonValue], message["content"]) + text: Final = cast(dict[str, JsonValue], content[0])["text"] + return ( + f"event: message_start\ndata: {json.dumps({**message, 'content': [], 'stop_reason': None, 'usage': {'input_tokens': 11, 'output_tokens': 0}})}\n\n".encode(), + b'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}\n\n', + f"event: content_block_delta\ndata: {json.dumps({'type': 'content_block_delta', 'index': 0, 'delta': {'type': 'text_delta', 'text': text}})}\n\n".encode(), + b'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n', + f"event: message_delta\ndata: {json.dumps({'type': 'message_delta', 'delta': {'stop_reason': 'end_turn', 'stop_sequence': None}, 'usage': {'output_tokens': 4}})}\n\n".encode(), + b'event: message_stop\ndata: {"type":"message_stop"}\n\n', + ) + + +def _responses_stream( + response: Mapping[str, JsonValue], output: Mapping[str, JsonValue], marker: str +) -> tuple[bytes, ...]: + events: Final[tuple[dict[str, JsonValue], ...]] = ( + {"type": "response.created", "response": {**response, "status": "in_progress", "output": []}}, + {"type": "response.in_progress", "response": {**response, "status": "in_progress", "output": []}}, + {"type": "response.output_item.added", "item": output, "output_index": 0}, + { + "type": "response.content_part.added", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "part": {"type": "output_text", "text": "", "annotations": []}, + }, + { + "type": "response.output_text.delta", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "delta": f"{PROVIDER_TEXT} {marker}", + }, + { + "type": "response.output_text.done", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "text": f"{PROVIDER_TEXT} {marker}", + }, + { + "type": "response.content_part.done", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "part": cast(list[JsonValue], output["content"])[0], + }, + {"type": "response.output_item.done", "item": output, "output_index": 0}, + {"type": "response.completed", "response": response}, + ) + return tuple( + f"event: {event['type']}\ndata: {json.dumps({**event, 'sequence_number': index}, separators=(',', ':'))}\n\n".encode() + for index, event in enumerate(events) + ) + + +def _provider_reply(target: str, body: Mapping[str, JsonValue]) -> Reply: + marker: Final = _marker(body) + streamed: Final = bool(body.get("stream")) + if target == "/v1/chat/completions": + if streamed: + common: Final = {"id": f"chatcmpl-{marker}", "object": "chat.completion.chunk", "created": 1} + return Reply( + content_type="text/event-stream", + chunks=_sse( + ( + { + **common, + "choices": [{"index": 0, "delta": {"role": "assistant"}, "finish_reason": None}], + }, + { + **common, + "choices": [ + {"index": 0, "delta": {"content": f"{PROVIDER_TEXT} {marker}"}, "finish_reason": None} + ], + }, + {**common, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + ) + ), + ) + return Reply( + body=_json( + { + "id": f"chatcmpl-{marker}", + "object": "chat.completion", + "created": 1, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": f"{PROVIDER_TEXT} {marker}"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}, + } + ) + ) + if target == "/v1/messages": + message: Final = { + "id": f"msg-{marker}", + "type": "message", + "role": "assistant", + "model": "synthetic-anthropic-model", + "content": [{"type": "text", "text": f"{PROVIDER_TEXT} {marker}"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 4}, + } + return ( + Reply(content_type="text/event-stream", chunks=_messages_stream(message)) + if streamed + else Reply(body=_json(message)) + ) + if target == "/v1/responses": + output: Final = { + "type": "message", + "id": f"msg-{marker}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": f"{PROVIDER_TEXT} {marker}", "annotations": []}], + } + response: Final = { + "id": f"resp-{marker}", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [output], + "usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15}, + } + return ( + Reply(content_type="text/event-stream", chunks=_responses_stream(response, output, marker)) + if streamed + else Reply(body=_json(response)) + ) + if target.endswith(":generateContent"): + return Reply( + body=_json( + {"candidates": [{"content": {"role": "model", "parts": [{"text": f"{PROVIDER_TEXT} {marker}"}]}}]} + ) + ) + if target.endswith(":streamGenerateContent"): + return Reply( + content_type="text/event-stream", + chunks=( + f"data: {_json({'candidates': [{'content': {'role': 'model', 'parts': [{'text': f'{PROVIDER_TEXT} {marker}'}]}}]}).decode()}\n\n".encode(), + ), + ) + return Reply(status=404, body=_json({"error": {"message": f"unexpected target {target}"}})) + + +def _provider(request: Request) -> Reply: + if not request.body: + return Reply(status=400, body=_json({"error": "empty request body"})) + body: Final = JSON_OBJECT.validate_json(request.body) + if request.target.startswith("/passthrough"): + marker: Final = _marker(body) + if bool(body.get("stream")): + return Reply( + content_type="text/event-stream", + chunks=_sse(({"text": f"{PROVIDER_TEXT} {marker}"},)), + ) + return Reply(body=_json({"received": body})) + if "audit-g3-401-" in request.body.decode(): + return Reply(status=401, body=_json({"error": {"message": "synthetic provider unauthorized"}})) + target: Final = request.target.split("?", 1)[0] + provider_target: Final = "/v1/messages" if target.endswith("/anthropic/v1/messages") else target + return _provider_reply(provider_target, body) + + +def _sink(request: Request) -> Reply: + assert request.target.endswith(GUARDRAIL_PATH), request.target + JSON_OBJECT.validate_json(request.body) + if "/g1-" in request.target: + return Reply(status=500, body=_json({"error": "synthetic sink failure"})) + if "/g2-" in request.target or "/f2-block-" in request.target: + return Reply(body=_json({"action": "BLOCKED", "blocked_reason": "synthetic policy block"})) + return Reply(body=_json({"action": "NONE"})) + + +def _rail( + name: str, + sink: Wire, + *, + mode: str | list[str], + scope: str | Mapping[str, str] | None = None, + default_on: bool = False, +) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": mode, + "default_on": default_on, + **({"stream_scope": dict(scope)} if isinstance(scope, Mapping) else {}), + **({"stream_scope": scope} if isinstance(scope, str) else {}), + "api_base": f"{sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + }, + } + + +def _realtime_filter_rail(name: str, scope: Scope, keyword: str) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "litellm_content_filter", + "mode": "realtime_input_transcription", + "default_on": True, + "stream_scope": scope, + "blocked_words": [{"keyword": keyword, "action": "BLOCK"}], + }, + } + + +def _mcp_filter_rail(name: str, scope: MatrixScope) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "litellm_content_filter", + "mode": "pre_mcp_call", + "default_on": False, + **({"stream_scope": scope} if scope is not None else {}), + "blocked_words": [{"keyword": F1_BLOCKED_WORD, "action": "BLOCK"}], + }, + } + + +def _mode_scope_pairs() -> tuple[tuple[Mode, MatrixScope], ...]: + return tuple(chain.from_iterable(tuple((mode, scope) for scope in Scopes) for mode in Modes)) + + +def _rail_names() -> Mapping[str, str]: + return MappingProxyType( + { + **{f"a_{mode}_{scope or 'unset'}": f"a_{mode}_{scope or 'unset'}" for mode, scope in _mode_scope_pairs()}, + "c1_default": "c1_default", + "c2_key": "c2_key", + "c3_team": "c3_team", + "c4_modes": "c4_modes", + "c5_omitted": "c5_omitted", + "c6_both": "c6_both", + "c6_unset": "c6_unset", + "e_stream": "e_stream", + "e_non_stream": "e_non_stream", + "f1_stream": "f1_stream", + "f1_non_stream": "f1_non_stream", + "f1_unset": "f1_unset", + "f2_stream": "f2_stream", + "f2_non_stream": "f2-block-non-stream", + "f3_stream": "f3_stream", + "f3_non_stream": "f3_non_stream", + "g1_failure": "g1-failure", + "g2_block": "g2-block", + "g3_provider": "g3-provider", + "z_logging_only_barrier": "z_logging_only_barrier", + } + ) + + +def _configured_rails(names: Mapping[str, str], sink: Wire) -> tuple[dict[str, JsonValue], ...]: + return tuple( + _rail(names[f"a_{mode}_{scope or 'unset'}"], sink, mode=mode, scope=scope) + for mode, scope in _mode_scope_pairs() + ) + ( + _rail(names["c1_default"], sink, mode="post_call", scope="streaming"), + _rail(names["c2_key"], sink, mode="post_call", scope="streaming"), + _rail(names["c3_team"], sink, mode="post_call", scope="streaming"), + _rail( + names["c4_modes"], + sink, + mode=["pre_call", "post_call"], + scope={"pre_call": "non_streaming", "post_call": "streaming"}, + ), + _rail( + names["c5_omitted"], + sink, + mode=["pre_call", "post_call"], + scope={"post_call": "streaming"}, + ), + _rail(names["c6_both"], sink, mode="post_call", scope="both"), + _rail(names["c6_unset"], sink, mode="post_call"), + _rail(names["e_stream"], sink, mode="pre_call", scope="streaming"), + _rail(names["e_non_stream"], sink, mode="pre_call", scope="non_streaming"), + _mcp_filter_rail(names["f1_stream"], "streaming"), + _mcp_filter_rail(names["f1_non_stream"], "non_streaming"), + _mcp_filter_rail(names["f1_unset"], None), + _rail(names["f2_stream"], sink, mode="post_call", scope="streaming"), + _rail(names["f2_non_stream"], sink, mode="post_call", scope="non_streaming"), + _realtime_filter_rail(names["f3_stream"], "streaming", F3_STREAM_BLOCKED_WORD), + _realtime_filter_rail(names["f3_non_stream"], "non_streaming", F3_NON_STREAM_BLOCKED_WORD), + _rail(names["g1_failure"], sink, mode="pre_call", scope="streaming"), + _rail(names["g2_block"], sink, mode="pre_call", scope="streaming"), + _rail(names["g3_provider"], sink, mode="pre_call", scope="both"), + _rail(names["z_logging_only_barrier"], sink, mode="logging_only"), + ) + + +def _realtime_transcription_response(transcript: str) -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", + events=( + { + "type": "conversation.item.input_audio_transcription.completed", + "event_id": "evt_$REQUEST_ID", + "item_id": "item_$REQUEST_ID", + "content_index": 0, + "transcript": transcript, + }, + { + "type": "response.done", + "event_id": "evt_$REQUEST_ID", + "response": { + "id": "resp_$REQUEST_ID", + "object": "realtime.response", + "status": "completed", + "output": [], + "usage": { + "total_tokens": 0, + "input_tokens": 0, + "output_tokens": 0, + "input_token_details": { + "text_tokens": 0, + "audio_tokens": 0, + "cached_tokens": 0, + "cached_tokens_details": {"text_tokens": 0, "audio_tokens": 0}, + }, + "output_token_details": {"text_tokens": 0, "audio_tokens": 0}, + }, + }, + }, + ), + ) + + +async def _collect_realtime_events(websocket: websockets.ClientConnection) -> tuple[dict[str, JsonValue], ...]: + event: Final = JSON_OBJECT.validate_json(await websocket.recv()) + if event.get("type") == "response.done": + return (event,) + return (event, *await _collect_realtime_events(websocket)) + + +async def _realtime_transcription_events( + url: str, + key: str, + model: str, +) -> tuple[dict[str, JsonValue], ...]: + websocket_url: Final = ( + f"{url.replace('http://', 'ws://').replace('https://', 'wss://').rstrip('/')}/v1/realtime?model={model}" + ) + async with websockets.connect(websocket_url, additional_headers={"Authorization": f"Bearer {key}"}) as websocket: + session: Final = JSON_OBJECT.validate_json(await websocket.recv()) + assert session.get("type") == "session.created", session + await websocket.send(json.dumps({"type": "response.create"})) + return await _collect_realtime_events(websocket) + + +@dataclass(frozen=True, slots=True) +class MatrixRig: + candidate: Gateway + scenario: Scenario + models: Mapping[str, str] + rails: Mapping[str, str] + key: str + provider: Wire + sink: Wire + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[MatrixRig]: + with httpx.Client( + base_url=os.environ["INTEGRATION_PROXY_URL"], + timeout=30, + trust_env=False, + ) as root_client: + root_gateway: Final = Gateway( + root_client, + os.environ.get("INTEGRATION_MASTER_KEY", "sk-integration-master"), + os.environ["INTEGRATION_UPSTREAM_URL"], + ) + directory: Final = tmp_path_factory.mktemp("guardrail-stream-scope-matrix") + with wire_server(_provider) as provider, wire_server(_sink) as sink: + names: Final = _rail_names() + config_base: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + pass_through_paths: Final = ( + { + "path": "/pt", + "target": f"{provider.url}/passthrough", + "include_subpath": True, + "guardrails": {names["e_stream"]: None, names["e_non_stream"]: None}, + }, + { + "path": "/anthropic/v1/messages", + "target": f"{provider.url}/v1/messages", + "guardrails": {names["e_stream"]: None, names["e_non_stream"]: None}, + }, + { + "path": "/gemini/v1beta/models", + "target": "", + "include_subpath": True, + "guardrails": {names["e_stream"]: None, names["e_non_stream"]: None}, + }, + ) + config: Final = { + **config_base, + "guardrails": _configured_rails(names, sink), + "model_list": [ + { + "model_name": "matrix-pipeline-model", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider.url}/v1", + "api_key": "synthetic-provider-key", + }, + } + ], + "environment_variables": { + **config_base.get("environment_variables", {}), + "ANTHROPIC_API_BASE": provider.url, + "ANTHROPIC_API_KEY": "synthetic-anthropic-key", + "GEMINI_API_BASE": provider.url, + "GEMINI_API_KEY": "synthetic-gemini-key", + }, + "general_settings": { + **config_base["general_settings"], + "pass_through_endpoints": pass_through_paths, + }, + "policies": { + "matrix-f2": { + "guardrails": {"add": [names["f2_stream"], names["f2_non_stream"]]}, + "pipeline": { + "mode": "post_call", + "steps": [ + {"guardrail": names["f2_stream"], "on_pass": "next", "on_fail": "block"}, + {"guardrail": names["f2_non_stream"], "on_pass": "next", "on_fail": "block"}, + ], + }, + } + }, + "policy_attachments": [{"policy": "matrix-f2", "models": ["matrix-pipeline-model"]}], + } + config_path: Final = directory / "stream-scope-matrix.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy_process(root_gateway, directory, {}, config=config_path, workers=1) as owned: + with owned.gateway.scenario() as scenario: + models: Final = MappingProxyType( + { + "chat": scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{provider.url}/v1", + api_key="synthetic-provider-key", + ), + "messages": scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=provider.url, + api_key="synthetic-provider-key", + ), + "responses": scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{provider.url}/v1", + api_key="synthetic-provider-key", + ), + } + ) + key: Final = scenario.key(guardrails=[names["c2_key"]]) + yield MatrixRig(owned.gateway, scenario, models, names, key, provider, sink) + + +def _cell_body( + rig: MatrixRig, + endpoint: Endpoint, + marker: str, + streamed: bool, + rail: str | None, +) -> tuple[str, dict[str, JsonValue]]: + selected: Final = [] if rail is None else [rail] + if endpoint == "chat": + body: Final = { + "model": rig.models["chat"], + "messages": [{"role": "user", "content": marker}], + "guardrails": selected, + **({"stream": True} if streamed else {}), + } + return "/v1/chat/completions", body + if endpoint == "messages": + body = { + "model": rig.models["messages"], + "max_tokens": 32, + "messages": [{"role": "user", "content": marker}], + "guardrails": selected, + **({"stream": True} if streamed else {}), + } + return "/v1/messages", body + body = { + "model": rig.models["responses"], + "input": marker, + "guardrails": selected, + **({"stream": True} if streamed else {}), + } + return "/v1/responses", body + + +def _matching_requests(wire: Wire, marker: str) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if marker.encode() in request.body) + + +def _rail_scans(rows: Sequence[Request], rail_name: str, marker: str) -> tuple[Request, ...]: + return tuple( + request for request in rows if request.target.startswith(f"/{rail_name}/") and marker.encode() in request.body + ) + + +def _logging_phases(rows: Sequence[Request]) -> tuple[str, ...]: + return tuple( + sorted(LOGGING_PHASE.validate_python(JSON_OBJECT.validate_json(request.body)["input_type"]) for request in rows) + ) + + +@dataclass(slots=True) +class _SinkRowsAccumulator: + sink: Wire + rows: tuple[Request, ...] = () + + def drain(self) -> tuple[Request, ...]: + self.rows = (*self.rows, *self.sink.drain()) + return self.rows + + +def _logging_only_scans( + sink: Wire, + marker: str, + rail: str, + barrier: str, +) -> tuple[Request, ...]: + accumulator: Final = _SinkRowsAccumulator(sink) + eventually( + accumulator.drain, + lambda rows: bool(_rail_scans(rows, barrier, marker)), + seconds=70, + ) + return _rail_scans(accumulator.rows, rail, marker) + + +def _scope_scan_count(scope: MatrixScope, streamed: bool) -> int: + if scope is None or scope == "both": + return 1 + return int((scope == "streaming") == streamed) + + +def _expected_logging_phases(scope: MatrixScope, streamed: bool) -> tuple[str, ...]: + return ("request", "response") if _scope_scan_count(scope, streamed) else () + + +def _spend_row_for_call(call_id: str, content: bytes) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + "SELECT request_id, litellm_call_id, spend, prompt_tokens, completion_tokens, metadata " + 'FROM "LiteLLM_SpendLogs" WHERE request_id=%s OR litellm_call_id=%s', + (call_id, call_id), + ), + lambda values: len(values) >= 1, + seconds=70, + ) + assert len(rows) == 1, (call_id, rows, content) + assert call_id in (rows[0]["request_id"], rows[0]["litellm_call_id"]), content + return rows[0] + + +def _cell_rows( + rig: MatrixRig, + marker: str, + response: httpx.Response, + mode: Mode, + streamed: bool, + scope: MatrixScope, + call_id: str, + rail: str, +) -> tuple[tuple[Request, ...], tuple[Request, ...]]: + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows, response.text) + _spend_row_for_call(call_id, response.content) + sink_rows: Final = ( + _logging_only_scans( + rig.sink, + marker, + rail, + rig.rails["z_logging_only_barrier"], + ) + if mode == "logging_only" + else _rail_scans(rig.sink.drain(), rail, marker) + ) + return provider_rows, sink_rows + + +def _run_matrix_cell( + rig: MatrixRig, + endpoint: Endpoint, + streamed: bool, + mode: Mode, + scope: MatrixScope, +) -> None: + marker: Final = f"audit-a-{endpoint}-{int(streamed)}-{mode}-{scope or 'unset'}-{uuid.uuid4().hex}" + call_id: Final = f"matrix-a-{uuid.uuid4().hex}" + rail: Final = rig.rails[f"a_{mode}_{scope or 'unset'}"] + path, body = _cell_body(rig, endpoint, marker, streamed, rail) + response: Final = rig.candidate.request( + "POST", + path, + body, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 200, response.text + assert PROVIDER_TEXT in response.text and marker in response.text, response.text + provider_rows, sink_rows = _cell_rows(rig, marker, response, mode, streamed, scope, call_id, rail) + assert len(provider_rows) == 1, (marker, provider_rows, response.text) + if mode == "logging_only": + assert _logging_phases(sink_rows) == _expected_logging_phases(scope, streamed), ( + marker, + sink_rows, + response.text, + ) + else: + expected_scans: Final = _scope_scan_count(scope, streamed) + assert len(sink_rows) == expected_scans, (marker, sink_rows, response.text) + + +@pytest.mark.parametrize("endpoint", Endpoints) +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +@pytest.mark.parametrize("mode", Modes) +@pytest.mark.parametrize("scope", Scopes, ids=("streaming", "non_streaming", "unset")) +def test_a_stream_scope_matrix( + rig: MatrixRig, + endpoint: Endpoint, + streamed: bool, + mode: Mode, + scope: MatrixScope, +) -> None: + _run_matrix_cell(rig, endpoint, streamed, mode, scope) + + +SDK_CASES: Final[tuple[str, ...]] = ( + "openai-chat-sync", + "openai-chat-async", + "openai-responses-sync", + "openai-responses-async", + "anthropic-messages-sync", + "anthropic-messages-async", +) + + +def _openai_sdk_base_url(rig: MatrixRig) -> str: + return f"{str(rig.candidate.client.base_url).rstrip('/')}/v1" + + +def _anthropic_sdk_base_url(rig: MatrixRig) -> str: + return str(rig.candidate.client.base_url).rstrip("/") + + +def _verify_sdk_call(rig: MatrixRig, marker: str, call_id: str, text: str, streamed: bool) -> None: + assert marker in text and PROVIDER_TEXT in text, text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + _spend_row_for_call(call_id, text.encode()) + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails["c2_key"], marker) + assert len(sink_rows) == int(streamed), (marker, streamed, sink_rows) + + +def _run_sync_sdk(rig: MatrixRig, sdk: str, marker: str, call_id: str, streamed: bool) -> str: + if sdk == "openai-chat-sync": + with openai.OpenAI( + api_key=rig.key, + base_url=_openai_sdk_base_url(rig), + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) as client: + response: Final = client.chat.completions.create( + model=rig.models["chat"], + messages=[{"role": "user", "content": marker}], + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + return "".join( + chunk.choices[0].delta.content or "" for chunk in response if chunk.choices[0].delta.content + ) + return cast(str, response.choices[0].message.content) + if sdk == "openai-responses-sync": + with openai.OpenAI( + api_key=rig.key, + base_url=_openai_sdk_base_url(rig), + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) as client: + response = client.responses.create( + model=rig.models["responses"], + input=marker, + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + return "".join(event.delta for event in response if event.type == "response.output_text.delta") + return response.output_text + if sdk == "anthropic-messages-sync": + with anthropic.Anthropic( + api_key=rig.key, + base_url=_anthropic_sdk_base_url(rig), + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) as client: + response = client.messages.create( + model=rig.models["messages"], + max_tokens=32, + messages=[{"role": "user", "content": marker}], + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + return "".join( + event.delta.text + for event in response + if event.type == "content_block_delta" and event.delta.type == "text_delta" + ) + return "".join(block.text for block in response.content if block.type == "text") + raise AssertionError(sdk) + + +async def _run_async_sdk(rig: MatrixRig, sdk: str, marker: str, call_id: str, streamed: bool) -> str: + if sdk == "openai-chat-async": + async with openai.AsyncOpenAI( + api_key=rig.key, + base_url=_openai_sdk_base_url(rig), + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=30), + ) as client: + response = await client.chat.completions.create( + model=rig.models["chat"], + messages=[{"role": "user", "content": marker}], + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + chunks: Final = [ + chunk.choices[0].delta.content or "" async for chunk in response if chunk.choices[0].delta.content + ] + return "".join(chunks) + return cast(str, response.choices[0].message.content) + if sdk == "openai-responses-async": + async with openai.AsyncOpenAI( + api_key=rig.key, + base_url=_openai_sdk_base_url(rig), + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=30), + ) as client: + response = await client.responses.create( + model=rig.models["responses"], + input=marker, + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + events: Final = [event.delta async for event in response if event.type == "response.output_text.delta"] + return "".join(events) + return response.output_text + if sdk == "anthropic-messages-async": + async with anthropic.AsyncAnthropic( + api_key=rig.key, + base_url=_anthropic_sdk_base_url(rig), + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=30), + ) as client: + response = await client.messages.create( + model=rig.models["messages"], + max_tokens=32, + messages=[{"role": "user", "content": marker}], + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + events: Final = [ + event.delta.text + async for event in response + if event.type == "content_block_delta" and event.delta.type == "text_delta" + ] + return "".join(events) + return "".join(block.text for block in response.content if block.type == "text") + raise AssertionError(sdk) + + +@pytest.mark.parametrize("sdk", SDK_CASES) +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +def test_b_streaming_scope_classifies_sdk_streams(rig: MatrixRig, sdk: str, streamed: bool) -> None: + marker: Final = f"audit-b-{sdk}-{int(streamed)}-{uuid.uuid4().hex}" + call_id: Final = f"matrix-b-{uuid.uuid4().hex}" + if sdk.endswith("-async"): + text: Final = run(_run_async_sdk(rig, sdk, marker, call_id, streamed)) + else: + text = _run_sync_sdk(rig, sdk, marker, call_id, streamed) + _verify_sdk_call(rig, marker, call_id, text, streamed) + + +def _yaml_proxy_config( + rig: MatrixRig, + guardrail_name: str, + parameters: Mapping[str, JsonValue], +) -> dict[str, JsonValue]: + config_base: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + return cast( + dict[str, JsonValue], + { + **config_base, + "guardrails": [{"guardrail_name": guardrail_name, "litellm_params": dict(parameters)}], + "model_list": [ + { + "model_name": "scope-yaml-chat", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{rig.provider.url}/v1", + "api_key": "synthetic-provider-key", + }, + } + ], + "environment_variables": { + **config_base.get("environment_variables", {}), + "OPENAI_API_BASE": rig.provider.url, + "OPENAI_API_KEY": "synthetic-provider-key", + }, + }, + ) + + +def _management_params( + name: str, + rig: MatrixRig, + *, + mode: str | list[str] = "pre_call", + scope: JsonValue = "streaming", + default_on: bool = False, +) -> dict[str, JsonValue]: + return { + "guardrail": "generic_guardrail_api", + "mode": mode, + "default_on": default_on, + "api_base": f"{rig.sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + "stream_scope": scope, + } + + +def _raw_chat( + gateway: Gateway, + marker: str, + streamed: bool, + model: str, + *, + rails: Sequence[str] = (), + key: str | None = None, + call_id: str | None = None, +) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": marker}], + "guardrails": list(rails), + **({"stream": True} if streamed else {}), + }, + key=key, + headers={} if call_id is None else {"x-litellm-call-id": call_id}, + ) + + +def _assert_raw_call( + rig: MatrixRig, + marker: str, + response: httpx.Response, + streamed: bool, + expected_scans: int, + rail: str, + call_id: str | None = None, +) -> tuple[Request, ...]: + assert response.status_code == 200, response.text + assert PROVIDER_TEXT in response.text and marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + if call_id is not None: + _spend_row_for_call(call_id, response.content) + sink_rows: Final = _rail_scans(rig.sink.drain(), rail, marker) + assert len(sink_rows) == expected_scans, (marker, streamed, expected_scans, sink_rows) + return sink_rows + + +def test_c1_default_on_rail_respects_stream_scope(rig: MatrixRig, tmp_path: Path) -> None: + name: Final = f"c1-default-{uuid.uuid4().hex}" + parameters: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "stream_scope": "streaming", + "api_base": f"{rig.sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + } + config_path: Final = tmp_path / "default-on-stream-scope.yaml" + config_path.write_text(yaml.safe_dump(_yaml_proxy_config(rig, name, parameters))) + with owned_proxy_process(rig.candidate, tmp_path, {}, config=config_path, workers=1) as owned: + for streamed in (False, True): + marker: Final = f"audit-c1-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat(owned.gateway, marker, streamed, "scope-yaml-chat") + _assert_raw_call(rig, marker, response, streamed, int(streamed), name) + + +def test_c2_key_attached_rail_respects_stream_scope(rig: MatrixRig) -> None: + marker0: Final = f"audit-c2-0-{uuid.uuid4().hex}" + response0: Final = _raw_chat(rig.candidate, marker0, False, rig.models["chat"], key=rig.key) + _assert_raw_call(rig, marker0, response0, False, 0, rig.rails["c2_key"]) + marker1: Final = f"audit-c2-1-{uuid.uuid4().hex}" + response1: Final = _raw_chat(rig.candidate, marker1, True, rig.models["chat"], key=rig.key) + _assert_raw_call(rig, marker1, response1, True, 1, rig.rails["c2_key"]) + + +def test_c3_team_attached_rail_respects_stream_scope(rig: MatrixRig) -> None: + team: Final = rig.scenario.team(guardrails=[rig.rails["c3_team"]]) + key: Final = rig.scenario.key(team_id=team) + marker0: Final = f"audit-c3-0-{uuid.uuid4().hex}" + response0: Final = _raw_chat(rig.candidate, marker0, False, rig.models["chat"], key=key) + _assert_raw_call(rig, marker0, response0, False, 0, rig.rails["c3_team"]) + marker1: Final = f"audit-c3-1-{uuid.uuid4().hex}" + response1: Final = _raw_chat(rig.candidate, marker1, True, rig.models["chat"], key=key) + _assert_raw_call(rig, marker1, response1, True, 1, rig.rails["c3_team"]) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +def test_c4_per_mode_scope_map_selects_each_mode(rig: MatrixRig, streamed: bool) -> None: + marker: Final = f"audit-c4-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat( + rig.candidate, + marker, + streamed, + rig.models["chat"], + rails=(rig.rails["c4_modes"],), + ) + assert response.status_code == 200, response.text + assert PROVIDER_TEXT in response.text and marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails["c4_modes"], marker) + request_bodies: Final = tuple(JSON_OBJECT.validate_json(row.body) for row in sink_rows) + request_texts: Final = tuple(chain.from_iterable(_strings(body.get("texts", [])) for body in request_bodies)) + assert len(sink_rows) == 1, (marker, streamed, sink_rows) + assert any(marker in text for text in request_texts), (marker, request_texts) + assert (any(PROVIDER_TEXT in text for text in request_texts)) == streamed, ( + marker, + streamed, + request_texts, + ) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +def test_c5_mode_omitted_from_scope_map_means_both(rig: MatrixRig, streamed: bool) -> None: + marker: Final = f"audit-c5-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat( + rig.candidate, + marker, + streamed, + rig.models["chat"], + rails=(rig.rails["c5_omitted"],), + ) + sink_rows: Final = _assert_raw_call( + rig, + marker, + response, + streamed, + 1 + int(streamed), + rig.rails["c5_omitted"], + ) + request_bodies: Final = tuple(JSON_OBJECT.validate_json(row.body) for row in sink_rows) + request_texts: Final = tuple(chain.from_iterable(_strings(body.get("texts", [])) for body in request_bodies)) + assert sum(marker in text and PROVIDER_TEXT not in text for text in request_texts) == 1, (marker, request_texts) + assert sum(PROVIDER_TEXT in text for text in request_texts) == int(streamed), (marker, request_texts) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +def test_c6_explicit_both_matches_unset_scope(rig: MatrixRig, streamed: bool) -> None: + marker: Final = f"audit-c6-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat( + rig.candidate, + marker, + streamed, + rig.models["chat"], + rails=(rig.rails["c6_both"], rig.rails["c6_unset"]), + ) + assert response.status_code == 200, response.text + assert PROVIDER_TEXT in response.text and marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + sink_rows: Final = rig.sink.drain() + both_rows: Final = _rail_scans(sink_rows, rig.rails["c6_both"], marker) + unset_rows: Final = _rail_scans(sink_rows, rig.rails["c6_unset"], marker) + assert (len(both_rows), len(unset_rows)) == (1, 1), (marker, streamed, both_rows, unset_rows) + + +@pytest.mark.parametrize( + ("case", "scope"), + ( + ("missing", None), + ("null", None), + ("empty", ""), + ("invalid-scalar", "sometimes"), + ("invalid-map", {"unknown_mode": "streaming"}), + ), + ids=("missing", "null", "empty", "invalid-scalar", "invalid-map"), +) +def test_c7_yaml_unset_and_invalid_scope_values_are_tolerated( + rig: MatrixRig, + tmp_path: Path, + case: str, + scope: JsonValue | None, +) -> None: + name: Final = f"c7-yaml-{case}-{uuid.uuid4().hex}" + parameters: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": f"{rig.sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + **({} if case == "missing" else {"stream_scope": scope}), + } + config_path: Final = tmp_path / f"{case}-stream-scope.yaml" + config_path.write_text(yaml.safe_dump(_yaml_proxy_config(rig, name, parameters))) + with owned_proxy_process(rig.candidate, tmp_path, {}, config=config_path, workers=1) as owned: + for streamed in (False, True): + marker: Final = f"audit-c7-{case}-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat(owned.gateway, marker, streamed, "scope-yaml-chat") + _assert_raw_call(rig, marker, response, streamed, 1, name) + if case == "empty": + assert "Ignoring invalid stored stream_scope value of type str" in owned.log.read_text() + + +def test_c8_identical_requests_have_one_scan_and_spend_each(rig: MatrixRig) -> None: + marker: Final = f"audit-c8-identical-{uuid.uuid4().hex}" + call_ids: Final = tuple(f"matrix-c8-{uuid.uuid4().hex}" for _ in range(3)) + responses: Final = tuple( + _raw_chat( + rig.candidate, + marker, + True, + rig.models["chat"], + rails=(rig.rails["a_post_call_streaming"],), + call_id=call_id, + ) + for call_id in call_ids + ) + assert tuple(response.status_code for response in responses) == (200, 200, 200), tuple( + response.text for response in responses + ) + for call_id, response in zip(call_ids, responses): + _spend_row_for_call(call_id, response.content) + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails["a_post_call_streaming"], marker) + assert len(sink_rows) == 3, (marker, sink_rows) + guardrail_call_ids: Final = tuple( + cast(str, JSON_OBJECT.validate_json(row.body).get("litellm_call_id")) for row in sink_rows + ) + assert set(guardrail_call_ids) == set(call_ids), (call_ids, guardrail_call_ids) + + +def _create_management_rail(rig: MatrixRig, name: str, params: Mapping[str, JsonValue]) -> str: + created: Final = rig.candidate.request( + "POST", + "/guardrails", + {"guardrail": {"guardrail_name": name, "litellm_params": dict(params)}}, + ) + assert created.status_code == 200, created.text + payload: Final = JSON_OBJECT.validate_json(created.content) + identity: Final = payload.get("guardrail_id") + assert isinstance(identity, str), payload + return identity + + +def _management_observation( + rig: MatrixRig, + name: str, + streamed: bool, +) -> tuple[int, int, int]: + marker: Final = f"audit-d-management-{int(streamed)}-{uuid.uuid4().hex}" + call_id: Final = f"matrix-d1-{uuid.uuid4().hex}" + response: Final = _raw_chat( + rig.candidate, + marker, + streamed, + rig.models["chat"], + rails=(name,), + call_id=call_id, + ) + provider_rows: Final = _matching_requests(rig.provider, marker) + sink_rows: Final = _rail_scans(rig.sink.drain(), name, marker) + if response.status_code == 200 and len(provider_rows) == 1: + assert marker in response.text, response.text + _spend_row_for_call(call_id, response.content) + return response.status_code, len(provider_rows), len(sink_rows) + + +def _eventually_management_scope( + rig: MatrixRig, + name: str, + expected: tuple[int, int], +) -> tuple[tuple[int, int, int], tuple[int, int, int]]: + return eventually( + lambda: ( + _management_observation(rig, name, False), + _management_observation(rig, name, True), + ), + lambda values: ( + (values[0][0], values[0][2]) == (200, expected[0]) + and (values[1][0], values[1][2]) == (200, expected[1]) + and values[0][1] == values[1][1] == 1 + ), + seconds=30, + ) + + +def test_d1_management_create_read_and_runtime_scope(rig: MatrixRig) -> None: + name: Final = f"d1-management-{uuid.uuid4().hex}" + params: Final = _management_params(name, rig) + identity: Final = _create_management_rail(rig, name, params) + try: + info: Final = rig.candidate.request("GET", f"/guardrails/{identity}/info") + listing: Final = rig.candidate.request("GET", "/v2/guardrails/list") + info_payload: Final = JSON_OBJECT.validate_json(info.content) + list_payload: Final = JSON_OBJECT.validate_json(listing.content) + list_rows: Final = list_payload.get("guardrails") + listed: Final = ( + tuple(row for row in list_rows if isinstance(row, dict) and row.get("guardrail_id") == identity) + if isinstance(list_rows, list) + else () + ) + info_params: Final = info_payload.get("litellm_params") + list_params: Final = listed[0].get("litellm_params") if listed else None + assert ( + info.status_code == 200 + and listing.status_code == 200 + and isinstance(info_params, dict) + and info_params.get("stream_scope") == "streaming" + and len(listed) == 1 + and isinstance(list_params, dict) + and list_params.get("stream_scope") == "streaming" + ), (info.status_code, info.text, listing.status_code, listing.text) + observed: Final = _eventually_management_scope(rig, name, (0, 1)) + assert observed[0][2] == 0 and observed[1][2] == 1, observed + finally: + deleted: Final = rig.candidate.request("DELETE", f"/guardrails/{identity}") + assert deleted.status_code == 200, deleted.text + + +def test_d2_patch_then_put_updates_runtime_scope(rig: MatrixRig) -> None: + name: Final = f"d2-management-{uuid.uuid4().hex}" + identity: Final = _create_management_rail(rig, name, _management_params(name, rig)) + try: + patched: Final = rig.candidate.request( + "PATCH", + f"/guardrails/{identity}", + {"litellm_params": {"stream_scope": "non_streaming"}}, + ) + assert patched.status_code == 200, patched.text + after_patch: Final = _eventually_management_scope(rig, name, (1, 0)) + assert after_patch[0][2] == 1 and after_patch[1][2] == 0, after_patch + params: Final = _management_params(name, rig) + updated: Final = rig.candidate.request( + "PUT", + f"/guardrails/{identity}", + {"guardrail": {"guardrail_name": name, "litellm_params": params}}, + ) + assert updated.status_code == 200, updated.text + after_put: Final = _eventually_management_scope(rig, name, (0, 1)) + assert after_put[0][2] == 0 and after_put[1][2] == 1, after_put + finally: + deleted: Final = rig.candidate.request("DELETE", f"/guardrails/{identity}") + assert deleted.status_code == 200, deleted.text + + +HOSTILE_SCOPE_VALUES: Final[tuple[JsonValue, ...]] = ( + 1, + [], + "", + "x" * 5000, + {"unknown_mode": "streaming"}, + {"pre_call": 1}, +) + + +@pytest.mark.parametrize("operation", ("POST", "PUT", "PATCH")) +@pytest.mark.parametrize( + "scope", HOSTILE_SCOPE_VALUES, ids=("integer", "list", "empty", "long", "unknown-key", "wrong-value") +) +@pytest.mark.parametrize("repeat", (1, 2), ids=("first", "second")) +def test_d3_management_rejects_invalid_scope_values( + rig: MatrixRig, + operation: str, + scope: JsonValue, + repeat: int, +) -> None: + name: Final = f"d3-management-{operation.lower()}-{repeat}-{uuid.uuid4().hex}" + identity: Final = _create_management_rail(rig, name, _management_params(name, rig)) if operation != "POST" else None + payload: Final = ( + {"litellm_params": {"stream_scope": scope}} + if operation == "PATCH" + else { + "guardrail": { + "guardrail_name": name, + "litellm_params": _management_params(name, rig, scope=scope), + } + } + ) + path: Final = "/guardrails" if operation == "POST" else f"/guardrails/{identity}" + try: + response: Final = rig.candidate.request(operation, path, payload) + if operation == "POST" and response.status_code == 200: + created_payload: Final = JSON_OBJECT.validate_json(response.content) + created_identity: Final = created_payload.get("guardrail_id") + if isinstance(created_identity, str): + rig.candidate.request("DELETE", f"/guardrails/{created_identity}") + assert response.status_code == 422 and "stream_scope" in response.text, ( + operation, + scope, + repeat, + response.status_code, + response.text, + ) + if operation == "POST": + stored: Final = read_rows( + 'SELECT guardrail_id FROM "LiteLLM_GuardrailsTable" WHERE guardrail_name=%s', + (name,), + ) + assert stored == [], (operation, scope, stored) + else: + existing: Final = rig.candidate.request("GET", f"/guardrails/{identity}/info") + existing_payload: Final = JSON_OBJECT.validate_json(existing.content) + existing_params: Final = existing_payload.get("litellm_params") + assert ( + existing.status_code == 200 + and isinstance(existing_params, dict) + and existing_params.get("stream_scope") == "streaming" + ), (operation, scope, existing.status_code, existing.text) + finally: + if isinstance(identity, str): + deleted: Final = rig.candidate.request("DELETE", f"/guardrails/{identity}") + assert deleted.status_code == 200, deleted.text + + +def test_d3_management_normalizes_uppercase_scope(rig: MatrixRig) -> None: + name: Final = f"d3-uppercase-{uuid.uuid4().hex}" + identity: Final = _create_management_rail(rig, name, _management_params(name, rig, scope="STREAMING")) + try: + info: Final = rig.candidate.request("GET", f"/guardrails/{identity}/info") + payload: Final = JSON_OBJECT.validate_json(info.content) + parameters: Final = payload.get("litellm_params") + assert ( + info.status_code == 200 and isinstance(parameters, dict) and parameters.get("stream_scope") == "streaming" + ), (info.status_code, info.text) + finally: + deleted: Final = rig.candidate.request("DELETE", f"/guardrails/{identity}") + assert deleted.status_code == 200, deleted.text + + +def test_d4_unauthenticated_management_create_is_rejected(rig: MatrixRig) -> None: + name: Final = f"d4-unauthenticated-{uuid.uuid4().hex}" + response: Final = rig.candidate.client.request( + "POST", + "/guardrails", + json={"guardrail": {"guardrail_name": name, "litellm_params": _management_params(name, rig)}}, + ) + assert response.status_code == 401, response.text + + +def _scoped_rows( + rows: Sequence[Request], + marker: str, + streaming_name: str, + non_streaming_name: str, +) -> tuple[tuple[Request, ...], tuple[Request, ...]]: + return ( + _rail_scans(rows, streaming_name, marker), + _rail_scans(rows, non_streaming_name, marker), + ) + + +HOSTILE_CLASSIFICATION_VALUES: Final[tuple[JsonValue, ...]] = ( + True, + 1, + [], + "", + "litellm-server-streaming", + "x" * 5000, +) + + +@pytest.mark.parametrize( + "value", + HOSTILE_CLASSIFICATION_VALUES, + ids=("true", "integer", "list", "empty", "marker", "long"), +) +def test_ea_is_streaming_request_body_does_not_change_chat_classification( + rig: MatrixRig, + value: JsonValue, +) -> None: + marker: Final = f"audit-ea-{uuid.uuid4().hex}" + response: Final = rig.candidate.request( + "POST", + "/v1/chat/completions", + { + "model": rig.models["chat"], + "messages": [{"role": "user", "content": marker}], + "guardrails": [rig.rails["e_stream"], rig.rails["e_non_stream"]], + "is_streaming_request": value, + }, + ) + assert response.status_code == 200, response.text + assert PROVIDER_TEXT in response.text and marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + assert provider_body.get("is_streaming_request") == value, provider_rows[0].body.decode() + sink_rows: Final = rig.sink.drain() + streaming_rows, non_streaming_rows = _scoped_rows( + sink_rows, + marker, + rig.rails["e_stream"], + rig.rails["e_non_stream"], + ) + assert (len(streaming_rows), len(non_streaming_rows)) == (0, 1), (marker, streaming_rows, non_streaming_rows) + + +@pytest.mark.parametrize("endpoint", ("chat", "passthrough"), ids=("chat", "configured-pass-through")) +@pytest.mark.parametrize("value", (True, "litellm-server-streaming"), ids=("true", "marker")) +def test_eb_namespaced_caller_body_field_cannot_flip_classification( + rig: MatrixRig, + endpoint: str, + value: JsonValue, +) -> None: + marker: Final = f"audit-eb-{endpoint}-{uuid.uuid4().hex}" + body: Final = ( + { + "model": rig.models["chat"], + "messages": [{"role": "user", "content": marker}], + "guardrails": [rig.rails["e_stream"], rig.rails["e_non_stream"]], + "litellm_server_streaming_classification": value, + } + if endpoint == "chat" + else { + "marker": marker, + "litellm_server_streaming_classification": value, + } + ) + response: Final = rig.candidate.request( + "POST", + "/v1/chat/completions" if endpoint == "chat" else "/pt", + body, + ) + assert response.status_code == 200 and marker in response.text, response.text + actual_streamed: Final = response.headers.get("content-type", "").lower().startswith("text/event-stream") + assert not actual_streamed, (endpoint, response.headers, response.text) + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + if endpoint == "passthrough" and value == "litellm-server-streaming": + assert "litellm_server_streaming_classification" not in provider_body, provider_rows[0].body.decode() + else: + assert provider_body.get("litellm_server_streaming_classification") == value, provider_rows[0].body.decode() + sink_rows: Final = rig.sink.drain() + streaming_rows, non_streaming_rows = _scoped_rows( + sink_rows, + marker, + rig.rails["e_stream"], + rig.rails["e_non_stream"], + ) + assert (len(streaming_rows), len(non_streaming_rows)) == (0, 1), (marker, streaming_rows, non_streaming_rows) + + +HOSTILE_STREAM_VALUES: Final[tuple[JsonValue, ...]] = ("true", 1, [], "") + + +@pytest.mark.parametrize("value", HOSTILE_STREAM_VALUES, ids=("string", "integer", "list", "empty")) +def test_ec_hostile_stream_values_follow_observed_response_shape(rig: MatrixRig, value: JsonValue) -> None: + marker: Final = f"audit-ec-{uuid.uuid4().hex}" + response: Final = rig.candidate.request( + "POST", + "/v1/chat/completions", + { + "model": rig.models["chat"], + "messages": [{"role": "user", "content": marker}], + "guardrails": [rig.rails["e_stream"], rig.rails["e_non_stream"]], + "stream": value, + }, + ) + assert response.status_code == 200, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + observed_stream: Final = response.headers.get("content-type", "").startswith("text/event-stream") + assert marker in response.text, response.text + sink_rows: Final = rig.sink.drain() + streaming_rows, non_streaming_rows = _scoped_rows( + sink_rows, + marker, + rig.rails["e_stream"], + rig.rails["e_non_stream"], + ) + assert (len(streaming_rows), len(non_streaming_rows)) == ( + int(observed_stream), + int(not observed_stream), + ), (marker, value, response.headers, streaming_rows, non_streaming_rows) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("stream-absent", "stream-true")) +def test_ed_configured_passthrough_forwards_caller_flag_and_uses_route_body_stream( + rig: MatrixRig, + streamed: bool, +) -> None: + marker: Final = f"audit-ed-{int(streamed)}-{uuid.uuid4().hex}" + body: Final = { + "marker": marker, + "is_streaming_request": True, + **({"stream": True} if streamed else {}), + } + response: Final = rig.candidate.request("POST", "/pt", body) + assert response.status_code == 200 and marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + assert provider_body == body, (body, provider_body) + + +@pytest.mark.parametrize( + ("route", "body", "streamed"), + ( + ( + "/anthropic/v1/messages", + { + "model": "synthetic-anthropic-model", + "max_tokens": 32, + "messages": [{"role": "user", "content": "MARKER"}], + }, + False, + ), + ( + "/anthropic/v1/messages", + { + "model": "synthetic-anthropic-model", + "max_tokens": 32, + "messages": [{"role": "user", "content": "MARKER"}], + "stream": True, + }, + True, + ), + ( + "/gemini/v1beta/models/audit-model:generateContent", + {"contents": [{"parts": [{"text": "MARKER"}]}]}, + False, + ), + ( + "/gemini/v1beta/models/audit-model:streamGenerateContent?alt=sse", + {"contents": [{"parts": [{"text": "MARKER"}]}]}, + True, + ), + ), + ids=("anthropic-S0", "anthropic-S1", "gemini-generate", "gemini-stream"), +) +def test_eh_provider_passthrough_routes_classify_effective_streaming( + rig: MatrixRig, + route: str, + body: dict[str, JsonValue], + streamed: bool, +) -> None: + marker: Final = f"audit-eh-{uuid.uuid4().hex}" + request_body: Final = JSON_OBJECT.validate_python(json.loads(json.dumps(body).replace("MARKER", marker))) + team: Final = ( + rig.scenario.team( + metadata={ + "allowed_passthrough_routes": [ + "/gemini/v1beta/models/audit-model:generateContent", + "/gemini/v1beta/models/audit-model:streamGenerateContent", + ] + } + ) + if route.startswith("/gemini/") + else None + ) + key: Final = ( + rig.scenario.key(team_id=team, guardrails=[rig.rails["e_stream"], rig.rails["e_non_stream"]]) + if team is not None + else rig.scenario.key(guardrails=[rig.rails["e_stream"], rig.rails["e_non_stream"]]) + ) + headers: Final = {"x-goog-api-key": key} if route.startswith("/gemini/") else None + response: Final = rig.candidate.request( + "POST", + route, + request_body, + headers=headers, + ) + assert response.status_code == 200 and marker in response.text, response.text + actual_streamed: Final = response.headers.get("content-type", "").lower().startswith("text/event-stream") + assert actual_streamed is streamed, (route, streamed, response.headers, response.text) + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + sink_rows: Final = rig.sink.drain() + streaming_rows, non_streaming_rows = _scoped_rows( + sink_rows, + marker, + rig.rails["e_stream"], + rig.rails["e_non_stream"], + ) + assert (len(streaming_rows), len(non_streaming_rows)) == ((1, 0) if streamed else (0, 1)), ( + marker, + route, + streaming_rows, + non_streaming_rows, + ) + + +def test_ei_unauthenticated_scoped_request_has_no_upstream_or_guardrail_call(rig: MatrixRig) -> None: + marker: Final = f"audit-ei-{uuid.uuid4().hex}" + response: Final = rig.candidate.client.request( + "POST", + "/v1/chat/completions", + json={ + "model": rig.models["chat"], + "messages": [{"role": "user", "content": marker}], + "stream": True, + "guardrails": [rig.rails["e_stream"]], + }, + ) + assert response.status_code == 401, response.text + assert _matching_requests(rig.provider, marker) == () + assert _rail_scans(rig.sink.drain(), rig.rails["e_stream"], marker) == () + + +@pytest.mark.parametrize( + ("rail_key", "blocked"), + (("f1_stream", False), ("f1_non_stream", True), ("f1_unset", True)), + ids=("streaming", "non_streaming", "unset"), +) +def test_f1_mcp_content_filter_respects_stream_scope( + rig: MatrixRig, + rail_key: Literal["f1_stream", "f1_non_stream", "f1_unset"], + blocked: bool, +) -> None: + with scripted_peer(echo_tool("echo")) as peer: + alias: Final = f"streamscope{uuid.uuid4().hex}" + identity: Final = register_mcp(rig.scenario, peer, alias) + key: Final = rig.scenario.key( + object_permission={"mcp_servers": [identity]}, + guardrails=[rig.rails[rail_key]], + ) + marker: Final = f"audit-f1-{uuid.uuid4().hex}" + caller: Final = McpCaller(rig.candidate, key, "rest", alias) + outcome: Final = caller.call( + f"{alias}-echo", + {"text": f"{F1_BLOCKED_WORD} {marker}"}, + identity, + ) + peer_requests: Final = peer.drain() + marker_requests: Final = tuple(request for request in peer_requests if marker in json.dumps(request)) + if blocked: + assert outcome.error is not None, outcome.raw + assert "block" in outcome.raw.casefold() or F1_BLOCKED_WORD in outcome.raw.casefold(), outcome.raw + assert marker_requests == (), (marker, marker_requests) + return + assert outcome.ok and marker in (outcome.text or ""), outcome.raw + assert len(marker_requests) == 1, (marker, marker_requests) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +def test_f2_mismatched_policy_step_skips_matching_sibling_still_enforces( + rig: MatrixRig, + streamed: bool, +) -> None: + marker: Final = f"audit-f2-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat( + rig.candidate, + marker, + streamed, + "matrix-pipeline-model", + ) + provider_rows: Final = _matching_requests(rig.provider, marker) + sink_rows: Final = rig.sink.drain() + streaming_rows, non_streaming_rows = _scoped_rows( + sink_rows, + marker, + rig.rails["f2_stream"], + rig.rails["f2_non_stream"], + ) + assert len(provider_rows) == 1, (marker, provider_rows) + if streamed: + assert response.status_code == 200 and marker in response.text, response.text + assert (len(streaming_rows), len(non_streaming_rows)) == (1, 0), (marker, streaming_rows, non_streaming_rows) + return + assert response.status_code == 400 and "synthetic policy block" in response.text, response.text + assert (len(streaming_rows), len(non_streaming_rows)) == (0, 1), (marker, streaming_rows, non_streaming_rows) + + +def test_f3_realtime_transcription_uses_streaming_scope(rig: MatrixRig) -> None: + streaming_scenario_id: Final = f"matrix-f3-stream-{uuid.uuid4().hex}" + streaming_upstream: Final = register_scenario( + streaming_scenario_id, + _realtime_transcription_response(F3_STREAM_BLOCKED_WORD), + ) + rig.scenario.cleanups.callback(delete_scenario, streaming_upstream) + streaming_model: Final = rig.scenario.model( + model="openai/gpt-realtime-2", + api_key=streaming_scenario_id, + api_base=rig.candidate.upstream_url, + ) + non_streaming_scenario_id: Final = f"matrix-f3-non-stream-{uuid.uuid4().hex}" + non_streaming_upstream: Final = register_scenario( + non_streaming_scenario_id, + _realtime_transcription_response(F3_NON_STREAM_BLOCKED_WORD), + ) + rig.scenario.cleanups.callback(delete_scenario, non_streaming_upstream) + non_streaming_model: Final = rig.scenario.model( + model="openai/gpt-realtime-2", + api_key=non_streaming_scenario_id, + api_base=rig.candidate.upstream_url, + ) + key: Final = rig.scenario.key() + proxy_url: Final = str(rig.candidate.client.base_url).rstrip("/") + + streaming_events: Final = run(_realtime_transcription_events(proxy_url, key, streaming_model)) + streaming_transcriptions: Final = tuple( + event + for event in streaming_events + if event.get("type") == "conversation.item.input_audio_transcription.completed" + ) + streaming_errors: Final = tuple(event for event in streaming_events if event.get("type") == "error") + assert len(streaming_transcriptions) == 1, streaming_events + assert streaming_transcriptions[0].get("transcript") == F3_STREAM_BLOCKED_WORD, streaming_events + assert len(streaming_errors) == 1, streaming_events + streaming_error: Final = streaming_errors[0].get("error") + assert isinstance(streaming_error, dict) and streaming_error.get("type") == "guardrail_violation", streaming_events + + non_streaming_events: Final = run(_realtime_transcription_events(proxy_url, key, non_streaming_model)) + non_streaming_transcriptions: Final = tuple( + event + for event in non_streaming_events + if event.get("type") == "conversation.item.input_audio_transcription.completed" + ) + non_streaming_errors: Final = tuple(event for event in non_streaming_events if event.get("type") == "error") + completions: Final = tuple(event for event in non_streaming_events if event.get("type") == "response.done") + assert len(non_streaming_transcriptions) == 1, non_streaming_events + assert non_streaming_transcriptions[0].get("transcript") == F3_NON_STREAM_BLOCKED_WORD, non_streaming_events + assert non_streaming_errors == (), non_streaming_events + assert len(completions) == 1, non_streaming_events + + +def test_g1_sink_failure_fails_closed_only_when_rail_is_in_scope(rig: MatrixRig) -> None: + marker_s0: Final = f"audit-g1-S0-{uuid.uuid4().hex}" + response_s0: Final = _raw_chat( + rig.candidate, + marker_s0, + False, + rig.models["chat"], + rails=(rig.rails["g1_failure"],), + ) + assert response_s0.status_code == 200 and marker_s0 in response_s0.text, response_s0.text + provider_rows_s0: Final = _matching_requests(rig.provider, marker_s0) + sink_rows_s0: Final = _rail_scans(rig.sink.drain(), rig.rails["g1_failure"], marker_s0) + assert len(provider_rows_s0) == 1 and sink_rows_s0 == (), (marker_s0, provider_rows_s0, sink_rows_s0) + marker_s1: Final = f"audit-g1-S1-{uuid.uuid4().hex}" + response_s1: Final = _raw_chat( + rig.candidate, + marker_s1, + True, + rig.models["chat"], + rails=(rig.rails["g1_failure"],), + ) + assert response_s1.status_code == 500 and "Generic Guardrail API failed" in response_s1.text, response_s1.text + assert _matching_requests(rig.provider, marker_s1) == () + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails["g1_failure"], marker_s1) + assert len(sink_rows) == 1, (marker_s1, sink_rows) + + +def test_g2_blocked_verdict_blocks_only_when_rail_is_in_scope(rig: MatrixRig) -> None: + marker_s0: Final = f"audit-g2-S0-{uuid.uuid4().hex}" + response_s0: Final = _raw_chat( + rig.candidate, + marker_s0, + False, + rig.models["chat"], + rails=(rig.rails["g2_block"],), + ) + assert response_s0.status_code == 200 and marker_s0 in response_s0.text, response_s0.text + assert len(_matching_requests(rig.provider, marker_s0)) == 1 + assert _rail_scans(rig.sink.drain(), rig.rails["g2_block"], marker_s0) == () + marker_s1: Final = f"audit-g2-S1-{uuid.uuid4().hex}" + response_s1: Final = _raw_chat( + rig.candidate, + marker_s1, + True, + rig.models["chat"], + rails=(rig.rails["g2_block"],), + ) + assert response_s1.status_code == 400 and "synthetic policy block" in response_s1.text, response_s1.text + assert _matching_requests(rig.provider, marker_s1) == () + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails["g2_block"], marker_s1) + assert len(sink_rows) == 1, (marker_s1, sink_rows) + + +def test_g3_provider_errors_reach_caller_and_proxy_remains_usable(rig: MatrixRig) -> None: + marker: Final = f"audit-g3-401-{uuid.uuid4().hex}" + unauthorized: Final = _raw_chat( + rig.candidate, + marker, + False, + rig.models["chat"], + rails=(rig.rails["g3_provider"],), + ) + assert unauthorized.status_code == 401 and "synthetic provider unauthorized" in unauthorized.text, unauthorized.text + assert len(_matching_requests(rig.provider, marker)) == 1 + assert len(_rail_scans(rig.sink.drain(), rig.rails["g3_provider"], marker)) == 1 + unknown_marker: Final = f"audit-g3-unknown-{uuid.uuid4().hex}" + unknown: Final = _raw_chat( + rig.candidate, + unknown_marker, + False, + "audit-unknown-model", + rails=(rig.rails["g3_provider"],), + ) + assert unknown.status_code in (400, 404) and "model" in unknown.text.lower(), unknown.text + assert _matching_requests(rig.provider, unknown_marker) == () + healthy_marker: Final = f"audit-g3-healthy-{uuid.uuid4().hex}" + healthy: Final = _raw_chat(rig.candidate, healthy_marker, False, rig.models["chat"], key=rig.key) + assert healthy.status_code == 200 and healthy_marker in healthy.text, healthy.text + assert len(_matching_requests(rig.provider, healthy_marker)) == 1 + + +VERTEX_LIVE_HOST: Final = "us-central1-aiplatform.googleapis.com" +VERTEX_LIVE_PATH: Final = "/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" +VERTEX_LIVE_AUTHORITY: Final = f"{VERTEX_LIVE_HOST}:443" + + +@dataclass(frozen=True, slots=True) +class _VertexLivePeer: + port: int + paths: SimpleQueue[str] + authorizations: SimpleQueue[str] + frames: SimpleQueue[dict[str, JsonValue]] + + +@dataclass(frozen=True, slots=True) +class _VertexConnectTunnel: + url: str + authorities: SimpleQueue[str] + + +async def _answer_vertex_live_peer( + connection: ServerConnection, + paths: SimpleQueue[str], + authorizations: SimpleQueue[str], + frames: SimpleQueue[dict[str, JsonValue]], +) -> None: + request: Final = connection.request + assert request is not None + paths.put(request.path) + authorizations.put(request.headers.get("Authorization", "")) + async for raw_frame in connection: + frame: Final = JSON_OBJECT.validate_python(json.loads(raw_frame)) + frames.put(frame) + if frames.qsize() == 1: + await connection.send(json.dumps({"setupComplete": {}})) + + +async def _serve_vertex_live_peer( + tls: ssl.SSLContext, + paths: SimpleQueue[str], + authorizations: SimpleQueue[str], + frames: SimpleQueue[dict[str, JsonValue]], + ports: SimpleQueue[int], + stop: asyncio.Event, +) -> None: + async with serve( + lambda connection: _answer_vertex_live_peer(connection, paths, authorizations, frames), + "127.0.0.1", + 0, + ssl=tls, + ) as server: + ports.put(next(iter(server.sockets)).getsockname()[1]) + await stop.wait() + + +@contextmanager +def _vertex_live_peer(cert: tuple[Path, Path]) -> Iterator[_VertexLivePeer]: + loop: Final = asyncio.new_event_loop() + stop: Final = asyncio.Event() + paths: Final = SimpleQueue[str]() + authorizations: Final = SimpleQueue[str]() + frames: Final = SimpleQueue[dict[str, JsonValue]]() + ports: Final = SimpleQueue[int]() + thread: Final = threading.Thread( + target=loop.run_until_complete, + args=(_serve_vertex_live_peer(server_context(*cert), paths, authorizations, frames, ports, stop),), + daemon=True, + ) + thread.start() + try: + yield _VertexLivePeer(ports.get(timeout=10), paths, authorizations, frames) + finally: + loop.call_soon_threadsafe(stop.set) + thread.join(timeout=10) + loop.close() + + +def _pipe_vertex_socket(source: socket.socket, sink: socket.socket) -> None: + with contextlib.suppress(OSError): + for chunk in iter(lambda: source.recv(65536), b""): + sink.sendall(chunk) + with contextlib.suppress(OSError): + sink.shutdown(socket.SHUT_WR) + + +@contextmanager +def _vertex_connect_tunnel(peer: _VertexLivePeer) -> Generator[_VertexConnectTunnel, None, None]: + authorities: Final = SimpleQueue[str]() + + class ConnectHandler(socketserver.StreamRequestHandler): + rbufsize = 0 + request: socket.socket + + def handle(self) -> None: + request_line: Final = self.rfile.readline().decode().split() + authority: Final = request_line[1] if len(request_line) > 1 else "" + while self.rfile.readline() not in (b"\r\n", b""): + pass + authorities.put(authority) + if authority != VERTEX_LIVE_AUTHORITY: + self.wfile.write(b"HTTP/1.1 403 Forbidden\r\ncontent-length: 0\r\n\r\n") + return + self.wfile.write(b"HTTP/1.1 200 Connection established\r\n\r\n") + self.request.settimeout(10) + with socket.create_connection(("127.0.0.1", peer.port), timeout=10) as upstream: + outbound: Final = threading.Thread(target=_pipe_vertex_socket, args=(self.request, upstream)) + outbound.start() + _pipe_vertex_socket(upstream, self.request) + outbound.join(timeout=12) + + with socketserver.ThreadingTCPServer(("127.0.0.1", 0), ConnectHandler) as server: + thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05}) + thread.start() + try: + yield _VertexConnectTunnel(f"http://127.0.0.1:{server.server_address[1]}", authorities) + finally: + server.shutdown() + thread.join(timeout=6) + + +async def _exchange_vertex_live_frame( + url: str, key: str, vertex_project: str, vertex_location: str, frame: str +) -> str | bytes: + websocket_url: Final = ( + f"{url.replace('http://', 'ws://', 1).rstrip('/')}/vertex_ai/live" + f"?vertex_project={vertex_project}&vertex_location={vertex_location}" + ) + async with websockets.connect( + websocket_url, + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as websocket: + await websocket.send(frame) + reply: Final = await wait_for(websocket.recv(), timeout=10) + await websocket.close() + return reply + + +def test_ej_websocket_passthrough_runs_only_streaming_scoped_rails(rig: MatrixRig, tmp_path: Path) -> None: + vertex_project: Final = f"matrix-ej-{uuid.uuid4().hex}" + vertex_location: Final = "us-central1" + streaming_name: Final = f"ej-streaming-{uuid.uuid4().hex}" + non_streaming_name: Final = f"ej-non-streaming-{uuid.uuid4().hex}" + + def _token_response(_request: Request) -> Reply: + return Reply( + body=_json( + { + "access_token": "synthetic-vertex-access-token", + "expires_in": 3600, + "token_type": "Bearer", + } + ) + ) + + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + private_key_pem: Final = private_key.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ).decode("utf-8") + credentials_path: Final = tmp_path / "vertex-service-account.json" + + with ( + wire_server(_token_response) as token_double, + wire_server(lambda _request: Reply(body=_json({"action": "NONE"}))) as sink, + ): + credentials_path.write_text( + json.dumps( + { + "type": "service_account", + "project_id": vertex_project, + "private_key_id": uuid.uuid4().hex, + "private_key": private_key_pem, + "client_email": f"integration-test@{vertex_project}.iam.gserviceaccount.com", + "client_id": "123456789012345678901", + "token_uri": f"{token_double.url}/token", + "auth_uri": "https://accounts.google.com/o/oauth2/auth", + "auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs", + "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/integration-test", + } + ) + ) + rails: Final = ( + _rail(streaming_name, sink, mode="pre_call", scope="streaming", default_on=True), + _rail(non_streaming_name, sink, mode="pre_call", scope="non_streaming", default_on=True), + ) + config_data: Final = cast(object, yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + config_base: Final = JSON_OBJECT.validate_python(config_data) + config: Final = cast(dict[str, JsonValue], {**config_base, "guardrails": list(rails)}) + config_path: Final = tmp_path / "ej-vertex-live-stream-scope.yaml" + config_path.write_text(yaml.safe_dump(config)) + certificate: Final = write_self_signed_cert(tmp_path, (VERTEX_LIVE_HOST,)) + with _vertex_live_peer(certificate) as peer, _vertex_connect_tunnel(peer) as tunnel: + with owned_proxy_process( + rig.candidate, + tmp_path, + { + "DEFAULT_VERTEXAI_PROJECT": vertex_project, + "DEFAULT_VERTEXAI_LOCATION": vertex_location, + "DEFAULT_GOOGLE_APPLICATION_CREDENTIALS": str(credentials_path), + "HTTPS_PROXY": tunnel.url, + "https_proxy": tunnel.url, + "NO_PROXY": "127.0.0.1,localhost", + "no_proxy": "127.0.0.1,localhost", + "SSL_CERT_FILE": str(certificate[0]), + }, + config=config_path, + workers=1, + ) as owned: + with owned.gateway.scenario() as scenario: + key: Final = scenario.key() + marker: Final = f"vertex-live-{uuid.uuid4().hex}" + frame: Final = json.dumps( + { + "clientContent": { + "turns": [{"role": "user", "parts": [{"text": marker}]}], + "turnComplete": True, + } + } + ) + client_reply_raw: Final = run( + _exchange_vertex_live_frame( + str(owned.gateway.client.base_url), + key, + vertex_project, + vertex_location, + frame, + ) + ) + proxy_log: Final = owned.log + + tunnel_authorities: Final = tuple( + tunnel.authorities.get_nowait() for _ in range(tunnel.authorities.qsize()) + ) + peer_path: Final = peer.paths.get_nowait() + peer_authorization: Final = peer.authorizations.get_nowait() + peer_frames: Final = tuple(peer.frames.get_nowait() for _ in range(peer.frames.qsize())) + client_reply: Final = JSON_OBJECT.validate_python(json.loads(client_reply_raw)) + expected_frame: Final = JSON_OBJECT.validate_python(json.loads(frame)) + token_requests: Final = token_double.drain() + sink_rows: Final = sink.drain() + streaming_rows: Final = tuple(row for row in sink_rows if row.target.startswith(f"/{streaming_name}/")) + non_streaming_rows: Final = tuple( + row for row in sink_rows if row.target.startswith(f"/{non_streaming_name}/") + ) + token_request_count: Final = len(token_requests) + streaming_count: Final = len(streaming_rows) + non_streaming_count: Final = len(non_streaming_rows) + assert tunnel_authorities == (VERTEX_LIVE_AUTHORITY,), ( + tunnel_authorities, + proxy_log, + ) + assert peer_path == VERTEX_LIVE_PATH, ( + peer_path, + proxy_log, + ) + assert peer_authorization == "Bearer synthetic-vertex-access-token", (peer_authorization, proxy_log) + assert peer_frames == (expected_frame,), (marker, peer_frames, proxy_log) + assert client_reply == {"setupComplete": {}}, (client_reply, proxy_log) + assert token_request_count == 1, (token_request_count, proxy_log) + assert streaming_count == 1, (streaming_count, non_streaming_count, proxy_log) + assert non_streaming_count == 0, (streaming_count, non_streaming_count, proxy_log) diff --git a/tests/integration/observability/test_prometheus_request_metrics.py b/tests/integration/observability/test_prometheus_request_metrics.py new file mode 100644 index 00000000000..042e4350640 --- /dev/null +++ b/tests/integration/observability/test_prometheus_request_metrics.py @@ -0,0 +1,274 @@ +from __future__ import annotations + +import json +import uuid +from hashlib import sha256 +from collections.abc import Iterator, Sequence +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import pytest +import yaml +from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS +from litellm.types.integrations.prometheus import LATENCY_BUCKETS + +from tests.integration._support.client import Gateway, eventually, gateway_from_environment, object_value +from tests.integration._support.process import owned_proxy +from tests.integration._support.prometheus_series import Sample, label_values, scrape +from tests.integration._support.wire import Reply, Request, Wire, wire_server + +_GOOD: Final = "prometheus-good-endpoint" +_LIMITED: Final = "prometheus-rate-limited-endpoint" +_FAILING: Final = "prometheus-failing-endpoint" +_LATENCY: Final = "prometheus-latency-endpoint" +_END_USER: Final = f"prometheus-end-user-{uuid.uuid4().hex}" + + +@dataclass(frozen=True, slots=True) +class _Rig: + gateway: Gateway + upstream: Wire + + +def _respond(request: Request) -> Reply: + if json.loads(request.body)["model"] == "429": + return Reply( + status=429, + body=json.dumps({"error": {"message": "rate limited", "type": "rate_limit_error", "code": "429"}}).encode(), + ) + return Reply( + body=json.dumps( + { + "id": f"chatcmpl-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "metered"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4}, + } + ).encode() + ) + + +def _config(directory: Path, upstream: Wire) -> Path: + configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + params: Final = {"api_key": "sk-fixture", "api_base": f"{upstream.url}/v1"} + configuration["model_list"] = [ + *configuration["model_list"], + *( + { + "model_name": name, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + **params, + }, + } + for name in (_GOOD, _LATENCY) + ), + {"model_name": _LIMITED, "litellm_params": {"model": "openai/429", **params}}, + {"model_name": _FAILING, "litellm_params": {"model": "openai/429", **params}}, + ] + configuration["litellm_settings"]["callbacks"] = ["prometheus"] + configuration["litellm_settings"]["disable_end_user_cost_tracking_prometheus_only"] = True + configuration.setdefault("router_settings", {})["num_retries"] = 0 + path: Final = directory / "prometheus-request-metrics.yaml" + path.write_text(yaml.safe_dump(configuration)) + return path + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]: + directory: Final = tmp_path_factory.mktemp("prometheus-request-metrics") + with ( + wire_server(_respond) as upstream, + gateway_from_environment() as shared, + owned_proxy(shared, directory, {}, config=_config(directory, upstream)) as owned, + ): + yield _Rig(owned, upstream) + + +def _ask(rig: _Rig, model: str, key: str | None = None, **extra: list[str] | str) -> httpx.Response: + return rig.gateway.request( + "POST", + "/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"metrics {uuid.uuid4().hex}"}], **extra}, + key=key, + ) + + +def _series(samples: Sequence[Sample], name: str, **labels: str) -> tuple[Sample, ...]: + return tuple( + sample + for sample in samples + if sample.name == name and all(sample.labels.get(label) == value for label, value in labels.items()) + ) + + +def _until(rig: _Rig, name: str, **labels: str) -> tuple[Sample, ...]: + return eventually(lambda: _series(scrape(rig.gateway), name, **labels), bool, seconds=30) + + +def test_a_rate_limited_call_counts_as_a_failed_and_a_429_total_request(rig: _Rig) -> None: + response: Final = _ask(rig, _FAILING) + assert response.status_code == 429, response.text + assert len(rig.upstream.drain()) == 1 + failed: Final = _until( + rig, + "litellm_proxy_failed_requests_metric_total", + api_key_alias="None", + exception_class="Openai.RateLimitError", + exception_status="429", + hashed_api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, + requested_model=_FAILING, + route="/chat/completions", + ) + assert [sample.value for sample in failed] == [1.0] + totals: Final = _until( + rig, + "litellm_proxy_total_requests_metric_total", + hashed_api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, + requested_model=_FAILING, + status_code="429", + ) + assert [sample.value for sample in totals] == [1.0] + + +def test_a_good_call_exports_latency_histograms_on_the_shared_buckets_without_the_end_user(rig: _Rig) -> None: + response: Final = _ask(rig, _LATENCY, user=_END_USER, tags=["teamB"]) + assert response.status_code == 200, response.text + assert len(rig.upstream.drain()) == 1 + master: Final = { + "api_key_alias": "None", + "hashed_api_key": LITELLM_PROXY_MASTER_KEY_ALIAS, + "requested_model": _LATENCY, + } + _until(rig, "litellm_request_total_latency_metric_bucket", le="0.005", **master) + _until(rig, "litellm_llm_api_latency_metric_bucket", le="0.005", **master) + for name in ("litellm_request_total_latency_metric_count", "litellm_llm_api_latency_metric_count"): + assert [sample.value for sample in _until(rig, name, **master)] == [1.0], name + samples: Final = scrape(rig.gateway) + assert _END_USER not in label_values(samples) + expected: Final = {str(bucket).replace("inf", "+Inf") for bucket in LATENCY_BUCKETS} + for name in ( + "litellm_request_total_latency_metric_bucket", + "litellm_llm_api_latency_metric_bucket", + "litellm_overhead_latency_metric_bucket", + ): + assert {sample.labels["le"] for sample in _series(samples, name)} == expected, name + + +def test_client_side_fallbacks_count_one_success_and_one_failure(rig: _Rig) -> None: + recovered: Final = _ask(rig, _LIMITED, fallbacks=[_GOOD]) + assert recovered.status_code == 200, recovered.text + missing: Final = f"unknown-model-{uuid.uuid4().hex[:8]}" + failed: Final = _ask(rig, _LIMITED, fallbacks=[missing]) + assert failed.status_code >= 400, failed.text + rig.upstream.drain() + shared: Final = { + "api_key_alias": "None", + "exception_class": "Openai.RateLimitError", + "exception_status": "429", + "hashed_api_key": LITELLM_PROXY_MASTER_KEY_ALIAS, + "requested_model": _LIMITED, + } + succeeded: Final = _until(rig, "litellm_deployment_successful_fallbacks_total", fallback_model=_GOOD, **shared) + assert [sample.value for sample in succeeded] == [1.0] + lost: Final = _until(rig, "litellm_deployment_failed_fallbacks_total", fallback_model=missing, **shared) + assert [sample.value for sample in lost] == [1.0] + + +@dataclass(frozen=True, slots=True) +class _Budget: + remaining: float + total: float + hours: float + + +def _budget(samples: Sequence[Sample], scope: str, label: str, identity: str) -> _Budget | None: + remaining: Final = _series(samples, f"litellm_remaining_{scope}_budget_metric", **{label: identity}) + total: Final = _series(samples, f"litellm_{scope}_max_budget_metric", **{label: identity}) + hours: Final = _series(samples, f"litellm_{scope}_budget_remaining_hours_metric", **{label: identity}) + if len(remaining) != 1 or len(total) != 1 or len(hours) != 1: + return None + return _Budget(remaining[0].value, total[0].value, hours[0].value) + + +def _reconciled(rig: _Rig, scope: str, label: str, identity: str, info: str, field: str, lookup: str) -> _Budget: + def read() -> tuple[_Budget | None, float]: + record: Final = object_value(rig.gateway.get(info, {field: lookup})[_INFO_FIELDS[info]]) + return _budget(scrape(rig.gateway), scope, label, identity), float(str(record["max_budget"])) - float( + str(record["spend"]) + ) + + budget, remaining = eventually( + read, + lambda state: ( + state[0] is not None and state[0].remaining < 10.0 and abs(state[1] - state[0].remaining) <= 0.001 + ), + seconds=60, + ) + assert budget is not None + assert abs(remaining - budget.remaining) <= 0.001 + return budget + + +_INFO_FIELDS: Final = {"/team/info": "team_info", "/key/info": "info", "/user/info": "user_info"} + + +def test_a_team_call_exports_remaining_max_and_hours_gauges_matching_team_info(rig: _Rig) -> None: + with rig.gateway.scenario() as scenario: + team: Final = scenario.team(max_budget=10, budget_duration="7d") + key: Final = scenario.key(team_id=team) + assert _ask(rig, _GOOD, key).status_code == 200 + assert len(rig.upstream.drain()) == 1 + budget: Final = _reconciled(rig, "team", "team", team, "/team/info", "team_id", team) + assert budget.total == 10.0 + assert 0 < budget.hours <= 168 + + +def test_a_key_call_exports_remaining_max_and_hours_gauges_matching_key_info(rig: _Rig) -> None: + with rig.gateway.scenario() as scenario: + key: Final = scenario.key(max_budget=10, budget_duration="7d") + assert _ask(rig, _GOOD, key).status_code == 200 + assert len(rig.upstream.drain()) == 1 + hashed: Final = sha256(key.encode()).hexdigest() + budget: Final = _reconciled(rig, "api_key", "hashed_api_key", hashed, "/key/info", "key", key) + assert budget.total == 10.0 + assert 0 <= budget.hours <= 168 + + +def test_a_user_call_exports_remaining_max_and_hours_gauges_matching_user_info(rig: _Rig) -> None: + with rig.gateway.scenario() as scenario: + user: Final = f"prometheus-user-{uuid.uuid4().hex}" + scenario.user(user_id=user, max_budget=10, budget_duration="7d") + key: Final = scenario.key(user_id=user) + assert _ask(rig, _GOOD, key).status_code == 200 + assert len(rig.upstream.drain()) == 1 + budget: Final = _reconciled(rig, "user", "user", user, "/user/info", "user_id", user) + assert budget.total == 10.0 + assert 0 <= budget.hours <= 168 + + +def test_a_user_email_labels_the_spend_and_failed_request_series_of_its_keys(rig: _Rig) -> None: + with rig.gateway.scenario() as scenario: + email: Final = f"prometheus-{uuid.uuid4().hex}@example.com" + user: Final = f"prometheus-email-{uuid.uuid4().hex}" + scenario.user(user_id=user, user_email=email) + key: Final = scenario.key(user_id=user) + assert _ask(rig, _GOOD, key).status_code == 200 + assert len(rig.upstream.drain()) == 1 + spend: Final = _until(rig, "litellm_spend_metric_total", user_email=email) + assert [(sample.labels["user"], sample.value) for sample in spend] == [(user, pytest.approx(0.005))] + assert email in label_values(scrape(rig.gateway)) + assert _ask(rig, _FAILING, key).status_code == 429 + assert len(rig.upstream.drain()) == 1 + failed: Final = _until( + rig, "litellm_proxy_failed_requests_metric_total", user_email=email, requested_model=_FAILING + ) + assert [(sample.labels["user"], sample.value) for sample in failed] == [(user, 1.0)] diff --git a/tests/integration/providers/_mantle_gpt_prompt_cache_support.py b/tests/integration/providers/_mantle_gpt_prompt_cache_support.py new file mode 100644 index 00000000000..8f497fea74f --- /dev/null +++ b/tests/integration/providers/_mantle_gpt_prompt_cache_support.py @@ -0,0 +1,959 @@ +import base64 +import binascii +import json +import math +import os +import uuid +from collections.abc import Callable, Iterable, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Final, Literal, Protocol +from urllib.parse import urlsplit + +import anthropic +import httpx +import openai +from integration._support.client import Gateway, Scenario, eventually +from integration._support.database import read_rows +from integration._support.responses_vendor import answer, error, newest_marker, sse +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with + +Endpoint = Literal["chat", "responses", "messages"] + + +JSON: Final = TypeAdapter(dict[str, JsonValue]) + + +ITEMS: Final = TypeAdapter(list[dict[str, JsonValue]]) + + +SIGNING_KEY: Final = os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt") + + +TOKEN: Final = "synthetic-mantle-bearer" + + +GPT: Final = "bedrock_mantle/openai.gpt-5.6-sol" + + +GPT_REGION: Final = "bedrock_mantle/us-east-1/openai.gpt-5.6-sol" + +GPT_BARE: Final = "openai.gpt-5.6-sol" + + +GPT_UNFLAGGED_ROW: Final = "bedrock_mantle/openai.gpt-5.4" + + +GPT_FLAGGED_ROW: Final = "bedrock_mantle/openai.gpt-6-luna" + + +GPT_ODD_STRING_FLAG: Final = "bedrock_mantle/openai.gpt-5.5" + + +GPT_ODD_INT_FLAG: Final = "bedrock_mantle/openai.gpt-daybreak-blue-5.6-sol" +ODD_FLAGS: Final[tuple[tuple[str, str, JsonValue], ...]] = ( + ("string-true", GPT_ODD_STRING_FLAG, "true"), + ("int-one", GPT_ODD_INT_FLAG, 1), + ("int-zero", GPT_ODD_INT_FLAG, 0), + ("5kb-string", GPT_ODD_STRING_FLAG, "x" * 5120), +) + + +CLAUDE: Final = "bedrock_mantle/anthropic.claude-haiku-4-5" + + +AZURE: Final = "azure/gpt-5.6" + + +THIRD_PARTY: Final = "openai/gpt-5.6" + + +SYSTEM: Final = "Reply with the signature the user gives you." + + +SYSTEM_POINT: Final[list[JsonValue]] = [{"location": "message", "role": "system"}] + + +EXPLICIT: Final[dict[str, JsonValue]] = {"mode": "explicit"} + + +IMPLICIT: Final[dict[str, JsonValue]] = {"mode": "implicit"} + + +EPHEMERAL: Final[dict[str, JsonValue]] = {"type": "ephemeral"} + + +NO_CACHE: Final[dict[str, JsonValue]] = {"no-cache": True} + + +MAX_TOKENS: Final = 64 + + +INPUT_TOKENS: Final = 2730 + + +CACHED_TOKENS: Final = 1024 + + +WRITTEN_TOKENS: Final = 1700 + + +UNCACHED_TOKENS: Final = INPUT_TOKENS - CACHED_TOKENS - WRITTEN_TOKENS + + +OUTPUT_TOKENS: Final = 7 + + +USAGE: Final[dict[str, JsonValue]] = { + "input_tokens": INPUT_TOKENS, + "output_tokens": OUTPUT_TOKENS, + "total_tokens": INPUT_TOKENS + OUTPUT_TOKENS, + "input_tokens_details": {"cached_tokens": CACHED_TOKENS, "cache_write_tokens": WRITTEN_TOKENS}, + "output_tokens_details": {"reasoning_tokens": 0}, +} + + +ANTHROPIC_USAGE: Final[dict[str, JsonValue]] = { + "input_tokens": UNCACHED_TOKENS, + "output_tokens": OUTPUT_TOKENS, + "cache_read_input_tokens": CACHED_TOKENS, + "cache_creation_input_tokens": WRITTEN_TOKENS, +} + + +CHAT_USAGE: Final[dict[str, JsonValue]] = { + "prompt_tokens": INPUT_TOKENS, + "completion_tokens": OUTPUT_TOKENS, + "total_tokens": INPUT_TOKENS + OUTPUT_TOKENS, + "prompt_tokens_details": {"cached_tokens": CACHED_TOKENS}, +} + + +COST_MAP: Final = JSON.validate_json(Path("model_prices_and_context_window.json").read_bytes()) + + +RESPONSES_PATH: Final = "/openai/v1/responses" + + +ANTHROPIC_PATH: Final = "/anthropic/v1/messages" + + +BURST: Final = 30 + + +ENDPOINTS: Final[tuple[Endpoint, ...]] = ("chat", "responses", "messages") + + +def cost_rate(model: str, field: str) -> float: + value: Final = JSON.validate_python(COST_MAP[model])[field] + assert isinstance(value, int | float), (model, field, value) + return float(value) + + +def expected_spend(model: str) -> float: + return ( + UNCACHED_TOKENS * cost_rate(model, "input_cost_per_token") + + CACHED_TOKENS * cost_rate(model, "cache_read_input_token_cost") + + WRITTEN_TOKENS * cost_rate(model, "cache_creation_input_token_cost") + + OUTPUT_TOKENS * cost_rate(model, "output_cost_per_token") + ) + + +def fresh_marker() -> str: + return uuid.uuid4().hex + + +def prompt_text(marker: str) -> str: + return f"Return the signature marker-{marker}." + + +def valid_options(options: JsonValue) -> bool: + if options is None: + return True + if not isinstance(options, dict) or not set(options) <= {"mode", "ttl"}: + return False + return options.get("mode", "implicit") in ("implicit", "explicit") + + +def message_item(identity: str, text: str, status: str) -> dict[str, JsonValue]: + return { + "id": f"msg_{identity}", + "type": "message", + "role": "assistant", + "status": status, + "content": [{"type": "output_text", "text": text, "annotations": []}] if status == "completed" else [], + } + + +def responses_object(identity: str, model: str, text: str) -> dict[str, JsonValue]: + return { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": model, + "output": [message_item(identity, text, "completed")], + "usage": USAGE, + } + + +def responses_stream(identity: str, model: str, text: str) -> tuple[bytes, ...]: + response: Final = responses_object(identity, model, text) + return ( + sse({"type": "response.created", "sequence_number": 0, "response": {**response, "status": "in_progress"}}), + sse( + { + "type": "response.output_item.added", + "sequence_number": 1, + "output_index": 0, + "item": message_item(identity, text, "in_progress"), + } + ), + sse( + { + "type": "response.output_text.delta", + "sequence_number": 2, + "item_id": f"msg_{identity}", + "output_index": 0, + "content_index": 0, + "delta": text, + } + ), + sse( + { + "type": "response.output_item.done", + "sequence_number": 3, + "output_index": 0, + "item": message_item(identity, text, "completed"), + } + ), + sse({"type": "response.completed", "sequence_number": 4, "response": response}), + ) + + +def anthropic_message(identity: str, model: str, text: str) -> dict[str, JsonValue]: + return { + "id": identity, + "type": "message", + "role": "assistant", + "model": model, + "content": [{"type": "text", "text": text}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": ANTHROPIC_USAGE, + } + + +def anthropic_stream(identity: str, model: str, text: str) -> tuple[bytes, ...]: + started: Final = {**anthropic_message(identity, model, text), "content": [], "stop_reason": None} + return ( + sse({"type": "message_start", "message": {**started, "usage": {**ANTHROPIC_USAGE, "output_tokens": 0}}}), + sse({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}), + sse({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}}), + sse({"type": "content_block_stop", "index": 0}), + sse( + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": OUTPUT_TOKENS}, + } + ), + sse({"type": "message_stop"}), + ) + + +def chat_completion(identity: str, model: str, text: str) -> dict[str, JsonValue]: + return { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": model, + "choices": [{"index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop"}], + "usage": CHAT_USAGE, + } + + +def issued_id(prefix: str, marker: str | None) -> str: + return f"{prefix}{marker or 'unmarked'}-{uuid.uuid4().hex[:8]}" + + +def mantle_peer(*, pause: float = 0) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.headers.get("authorization") == f"Bearer {TOKEN}", request.headers + body: Final = JSON.validate_json(request.body) + model: Final = str(body["model"]) + marker: Final = newest_marker(request.body.decode()) + text: Final = answer(marker) + streaming: Final = body.get("stream") is True + path: Final = urlsplit(request.target).path + if path == ANTHROPIC_PATH: + identity: Final = issued_id("msg_", marker) + if streaming: + return Reply( + content_type="text/event-stream", + chunks=anthropic_stream(identity, model, text), + pause_between_chunks=pause, + ) + return Reply(body=json.dumps(anthropic_message(identity, model, text)).encode()) + assert path == RESPONSES_PATH, request.target + if not valid_options(body.get("prompt_cache_options")): + return error(400, "Invalid prompt_cache_options", "invalid_prompt_cache_options") + response_id: Final = issued_id("resp_", marker) + if streaming: + return Reply( + content_type="text/event-stream", + chunks=responses_stream(response_id, model, text), + pause_between_chunks=pause, + ) + return Reply(body=json.dumps(responses_object(response_id, model, text)).encode()) + + return respond + + +def failing_peer(status: int, message: str, code: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert urlsplit(request.target).path == RESPONSES_PATH, request.target + return error(status, message, code) + + return respond + + +def openai_shaped_peer() -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + body: Final = JSON.validate_json(request.body) + model: Final = str(body["model"]) + marker: Final = newest_marker(request.body.decode()) + text: Final = answer(marker) + path: Final = urlsplit(request.target).path + if path.endswith("/chat/completions"): + return Reply(body=json.dumps(chat_completion(issued_id("chatcmpl-", marker), model, text)).encode()) + assert path.endswith("/responses"), request.target + return Reply(body=json.dumps(responses_object(issued_id("resp_", marker), model, text)).encode()) + + return respond + + +def system_item(*, marked: bool, endpoint: Endpoint) -> dict[str, JsonValue]: + part: Final[dict[str, JsonValue]] = {"type": "input_text", "text": SYSTEM} + content: Final[list[JsonValue]] = [{**part, "prompt_cache_breakpoint": EXPLICIT} if marked else part] + if endpoint == "responses": + return {"role": "system", "content": content} + return {"type": "message", "role": "system", "content": content} + + +def user_item(text: str, endpoint: Endpoint) -> dict[str, JsonValue]: + if endpoint == "responses": + return {"role": "user", "content": text} + return {"type": "message", "role": "user", "content": [{"type": "input_text", "text": text}]} + + +def expected_wire( + model: str, + prompt: str, + *, + endpoint: Endpoint, + marked: bool, + options: JsonValue | None = IMPLICIT, + **extra: JsonValue, +) -> dict[str, JsonValue]: + body: Final[dict[str, JsonValue]] = { + "model": model.rsplit("/", 1)[-1], + "input": [system_item(marked=marked, endpoint=endpoint), user_item(prompt, endpoint)], + "max_output_tokens": MAX_TOKENS, + **extra, + } + return body if options is None else {**body, "prompt_cache_options": options} + + +def body_of(request: Request) -> dict[str, JsonValue]: + return JSON.validate_json(request.body) + + +def without_stream(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return {key: value for key, value in body.items() if key != "stream"} + + +def only_received(wire: Wire) -> Request: + (request,) = wire.drain() + return request + + +def assert_wire(request: Request, expected: Mapping[str, JsonValue], *, streaming: bool) -> None: + body: Final = body_of(request) + assert urlsplit(request.target).path == RESPONSES_PATH, request.target + assert without_stream(body) == expected, json.dumps(body, sort_keys=True) + assert (body.get("stream") is True) is streaming, body.get("stream") + + +def breakpoint_count(body: Mapping[str, JsonValue]) -> int: + parts: Final = ( + part + for item in ITEMS.validate_python(body["input"]) + for part in (item["content"] if isinstance(item["content"], list) else ()) + ) # comprehension-ok: nested input items + return sum(1 for part in parts if isinstance(part, dict) and "prompt_cache_breakpoint" in part) + + +def endpoint_path(endpoint: Endpoint) -> str: + match endpoint: + case "chat": + return "/v1/chat/completions" + case "responses": + return "/v1/responses" + case "messages": + return "/v1/messages" + + +def request_body( + endpoint: Endpoint, model: str, prompt: str, *, stream: bool = False, **extra: JsonValue +) -> dict[str, JsonValue]: + match endpoint: + case "chat": + return { + "model": model, + "messages": [{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + "max_tokens": MAX_TOKENS, + "stream": stream, + **({"stream_options": {"include_usage": True}} if stream else {}), + **extra, + } + case "responses": + return { + "model": model, + "input": [{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + "max_output_tokens": MAX_TOKENS, + "stream": stream, + **extra, + } + case "messages": + return { + "model": model, + "system": SYSTEM, + "messages": [{"role": "user", "content": prompt}], + "max_tokens": MAX_TOKENS, + "stream": stream, + **extra, + } + + +@dataclass(frozen=True, slots=True) +class Outcome: + status: int + call_id: str + response_id: str + text: str + usage: dict[str, JsonValue] + headers: Mapping[str, str] + raw: str + + +def sse_payloads(lines: Iterable[str]) -> tuple[dict[str, JsonValue], ...]: + return tuple(JSON.validate_json(line[6:]) for line in lines if line.startswith("data: ") and line != "data: [DONE]") + + +def unwrapped(identity: str) -> str | None: + try: + decoded: Final = base64.b64decode(identity.removeprefix("resp_"), validate=True).decode() + except (binascii.Error, UnicodeDecodeError): + return None + return decoded.rsplit("response_id:", 1)[1] if decoded.startswith("litellm:") else None + + +def upstream_id_of(identity: str) -> str: + managed: Final = decrypt_if_encrypted_with(identity.removeprefix("resp_"), SIGNING_KEY) + wrapped: Final = identity if managed is None else managed.rsplit("response_id:", 1)[1].split(";", 1)[0] + inner: Final = unwrapped(wrapped) + return wrapped if inner is None else inner + + +def row_key(request_id: str) -> str: + return upstream_id_of(request_id.split("_cache_hit", 1)[0]) + + +def caller_sees_upstream_id(endpoint: Endpoint, *, stream: bool) -> bool: + return (endpoint, stream) != ("messages", True) + + +def usage_of(payload: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + usage: Final = payload.get("usage") + return JSON.validate_python(usage) if isinstance(usage, dict) else {} + + +def chat_text(payload: Mapping[str, JsonValue]) -> str: + (choice,) = ITEMS.validate_python(payload["choices"]) + return str(JSON.validate_python(choice["message"])["content"]) + + +def chat_stream_fields(payloads: Sequence[Mapping[str, JsonValue]]) -> tuple[str, str, dict[str, JsonValue]]: + (identity,) = {str(chunk["id"]) for chunk in payloads if "id" in chunk} + choices: Final = tuple(ITEMS.validate_python(chunk["choices"]) for chunk in payloads if chunk.get("choices")) + deltas: Final = tuple(JSON.validate_python(choice[0]["delta"]) for choice in choices if choice) + text: Final = "".join(str(delta["content"]) for delta in deltas if isinstance(delta.get("content"), str)) + usages: Final = tuple(usage_of(chunk) for chunk in payloads if isinstance(chunk.get("usage"), dict)) + return identity, text, usages[-1] if usages else {} + + +def responses_text(payload: Mapping[str, JsonValue]) -> str: + items: Final = ITEMS.validate_python(payload["output"]) + parts: Final = tuple(ITEMS.validate_python(item["content"]) for item in items if item.get("type") == "message") + return "".join(str(part["text"]) for part in parts[0] if part.get("type") == "output_text") + + +def completed_response(events: Iterable[Mapping[str, JsonValue]]) -> dict[str, JsonValue]: + (completed,) = tuple(event for event in events if event.get("type") == "response.completed") + return JSON.validate_python(completed["response"]) + + +def messages_text(payload: Mapping[str, JsonValue]) -> str: + return "".join(str(block["text"]) for block in ITEMS.validate_python(payload["content"]) if "text" in block) + + +def messages_stream_fields(payloads: Sequence[Mapping[str, JsonValue]]) -> tuple[str, str, dict[str, JsonValue]]: + (started,) = tuple(payload for payload in payloads if payload.get("type") == "message_start") + message: Final = JSON.validate_python(started["message"]) + deltas: Final = tuple(JSON.validate_python(payload["delta"]) for payload in payloads if "delta" in payload) + text: Final = "".join(str(delta["text"]) for delta in deltas if delta.get("type") == "text_delta") + final_usages: Final = tuple(usage_of(payload) for payload in payloads if payload.get("type") == "message_delta") + return str(message["id"]), text, {**usage_of(message), **(final_usages[-1] if final_usages else {})} + + +def parse_outcome( + endpoint: Endpoint, *, stream: bool, status: int, headers: Mapping[str, str], lines: tuple[str, ...] +) -> Outcome: + call_id: Final = headers.get("x-litellm-call-id", "") + raw: Final = "\n".join(lines) + if status != 200: + return Outcome(status, call_id, "", "", {}, headers, raw) + payloads: Final = sse_payloads(lines) if stream else (JSON.validate_json(raw),) + match endpoint, stream: + case "chat", True: + identity, text, usage = chat_stream_fields(payloads) + return Outcome(status, call_id, upstream_id_of(identity), text, usage, headers, raw) + case "chat", False: + return Outcome( + status, + call_id, + upstream_id_of(str(payloads[0]["id"])), + chat_text(payloads[0]), + usage_of(payloads[0]), + headers, + raw, + ) + case "responses", True: + completed: Final = completed_response(payloads) + return Outcome( + status, + call_id, + upstream_id_of(str(completed["id"])), + responses_text(completed), + usage_of(completed), + headers, + raw, + ) + case "responses", False: + return Outcome( + status, + call_id, + upstream_id_of(str(payloads[0]["id"])), + responses_text(payloads[0]), + usage_of(payloads[0]), + headers, + raw, + ) + case "messages", True: + identity, text, usage = messages_stream_fields(payloads) + return Outcome(status, call_id, upstream_id_of(identity), text, usage, headers, raw) + case "messages", False: + return Outcome( + status, + call_id, + upstream_id_of(str(payloads[0]["id"])), + messages_text(payloads[0]), + usage_of(payloads[0]), + headers, + raw, + ) + raise AssertionError((endpoint, stream)) + + +def send(gateway: Gateway, endpoint: Endpoint, body: Mapping[str, JsonValue], *, key: str | None = None) -> Outcome: + stream: Final = body.get("stream") is True + headers: Final = {"Authorization": f"Bearer {gateway.key if key is None else key}"} + with gateway.client.stream("POST", endpoint_path(endpoint), json=body, headers=headers, timeout=60) as response: + lines: Final = tuple(line for line in response.iter_lines() if line) + return parse_outcome(endpoint, stream=stream, status=response.status_code, headers=response.headers, lines=lines) + + +def send_raw(gateway: Gateway, endpoint: Endpoint, content: str, *, key: str | None = None) -> Outcome: + headers: Final = { + "Authorization": f"Bearer {gateway.key if key is None else key}", + "Content-Type": "application/json", + } + response: Final = gateway.client.post( + endpoint_path(endpoint), content=content.encode(), headers=headers, timeout=60 + ) + lines: Final = tuple(line for line in response.text.splitlines() if line) + return parse_outcome(endpoint, stream=False, status=response.status_code, headers=response.headers, lines=lines) + + +def matching_rows(name: str, markers: frozenset[str]) -> tuple[dict[str, JsonValue], ...]: + rows: Final = read_rows( + 'SELECT request_id, status, spend, prompt_tokens, completion_tokens, cache_hit FROM "LiteLLM_SpendLogs" ' + "WHERE model_group = %s", + (name,), + ) + return tuple(row for row in rows if any(marker in row_key(str(row["request_id"])) for marker in markers)) + + +def spend_rows( + name: str, markers: frozenset[str], *, expected: int, seconds: float = 90 +) -> tuple[dict[str, JsonValue], ...]: + return eventually(lambda: matching_rows(name, markers), lambda found: len(found) >= expected, seconds=seconds) + + +def success_row(name: str, *needles: str) -> dict[str, JsonValue]: + (row,) = tuple(row for row in spend_rows(name, frozenset(needles), expected=1) if row["cache_hit"] != "True") + assert row["status"] == "success", row + return row + + +def failure_row(call_id: str) -> dict[str, JsonValue]: + assert call_id, "No call id to look the failure row up by" + rows: Final = eventually( + lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (call_id,)), + lambda found: len(found) >= 1, + seconds=90, + ) + (row,) = tuple(rows) + assert row["status"] == "failure", row + return row + + +def assert_priced_row(row: Mapping[str, JsonValue], model: str) -> None: + assert row["prompt_tokens"] == INPUT_TOKENS, row + assert row["completion_tokens"] == OUTPUT_TOKENS, row + assert math.isclose(float(str(row["spend"])), expected_spend(model), rel_tol=1e-9), (row, expected_spend(model)) + + +def assert_usage(endpoint: Endpoint, usage: Mapping[str, JsonValue]) -> None: + match endpoint: + case "chat": + assert usage["prompt_tokens"] == INPUT_TOKENS, usage + assert usage["completion_tokens"] == OUTPUT_TOKENS, usage + details: Final = JSON.validate_python(usage["prompt_tokens_details"]) + assert details["cached_tokens"] == CACHED_TOKENS, usage + assert details["cache_write_tokens"] == WRITTEN_TOKENS, usage + assert details["cache_creation_tokens"] == WRITTEN_TOKENS, usage + case "responses": + assert usage["input_tokens"] == INPUT_TOKENS, usage + assert usage["output_tokens"] == OUTPUT_TOKENS, usage + input_details: Final = JSON.validate_python(usage["input_tokens_details"]) + assert input_details["cached_tokens"] == CACHED_TOKENS, usage + assert input_details["cache_write_tokens"] == WRITTEN_TOKENS, usage + case "messages": + assert usage["input_tokens"] == UNCACHED_TOKENS, usage + assert usage["output_tokens"] == OUTPUT_TOKENS, usage + assert usage["cache_read_input_tokens"] == CACHED_TOKENS, usage + assert usage["cache_creation_input_tokens"] == WRITTEN_TOKENS, usage + + +def assert_answered(outcome: Outcome, marker: str) -> None: + assert outcome.status == 200, (outcome.status, outcome.raw) + assert outcome.text == answer(marker), outcome.raw + + +def deployment( + scenario: Scenario, + wire: Wire, + model: str, + *, + model_info: Mapping[str, JsonValue] | None = None, + points: JsonValue = SYSTEM_POINT, + **extra: JsonValue, +) -> str: + return scenario.model( + model_info=model_info, + model=model, + api_base=wire.url, + api_key=TOKEN, + aws_region_name="us-east-1", + cache_control_injection_points=points, + **extra, + ) + + +def settled(gateway: Gateway, name: str, wire: Wire, *, accepted: frozenset[int] = frozenset({200})) -> None: + eventually( + lambda: tuple( + gateway.request( + "POST", "/v1/chat/completions", request_body("chat", name, prompt_text(fresh_marker()), cache=NO_CACHE) + ).status_code + for _ in range(12) + ), + lambda codes: all(code in accepted for code in codes), + seconds=90, + ) + wire.drain() + + +def mantle_deployment( + gateway: Gateway, + scenario: Scenario, + wire: Wire, + model: str = GPT, + *, + model_info: Mapping[str, JsonValue] | None = None, + **extra: JsonValue, +) -> str: + name: Final = deployment(scenario, wire, model, model_info=model_info, **extra) + settled(gateway, name, wire) + return name + + +def observe( + gateway: Gateway, wire: Wire, endpoint: Endpoint, body: Mapping[str, JsonValue], *, key: str | None = None +) -> tuple[Outcome, Request]: + wire.drain() + outcome: Final = send(gateway, endpoint, body, key=key) + return outcome, only_received(wire) + + +def marked_cell(gateway: Gateway, endpoint: Endpoint, *, stream: bool) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + outcome, received = observe( + gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker), stream=stream) + ) + assert_answered(outcome, marker) + assert_wire(received, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=stream) + assert_usage(endpoint, outcome.usage) + row: Final = success_row(name, marker) + assert_priced_row(row, GPT) + if caller_sees_upstream_id(endpoint, stream=stream): + assert row_key(str(row["request_id"])) == outcome.response_id, (row, outcome.response_id) + if not stream: + assert math.isclose(float(outcome.headers["x-litellm-response-cost"]), expected_spend(GPT), rel_tol=1e-9), ( + outcome.headers + ) + + +def openai_client(gateway: Gateway) -> openai.OpenAI: + return openai.OpenAI( + base_url=f"{gateway.client.base_url}/v1", + api_key=gateway.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=60), + ) + + +def async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI( + base_url=f"{gateway.client.base_url}/v1", + api_key=gateway.key, + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=60), + ) + + +def anthropic_client(gateway: Gateway) -> anthropic.Anthropic: + return anthropic.Anthropic( + base_url=str(gateway.client.base_url), + api_key=gateway.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=60), + ) + + +def async_anthropic_client(gateway: Gateway) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic( + base_url=str(gateway.client.base_url), + api_key=gateway.key, + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=60), + ) + + +class Dumpable(Protocol): + def model_dump(self, *, exclude_none: bool = ...) -> Mapping[str, object]: ... + + +def dump_model(value: Dumpable) -> dict[str, JsonValue]: + return JSON.validate_python(value.model_dump(exclude_none=True)) + + +def sdk_outcome(endpoint: Endpoint, payloads: Sequence[Mapping[str, JsonValue]], *, stream: bool) -> Outcome: + match endpoint, stream: + case "chat", True: + identity, text, usage = chat_stream_fields(payloads) + return Outcome(200, "", upstream_id_of(identity), text, usage, {}, "") + case "chat", False: + return Outcome( + 200, "", upstream_id_of(str(payloads[0]["id"])), chat_text(payloads[0]), usage_of(payloads[0]), {}, "" + ) + case "responses", _: + completed: Final = completed_response(payloads) if stream else dict(payloads[0]) + return Outcome( + 200, + "", + upstream_id_of(str(completed["id"])), + responses_text(completed), + usage_of(completed), + {}, + "", + ) + case "messages", True: + identity, text, usage = messages_stream_fields(payloads) + return Outcome(200, "", upstream_id_of(identity), text, usage, {}, "") + case "messages", False: + return Outcome( + 200, + "", + upstream_id_of(str(payloads[0]["id"])), + messages_text(payloads[0]), + usage_of(payloads[0]), + {}, + "", + ) + raise AssertionError((endpoint, stream)) + + +def sdk_call(gateway: Gateway, endpoint: Endpoint, name: str, prompt: str, *, stream: bool) -> Outcome: + match endpoint, stream: + case "chat", False: + completion: Final = openai_client(gateway).chat.completions.create( + model=name, + messages=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_tokens=MAX_TOKENS, + ) + return sdk_outcome(endpoint, (dump_model(completion),), stream=False) + case "chat", True: + chunks: Final = openai_client(gateway).chat.completions.create( + model=name, + messages=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_tokens=MAX_TOKENS, + stream=True, + stream_options={"include_usage": True}, + ) + return sdk_outcome(endpoint, tuple(dump_model(chunk) for chunk in chunks), stream=True) + case "responses", False: + response: Final = openai_client(gateway).responses.create( + model=name, + input=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_output_tokens=MAX_TOKENS, + ) + return sdk_outcome(endpoint, (dump_model(response),), stream=False) + case "responses", True: + events: Final = openai_client(gateway).responses.create( + model=name, + input=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_output_tokens=MAX_TOKENS, + stream=True, + ) + return sdk_outcome(endpoint, tuple(dump_model(event) for event in events), stream=True) + case "messages", False: + message: Final = anthropic_client(gateway).messages.create( + model=name, system=SYSTEM, messages=[{"role": "user", "content": prompt}], max_tokens=MAX_TOKENS + ) + return sdk_outcome(endpoint, (dump_model(message),), stream=False) + case "messages", True: + with anthropic_client(gateway).messages.stream( + model=name, system=SYSTEM, messages=[{"role": "user", "content": prompt}], max_tokens=MAX_TOKENS + ) as events_stream: + raw_events: Final = tuple(dump_model(event) for event in events_stream) + return sdk_outcome(endpoint, raw_events, stream=True) + raise AssertionError((endpoint, stream)) + + +async def async_sdk_call(gateway: Gateway, endpoint: Endpoint, name: str, prompt: str) -> Outcome: + match endpoint: + case "chat": + completion: Final = await async_openai_client(gateway).chat.completions.create( + model=name, + messages=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_tokens=MAX_TOKENS, + ) + return sdk_outcome(endpoint, (dump_model(completion),), stream=False) + case "responses": + response: Final = await async_openai_client(gateway).responses.create( + model=name, + input=[{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], + max_output_tokens=MAX_TOKENS, + ) + return sdk_outcome(endpoint, (dump_model(response),), stream=False) + case "messages": + message: Final = await async_anthropic_client(gateway).messages.create( + model=name, system=SYSTEM, messages=[{"role": "user", "content": prompt}], max_tokens=MAX_TOKENS + ) + return sdk_outcome(endpoint, (dump_model(message),), stream=False) + + +def assert_marked_sdk_cell( + wire: Wire, endpoint: Endpoint, name: str, marker: str, outcome: Outcome, *, stream: bool +) -> None: + assert_answered(outcome, marker) + assert_wire( + only_received(wire), expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=stream + ) + assert_usage(endpoint, outcome.usage) + assert_priced_row(success_row(name, marker), GPT) + + +HOSTILE_OPTIONS: Final[dict[str, JsonValue]] = { + "int": 5, + "list": [{"mode": "explicit"}], + "empty-string": "", + "5kb-string": "x" * 5120, +} + + +MALFORMED_POINTS: Final[dict[str, JsonValue]] = { + "null": None, + "string": "system", + "int": 5, + "dict": {"location": "message", "role": "system"}, + "string-list": ["system"], + "no-location": [{"role": "system"}], +} +MIXED_POINTS: Final[list[JsonValue]] = ["system", *SYSTEM_POINT, 3] + + +Step = tuple[Endpoint, bool, str] + + +def plan_burst(size: int) -> tuple[Step, ...]: + return tuple((ENDPOINTS[index % 3], index % 3 == 0, fresh_marker()) for index in range(size)) + + +def burst(gateway: Gateway, name: str, size: int) -> tuple[tuple[str, Outcome], ...]: + def call(step: Step) -> tuple[str, Outcome]: + endpoint, stream, marker = step + return marker, send(gateway, endpoint, request_body(endpoint, name, prompt_text(marker), stream=stream)) + + with ThreadPoolExecutor(max_workers=10) as pool: + return tuple(pool.map(call, plan_burst(size))) + + +def assert_marked(request: Request, *, marked: bool) -> None: + body: Final = body_of(request) + assert breakpoint_count(body) == (1 if marked else 0), request.body + assert ("prompt_cache_options" in body) is marked, request.body + + +def assert_burst_landed(wire: Wire, name: str, outcomes: Sequence[tuple[str, Outcome]], *, marked: bool) -> None: + for marker, outcome in outcomes: + assert_answered(outcome, marker) + identities: Final = frozenset(outcome.response_id for _, outcome in outcomes) + assert len(identities) == len(outcomes), identities + received: Final = wire.drain() + assert len(received) == len(outcomes), (len(received), len(outcomes)) + for request in received: + assert_marked(request, marked=marked) + markers: Final = frozenset(marker for marker, _ in outcomes) + rows: Final = spend_rows(name, markers, expected=len(outcomes), seconds=120) + assert sorted(row_key(str(row["request_id"])) for row in rows) == sorted(identities), rows diff --git a/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py b/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py new file mode 100644 index 00000000000..f9d4f7f91bb --- /dev/null +++ b/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py @@ -0,0 +1,1072 @@ +import asyncio +import json +import re +import signal +import socket +import threading +import uuid +from collections.abc import Callable, Iterator, Mapping, Sequence +from concurrent.futures import Future, ThreadPoolExecutor +from contextlib import ExitStack +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final + +import anthropic +import httpx +import psutil +import pytest +import yaml +from anthropic.types import MessageCountTokensToolParam, MessageParam +from integration._support.bedrock_runtime_peer import answer, marker_of, target_of +from integration._support.bedrock_runtime_peer import respond as runtime_generation +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value +from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +pytestmark = pytest.mark.timeout(2 * graceful_stop_seconds() + 120) + +_OPUS: Final = "global.anthropic.claude-opus-4-8" +_OPUS_BASE: Final = "anthropic.claude-opus-4-8" +_SONNET: Final = "anthropic.claude-sonnet-4-6" +_NOVA: Final = "amazon.nova-lite-v1:0" +_REGION: Final = "us-east-1" +_OWNED_OPUS: Final = "mantle-opus" +_OWNED_SONNET: Final = "mantle-sonnet" +_OWNED_NOVA: Final = "mantle-nova" +_OWNED_MODELS: Final = MappingProxyType({_OWNED_OPUS: _OPUS, _OWNED_SONNET: _SONNET, _OWNED_NOVA: _NOVA}) +_ACCESS_KEY: Final = "AKIAINTEGRATIONMANTLE" +_SECRET_KEY: Final = "integration-mantle-secret" +_SIGV4_SCOPE: Final = f"/{_REGION}/bedrock/aws4_request" +_MANTLE_TARGET: Final = "/anthropic/v1/messages/count_tokens" +_MANTLE_VERSION: Final = "2023-06-01" +_MANTLE_COUNT: Final = 4242 +_RUNTIME_COUNT: Final = 1345 +_UNSUPPORTED: Final = "The provided model doesn't support counting tokens." +_REJECTION: Final = "scripted mantle rejection" +_COUNT_TARGET: Final = re.compile(r"^/model/(.+)/count-tokens$") +_INVOKE_TARGET: Final = re.compile(r"^/model/(.+)/invoke$") +_SCRIPTED_STATUS: Final = re.compile(r"status=(\d{3})") +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_JSON_LIST: Final = TypeAdapter(list[JsonValue]) + +_SDK_MESSAGES: Final[list[MessageParam]] = [{"role": "user", "content": "Count this message"}] +_MESSAGES: Final = _JSON_LIST.validate_json(json.dumps(_SDK_MESSAGES)) +_SYSTEM: Final = "You are a terse assistant that answers in one sentence" +_SYSTEM_BLOCKS: Final[list[JsonValue]] = [ + {"type": "text", "text": "You are a terse assistant"}, + {"type": "text", "text": "Answer in one sentence"}, +] +_SDK_TOOLS: Final[list[MessageCountTokensToolParam]] = [ + { + "name": "get_weather", + "description": "Look up the current weather for a city", + "input_schema": { + "type": "object", + "properties": {"city": {"type": "string", "description": "City to look up"}}, + "required": ["city"], + }, + } +] +_TOOLS: Final = _JSON_LIST.validate_json(json.dumps(_SDK_TOOLS)) +_GEMINI_BODY: Final[dict[str, JsonValue]] = {"contents": [{"role": "user", "parts": [{"text": "Count this"}]}]} +_GEMINI_MESSAGES: Final[list[JsonValue]] = [{"role": "user", "content": "Count this"}] + + +def _mantle_body(**fields: JsonValue) -> dict[str, JsonValue]: + return {"model": _OPUS_BASE, "messages": _MESSAGES, **fields} + + +_MANTLE_BARE: Final = _mantle_body() +_MANTLE_FULL: Final = _mantle_body(system=_SYSTEM, tools=_TOOLS) + + +def _json_reply(status: int, payload: Mapping[str, JsonValue]) -> Reply: + return Reply(status=status, body=json.dumps(payload).encode()) + + +def _runtime_count(request: Request, model: str) -> Reply: + scripted: Final = _SCRIPTED_STATUS.search(request.body.decode(errors="replace")) + if scripted is not None: + return _json_reply(int(scripted.group(1)), {"message": f"scripted {scripted.group(1)}"}) + if "sonnet" in model: + return _json_reply(200, {"inputTokens": _RUNTIME_COUNT}) + return _json_reply(400, {"message": _UNSUPPORTED}) + + +def _invoke_reply(marker: str) -> Reply: + return _json_reply( + 200, + { + "id": f"msg_{marker}", + "type": "message", + "role": "assistant", + "model": _OPUS_BASE, + "content": [{"type": "text", "text": answer(marker)}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 5, "output_tokens": 3}, + }, + ) + + +def _runtime(request: Request) -> Reply: + target: Final = target_of(request) + counted: Final = _COUNT_TARGET.match(target) + if counted is not None: + return _runtime_count(request, counted.group(1)) + if _INVOKE_TARGET.match(target): + return _invoke_reply(marker_of(request)) + return runtime_generation(request) + + +def _mantle_counted(_request: Request) -> Reply: + return _json_reply(200, {"input_tokens": _MANTLE_COUNT}) + + +def _rejected(status: int) -> Reply: + return _json_reply(status, {"type": "error", "error": {"type": "invalid_request_error", "message": _REJECTION}}) + + +def _rejecting(status: int) -> Callable[[Request], Reply]: + def count(_request: Request) -> Reply: + return _rejected(status) + + return count + + +def _anthropic_message(message: JsonValue) -> bool: + return isinstance(message, dict) and message.get("role") in ("user", "assistant") + + +def _anthropic_tool(tool: JsonValue) -> bool: + return isinstance(tool, dict) and isinstance(tool.get("name"), str) and isinstance(tool.get("input_schema"), dict) + + +def _strict(request: Request) -> Reply: + body: Final = _JSON_OBJECT.validate_json(request.body) + messages: Final = body.get("messages") + tools: Final = body.get("tools", []) + accepted: Final = ( + isinstance(messages, list) + and all(map(_anthropic_message, messages)) + and isinstance(body.get("system", ""), (str, list)) + and isinstance(tools, list) + and all(map(_anthropic_tool, tools)) + ) + return _mantle_counted(request) if accepted else _rejected(400) + + +def _mantle(count: Callable[[Request], Reply] = _mantle_counted) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + if request.target == _MANTLE_TARGET: + return count(request) + return _json_reply(404, {"error": f"unscripted mantle target {request.target}"}) + + return respond + + +def _mantle_environment(port: int) -> Mapping[str, str]: + return {"BEDROCK_MANTLE_API_BASE": f"http://127.0.0.1:{port}"} + + +_INHERITED_BEARER: Final = ("AWS_BEARER_TOKEN_BEDROCK",) + + +def _reserved_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return reserve.getsockname()[1] + + +def _closed_port_url() -> str: + return f"http://127.0.0.1:{_reserved_port()}" + + +@pytest.fixture(scope="module") +def mantle_port() -> int: + return _reserved_port() + + +@pytest.fixture(scope="module") +def counting_proxy(tmp_path_factory: pytest.TempPathFactory, mantle_port: int) -> Iterator[OwnedProxy]: + with ( + gateway_from_environment() as gateway, + owned_proxy_process( + gateway, + tmp_path_factory.mktemp("mantle-count"), + _mantle_environment(mantle_port), + workers=1, + remove_environment=_INHERITED_BEARER, + ) as owned, + ): + yield owned + + +def _litellm_params(model: str, api_base: str) -> dict[str, JsonValue]: + return { + "model": f"bedrock/{model}", + "api_base": api_base, + "aws_access_key_id": _ACCESS_KEY, + "aws_secret_access_key": _SECRET_KEY, + "aws_region_name": _REGION, + } + + +def _owned_config(path: Path, runtime_url: str, settings: Mapping[str, JsonValue]) -> Path: + config: Final = object_value(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + litellm_settings: Final = object_value(config["litellm_settings"]) + path.write_text( + yaml.safe_dump( + { + **config, + "model_list": [ + {"model_name": name, "litellm_params": _litellm_params(model, runtime_url)} + for name, model in _OWNED_MODELS.items() + ], + "litellm_settings": {**litellm_settings, **settings}, + } + ) + ) + return path + + +def _deployment(scenario: Scenario, api_base: str, model: str = _OPUS) -> str: + return scenario.model(model_info=None, api_key=None, **_litellm_params(model, api_base)) + + +def _bare(model: str) -> dict[str, JsonValue]: + return {"model": model, "messages": _MESSAGES} + + +def _full(model: str) -> dict[str, JsonValue]: + return {**_bare(model), "system": _SYSTEM, "tools": _TOOLS} + + +def _count(gateway: Gateway, body: Mapping[str, JsonValue]) -> httpx.Response: + return gateway.request("POST", "/v1/messages/count_tokens", body) + + +def _payload(response: httpx.Response) -> dict[str, JsonValue]: + assert response.status_code == 200, response.text + return _JSON_OBJECT.validate_json(response.content) + + +def _local_count(gateway: Gateway, body: Mapping[str, JsonValue]) -> int: + response: Final = gateway.request("POST", "/utils/token_counter", body, params={"call_endpoint": "false"}) + payload: Final = _payload(response) + total: Final = payload["total_tokens"] + assert payload["tokenizer_type"] not in ("bedrock_api", "bedrock_mantle_api"), response.text + assert isinstance(total, int) and total > 0 and total not in (_MANTLE_COUNT, _RUNTIME_COUNT), response.text + return total + + +def _assert_sigv4(request: Request) -> None: + authorization: Final = request.headers.get("authorization", "") + assert authorization.startswith(f"AWS4-HMAC-SHA256 Credential={_ACCESS_KEY}/"), request.headers + assert _SIGV4_SCOPE in authorization, authorization + + +def _mantle_requests(requests: Sequence[Request]) -> tuple[dict[str, JsonValue], ...]: + for request in requests: + assert (request.method, request.target) == ("POST", _MANTLE_TARGET), request.target + assert request.headers["anthropic-version"] == _MANTLE_VERSION, request.headers + assert request.headers["content-type"] == "application/json", request.headers + _assert_sigv4(request) + return tuple(_JSON_OBJECT.validate_json(request.body) for request in requests) + + +def _mantle_bodies(wire: Wire) -> tuple[dict[str, JsonValue], ...]: + return _mantle_requests(wire.drain()) + + +def _runtime_count_targets(wire: Wire) -> tuple[str, ...]: + counts: Final = tuple(request for request in wire.drain() if _COUNT_TARGET.match(target_of(request))) + for request in counts: + assert request.method == "POST", request.method + _assert_sigv4(request) + return tuple(target_of(request) for request in counts) + + +def _clients(stack: ExitStack, base_url: str, count: int, timeout: float = 30) -> tuple[httpx.Client, ...]: + return tuple( + stack.enter_context(httpx.Client(base_url=base_url, timeout=timeout, trust_env=False)) for _ in range(count) + ) + + +def _counted_on(client: httpx.Client, key: str, body: Mapping[str, JsonValue]) -> tuple[int, JsonValue]: + response: Final = client.post( + "/v1/messages/count_tokens", json=dict(body), headers={"Authorization": f"Bearer {key}"} + ) + return response.status_code, _JSON_OBJECT.validate_json(response.content).get("input_tokens") + + +def _generated_then_counted( + client: httpx.Client, key: str, model: str, body: Mapping[str, JsonValue] +) -> tuple[int, int, JsonValue]: + generated: Final = client.post( + "/v1/messages", + json={ + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": f"Generate before counting marker-{uuid.uuid4().hex}"}], + }, + headers={"Authorization": f"Bearer {key}"}, + ) + return generated.status_code, *_counted_on(client, key, body) + + +def _counted_or_dropped(client: httpx.Client, key: str, body: Mapping[str, JsonValue]) -> tuple[int, JsonValue] | None: + try: + return _counted_on(client, key, body) + except httpx.TransportError: + return None + + +def _probed_then_counted_or_dropped( + client: httpx.Client, key: str, body: Mapping[str, JsonValue], probed: SimpleQueue[int] +) -> tuple[int, tuple[int, JsonValue] | None]: + port: Final = _local_port(client) + probed.put(port) + return port, _counted_or_dropped(client, key, body) + + +def _local_port(client: httpx.Client) -> int: + with client.stream("GET", "/health/liveliness") as response: + port: Final = int(response.extensions["network_stream"].get_extra_info("client_addr")[1]) + response.read() + assert response.status_code == 200, response.text + return port + + +def _accepted_client_ports(pid: int, proxy_port: int) -> frozenset[int]: + return frozenset( + connection.raddr.port + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.raddr and connection.laddr.port == proxy_port + ) + + +def _holding(held: SimpleQueue[str], release: threading.Event, seconds: float) -> Callable[[Request], Reply]: + def hold(request: Request) -> Reply: + held.put(request.target) + assert release.wait(timeout=seconds), "Held count was never released" + return _mantle_counted(request) + + return hold + + +def _anthropic_client(gateway: Gateway) -> anthropic.Anthropic: + return anthropic.Anthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0) + + +def _async_anthropic_client(gateway: Gateway) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0) + + +@pytest.mark.parametrize("system", [_SYSTEM, _SYSTEM_BLOCKS], ids=["string", "blocks"]) +def test_messages_count_tokens_counts_through_mantle_when_the_runtime_cannot( + counting_proxy: OwnedProxy, mantle_port: int, system: JsonValue +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, {**_bare(model), "system": system, "tools": _TOOLS}) + assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text + assert _runtime_count_targets(runtime) == (f"/model/{_OPUS_BASE}/count-tokens",) + assert _mantle_bodies(mantle) == (_mantle_body(system=system, tools=_TOOLS),) + + +def test_messages_count_tokens_without_system_or_tools_sends_a_bare_body_to_mantle( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, _bare(model)) + assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_BARE,) + + +def test_utils_token_counter_call_endpoint_reports_the_mantle_tokenizer( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request( + "POST", "/utils/token_counter", _full(model), params={"call_endpoint": "true"} + ) + payload: Final = _payload(response) + assert (payload["total_tokens"], payload["tokenizer_type"]) == (_MANTLE_COUNT, "bedrock_mantle_api") + assert payload["original_response"] == {"input_tokens": _MANTLE_COUNT}, response.text + assert (payload["request_model"], payload["model_used"]) == (model, _OPUS), response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_FULL,) + + +def test_gemini_count_tokens_route_reports_the_mantle_total(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request("POST", f"/v1beta/models/{model}:countTokens", _GEMINI_BODY) + assert _payload(response) == {"totalTokens": _MANTLE_COUNT, "promptTokensDetails": []}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == ({"model": _OPUS_BASE, "messages": _GEMINI_MESSAGES},) + + +def test_gemini_count_tokens_route_reports_the_runtime_total_for_a_model_bedrock_counts( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url, _SONNET) + response: Final = gateway.request("POST", f"/v1beta/models/{model}:countTokens", _GEMINI_BODY) + assert _payload(response) == {"totalTokens": _RUNTIME_COUNT, "promptTokensDetails": []}, response.text + assert _runtime_count_targets(runtime) == (f"/model/{_SONNET}/count-tokens",) + assert mantle.drain() == () + + +def test_responses_input_tokens_counts_through_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request( + "POST", "/v1/responses/input_tokens", {"model": model, "input": "Count this message"} + ) + assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _MANTLE_COUNT}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_BARE,) + + +def test_anthropic_sdk_count_tokens_counts_through_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + counted: Final = _anthropic_client(gateway).messages.count_tokens( + model=model, messages=_SDK_MESSAGES, system=_SYSTEM, tools=_SDK_TOOLS + ) + assert counted.input_tokens == _MANTLE_COUNT, counted + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_FULL,) + + +def test_async_anthropic_sdk_count_tokens_counts_through_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + counted: Final = asyncio.run( + _async_anthropic_client(gateway).messages.count_tokens(model=model, messages=_SDK_MESSAGES, system=_SYSTEM) + ) + assert counted.input_tokens == _MANTLE_COUNT, counted + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_mantle_body(system=_SYSTEM),) + + +def test_model_the_runtime_counts_never_reaches_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url, _SONNET) + assert _payload(_count(gateway, _full(model))) == {"input_tokens": _RUNTIME_COUNT} + detailed: Final = gateway.request( + "POST", "/utils/token_counter", _full(model), params={"call_endpoint": "true"} + ) + payload: Final = _payload(detailed) + assert (payload["total_tokens"], payload["tokenizer_type"]) == (_RUNTIME_COUNT, "bedrock_api"), detailed.text + assert _runtime_count_targets(runtime) == (f"/model/{_SONNET}/count-tokens",) * 2 + assert mantle.drain() == () + + +def test_non_claude_model_the_runtime_cannot_count_falls_back_locally_without_mantle( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url, _NOVA) + response: Final = _count(gateway, _bare(model)) + assert _payload(response) == {"input_tokens": _local_count(gateway, _bare(model))}, response.text + assert _runtime_count_targets(runtime) == (f"/model/{_NOVA}/count-tokens",) + assert mantle.drain() == () + + +def test_runtime_403_on_a_claude_model_never_reaches_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + body: Final[dict[str, JsonValue]] = { + "model": model, + "messages": [{"role": "user", "content": "status=403 Count this message"}], + } + response: Final = _count(gateway, body) + assert _payload(response) == {"input_tokens": _local_count(gateway, body)}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert mantle.drain() == () + + +def test_unreachable_runtime_falls_back_locally_without_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with wire_server(_mantle(), port=mantle_port) as mantle, gateway.scenario() as scenario: + model: Final = _deployment(scenario, _closed_port_url()) + response: Final = _count(gateway, _full(model)) + assert _payload(response) == {"input_tokens": _local_count(gateway, _full(model))}, response.text + assert mantle.drain() == () + + +def test_bedrock_passthrough_count_tokens_still_answers_the_runtime_rejection( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request("POST", "/bedrock/v1/messages/count_tokens", _full(model)) + assert response.status_code == 400, response.text + assert _UNSUPPORTED in response.text and "input_tokens" not in response.text, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert mantle.drain() == () + + +def test_responses_input_tokens_with_instructions_still_counts_locally( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(_strict), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request( + "POST", + "/v1/responses/input_tokens", + {"model": model, "input": "Count this message", "instructions": "Be terse"}, + ) + payload: Final = _payload(response) + assert len(_runtime_count_targets(runtime)) == 1 + (sent,) = _mantle_bodies(mantle) + messages: Final = sent["messages"] + assert isinstance(messages, list) and messages[0] == {"role": "system", "content": "Be terse"}, sent + assert "system" not in sent, sent + local: Final = _local_count(gateway, {"model": model, "messages": messages}) + assert payload == {"object": "response.input_tokens", "input_tokens": local}, response.text + + +@pytest.mark.parametrize("status", [400, 403, 404, 500, 503]) +def test_messages_count_tokens_falls_back_locally_when_mantle_errors( + counting_proxy: OwnedProxy, mantle_port: int, status: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(_rejecting(status)), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, _full(model)) + assert _payload(response) == {"input_tokens": _local_count(gateway, _full(model))}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_FULL,) + + +@pytest.mark.parametrize( + "body", + [b'{"inputTokens": 7}', b"not json at all", b"{}"], + ids=["runtime_key", "not_json", "empty_object"], +) +def test_messages_count_tokens_falls_back_locally_when_mantle_answers_without_input_tokens( + counting_proxy: OwnedProxy, mantle_port: int, body: bytes +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(lambda _request: Reply(body=body)), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, _bare(model)) + assert _payload(response) == {"input_tokens": _local_count(gateway, _bare(model))}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_BARE,) + + +def test_messages_count_tokens_falls_back_locally_when_mantle_is_unreachable(counting_proxy: OwnedProxy) -> None: + gateway: Final = counting_proxy.gateway + with wire_server(_runtime) as runtime, gateway.scenario() as scenario: + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, _full(model)) + assert _payload(response) == {"input_tokens": _local_count(gateway, _full(model))}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + + +def test_messages_count_tokens_falls_back_locally_when_mantle_never_answers( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + held: Final[SimpleQueue[str]] = SimpleQueue() + release: Final = threading.Event() + + def hold(request: Request) -> Reply: + held.put(request.target) + release.wait(timeout=120) + return Reply(drop_connection=True) + + with ExitStack() as stack: + (client,) = _clients(stack, str(gateway.client.base_url), 1, timeout=120) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=1)) + runtime: Final = stack.enter_context(wire_server(_runtime)) + mantle: Final = stack.enter_context(wire_server(_mantle(hold), port=mantle_port)) + stack.callback(release.set) + scenario: Final = stack.enter_context(gateway.scenario()) + model: Final = _deployment(scenario, runtime.url) + body: Final = _full(model) + local: Final = _local_count(gateway, body) + future: Final = pool.submit(_counted_on, client, gateway.key, body) + eventually(held.qsize, lambda size: size == 1, seconds=30) + assert gateway.request("GET", "/health/liveliness").status_code == 200 + assert _local_count(gateway, body) == local + assert not future.done() + assert future.result(timeout=120) == (200, local) + assert len(_runtime_count_targets(runtime)) == 1 + assert len(mantle.drain()) == 1 + + +def test_messages_count_tokens_falls_back_locally_when_mantle_rejects_a_non_text_system( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(_strict), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, {**_bare(model), "system": 5}) + assert _payload(response) == {"input_tokens": _local_count(gateway, _bare(model))}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_mantle_body(system=5),) + + +@pytest.mark.parametrize( + "tools", + [5, "", "x" * 5120, ["get_weather"]], + ids=["int", "empty_string", "5kb_string", "list_of_strings"], +) +def test_messages_count_tokens_rejects_malformed_tools_without_calling_either_peer( + counting_proxy: OwnedProxy, mantle_port: int, tools: JsonValue +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + refused: Final = _count(gateway, {**_bare(model), "tools": tools}) + assert 400 <= refused.status_code < 600, refused.text + assert "input_tokens" not in refused.text, refused.text + assert runtime.drain() == () and mantle.drain() == () + assert _generated_then_counted(gateway.client, gateway.key, model, _bare(model)) == (200, 200, _MANTLE_COUNT) + assert _mantle_bodies(mantle) == (_MANTLE_BARE,) + + +@pytest.mark.parametrize("fields", [{}, {"messages": []}], ids=["missing", "empty"]) +def test_messages_count_tokens_without_messages_is_rejected( + counting_proxy: OwnedProxy, mantle_port: int, fields: dict[str, JsonValue] +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, {"model": model, "system": _SYSTEM, "tools": _TOOLS, **fields}) + assert response.status_code == 400, response.text + assert "messages parameter is required" in response.text, response.text + assert runtime.drain() == () and mantle.drain() == () + + +def test_messages_count_tokens_unauthenticated_request_never_reaches_either_peer( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request( + "POST", "/v1/messages/count_tokens", _full(model), key="sk-not-a-key-this-proxy-issued" + ) + assert response.status_code == 401, response.text + assert "input_tokens" not in response.text, response.text + assert runtime.drain() == () and mantle.drain() == () + + +def test_messages_count_tokens_forwards_a_5kb_system_verbatim_to_mantle( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + system: Final = "Answer in one sentence. " * 214 + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, {**_bare(model), "system": system}) + assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_mantle_body(system=system),) + + +def test_messages_count_tokens_duplicate_system_and_tools_keys_forward_one_value_each_to_mantle( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + fields: Final = f'"system": {json.dumps(_SYSTEM)}, "tools": {json.dumps(_TOOLS)}' + response: Final = gateway.client.post( + "/v1/messages/count_tokens", + content=f'{{"model": "{model}", "messages": {json.dumps(_MESSAGES)}, {fields}, {fields}}}', + headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"}, + ) + assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + (sent,) = mantle.drain() + assert _mantle_requests((sent,)) == (_MANTLE_FULL,) + assert (sent.body.count(b'"system"'), sent.body.count(b'"tools"')) == (1, 1), sent.body + + +def test_messages_count_tokens_leaves_empty_tools_and_system_out_of_the_mantle_body( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, {**_bare(model), "tools": [], "system": ""}) + assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_BARE,) + + +def test_messages_count_tokens_repeated_request_reaches_both_peers_each_time( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + answers: Final = tuple(_payload(_count(gateway, _bare(model))) for _ in range(2)) + assert answers == ({"input_tokens": _MANTLE_COUNT},) * 2 + assert _runtime_count_targets(runtime) == (f"/model/{_OPUS_BASE}/count-tokens",) * 2 + assert _mantle_bodies(mantle) == (_MANTLE_BARE,) * 2 + + +def test_disabled_token_counter_surfaces_the_mantle_error_instead_of_counting_locally( + gateway: Gateway, mantle_port: int, tmp_path: Path +) -> None: + with ExitStack() as stack: + runtime: Final = stack.enter_context(wire_server(_runtime)) + config: Final = _owned_config( + tmp_path / "disabled-token-counter.yaml", runtime.url, {"disable_token_counter": True} + ) + owned: Final = stack.enter_context( + owned_proxy_process( + gateway, + tmp_path, + _mantle_environment(mantle_port), + config=config, + workers=2, + remove_environment=_INHERITED_BEARER, + ) + ) + with wire_server(_mantle(_rejecting(403)), port=mantle_port) as refusing: + refused: Final = _count(owned.gateway, _full(_OWNED_OPUS)) + assert refused.status_code == 403, refused.text + assert _REJECTION in refused.text and "input_tokens" not in refused.text, refused.text + assert _mantle_bodies(refusing) == (_MANTLE_FULL,) + with wire_server(_mantle(), port=mantle_port) as counting: + assert _payload(_count(owned.gateway, _full(_OWNED_OPUS))) == {"input_tokens": _MANTLE_COUNT} + assert _mantle_bodies(counting) == (_MANTLE_FULL,) + assert _payload(_count(owned.gateway, _full(_OWNED_SONNET))) == {"input_tokens": _RUNTIME_COUNT} + rejected: Final = _count(owned.gateway, _full(_OWNED_NOVA)) + assert rejected.status_code == 400, rejected.text + assert _UNSUPPORTED in rejected.text and "input_tokens" not in rejected.text, rejected.text + assert counting.drain() == () + assert len(_runtime_count_targets(runtime)) == 4 + + +def test_mantle_outage_between_concurrent_waves_falls_back_then_recovers( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ExitStack() as stack: + clients: Final = _clients(stack, str(gateway.client.base_url), 8) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients))) + runtime: Final = stack.enter_context(wire_server(_runtime)) + scenario: Final = stack.enter_context(gateway.scenario()) + model: Final = _deployment(scenario, runtime.url) + body: Final = _full(model) + local: Final = _local_count(gateway, body) + + def generate_then_count(client: httpx.Client) -> tuple[int, int, JsonValue]: + return _generated_then_counted(client, gateway.key, model, body) + + def count_only(client: httpx.Client) -> tuple[int, JsonValue]: + return _counted_on(client, gateway.key, body) + + with wire_server(_mantle(), port=mantle_port) as mantle: + assert tuple(pool.map(generate_then_count, clients)) == ((200, 200, _MANTLE_COUNT),) * len(clients) + assert _mantle_bodies(mantle) == (_MANTLE_FULL,) * len(clients) + outage: Final = tuple(pool.map(count_only, clients)) + assert outage == ((200, local),) * len(clients) + with wire_server(_mantle(), port=mantle_port) as revived: + assert tuple(pool.map(generate_then_count, clients)) == ((200, 200, _MANTLE_COUNT),) * len(clients) + assert _mantle_bodies(revived) == (_MANTLE_FULL,) * len(clients) + assert len(_runtime_count_targets(runtime)) == 3 * len(clients) + + +def test_slow_mantle_holds_concurrent_counts_without_stalling_the_proxy( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + held: Final[SimpleQueue[str]] = SimpleQueue() + release: Final = threading.Event() + with ExitStack() as stack: + clients: Final = _clients(stack, str(gateway.client.base_url), 6) + runtime: Final = stack.enter_context(wire_server(_runtime)) + mantle: Final = stack.enter_context(wire_server(_mantle(_holding(held, release, 20)), port=mantle_port)) + scenario: Final = stack.enter_context(gateway.scenario()) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients))) + stack.callback(release.set) + model: Final = _deployment(scenario, runtime.url) + futures: Final = tuple( + pool.submit(_generated_then_counted, client, gateway.key, model, _full(model)) for client in clients + ) + eventually(held.qsize, lambda size: size == len(clients), seconds=30) + assert gateway.request("GET", "/health/liveliness").status_code == 200 + assert _local_count(gateway, _full(model)) > 0 + assert not any(future.done() for future in futures) + release.set() + assert tuple(future.result(timeout=30) for future in futures) == ((200, 200, _MANTLE_COUNT),) * len(clients) + assert _mantle_bodies(mantle) == (_MANTLE_FULL,) * len(clients) + assert len(_runtime_count_targets(runtime)) == len(clients) + + +def test_worker_sigkill_mid_burst_leaves_the_sibling_counting( + gateway: Gateway, mantle_port: int, tmp_path: Path +) -> None: + held: Final[SimpleQueue[str]] = SimpleQueue() + release: Final = threading.Event() + with ExitStack() as stack: + runtime: Final = stack.enter_context(wire_server(_runtime)) + mantle: Final = stack.enter_context(wire_server(_mantle(_holding(held, release, 60)), port=mantle_port)) + config: Final = _owned_config(tmp_path / "worker-kill.yaml", runtime.url, {}) + owned: Final = stack.enter_context( + owned_proxy_process( + gateway, + tmp_path, + _mantle_environment(mantle_port), + config=config, + workers=2, + remove_environment=_INHERITED_BEARER, + ) + ) + body: Final = _full(_OWNED_OPUS) + proxy_url: Final = owned.gateway.client.base_url + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + clients: Final = _clients(stack, str(proxy_url), 12) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients))) + stack.callback(release.set) + probed: Final[SimpleQueue[int]] = SimpleQueue() + futures: Final[tuple[Future[tuple[int, tuple[int, JsonValue] | None]], ...]] = tuple( + pool.submit(_probed_then_counted_or_dropped, client, owned.gateway.key, body, probed) for client in clients + ) + eventually(held.qsize, lambda size: size == len(clients), seconds=30) + ports: Final = frozenset(probed.get_nowait() for _ in clients) + shares: Final = {pid: _accepted_client_ports(pid, proxy_url.port or 0) & ports for pid in workers} + assert sum(map(len, shares.values())) == len(clients), shares + victim: Final = min((pid for pid in workers if shares[pid]), key=lambda pid: len(shares[pid])) + psutil.Process(victim).send_signal(signal.SIGKILL) + release.set() + for port, result in (future.result(timeout=60) for future in futures): + assert result == (None if port in shares[victim] else (200, _MANTLE_COUNT)), (port, result, shares) + second_wave: Final = _clients(stack, str(proxy_url), 6) + assert tuple(_counted_on(client, owned.gateway.key, body) for client in second_wave) == ( + (200, _MANTLE_COUNT), + ) * len(second_wave) + assert len(_mantle_bodies(mantle)) == len(clients) + len(second_wave) + eventually(lambda: len(_STARTED_WORKER.findall(owned.log.read_text())), lambda started: started >= 3, 120) + assert owned.process.poll() is None + + +def _streamed_text(text: str) -> str: + events: Final = tuple( + _JSON_OBJECT.validate_json(line.removeprefix("data: ")) + for line in text.splitlines() + if line.startswith("data: {") + ) + return "".join(_delta_content(event) for event in events) + + +def _delta_content(event: dict[str, JsonValue]) -> str: + choices: Final = event.get("choices") + if not isinstance(choices, list) or not choices: + return "" + delta: Final = object_value(choices[0]).get("delta") + if not isinstance(delta, dict): + return "" + content: Final = delta.get("content") + return content if isinstance(content, str) else "" + + +@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"]) +def test_chat_completions_on_the_same_deployment_still_generate( + counting_proxy: OwnedProxy, mantle_port: int, stream: bool +) -> None: + gateway: Final = counting_proxy.gateway + marker: Final = uuid.uuid4().hex + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": f"chat control marker-{marker}"}], + "stream": stream, + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == 200, response.text + generated: Final = _streamed_text(response.text) if stream else response.text + assert answer(marker) in generated, response.text + assert not stream or response.text.rstrip().endswith("data: [DONE]"), response.text + (sent,) = runtime.drain() + assert target_of(sent) == f"/model/{_OPUS}/{'converse-stream' if stream else 'converse'}", sent.target + assert mantle.drain() == () + + +def test_messages_endpoint_on_the_same_deployment_still_generates(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + marker: Final = uuid.uuid4().hex + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": f"marker-{marker}"}]}, + ) + assert response.status_code == 200, response.text + assert answer(marker) in response.text, response.text + (sent,) = runtime.drain() + assert target_of(sent) == f"/model/{_OPUS}/invoke", sent.target + assert mantle.drain() == () + + +def test_responses_endpoint_on_the_same_deployment_still_generates( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + marker: Final = uuid.uuid4().hex + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request("POST", "/v1/responses", {"model": model, "input": f"marker-{marker}"}) + assert response.status_code == 200, response.text + assert answer(marker) in response.text, response.text + (sent,) = runtime.drain() + assert target_of(sent) == f"/model/{_OPUS}/converse", sent.target + assert mantle.drain() == () + + +def test_utils_token_counter_local_mode_never_calls_either_peer(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + assert _local_count(gateway, _full(model)) > 0 + assert runtime.drain() == () and mantle.drain() == () diff --git a/tests/integration/providers/test_bedrock_mantle_gpt_prompt_cache_breakpoint_wire.py b/tests/integration/providers/test_bedrock_mantle_gpt_prompt_cache_breakpoint_wire.py new file mode 100644 index 00000000000..2af41c2b683 --- /dev/null +++ b/tests/integration/providers/test_bedrock_mantle_gpt_prompt_cache_breakpoint_wire.py @@ -0,0 +1,531 @@ +import asyncio +import json +from typing import Final +from urllib.parse import urlsplit + +import pytest +from integration._support.client import Gateway, eventually +from integration._support.wire import Request, wire_server +from integration.providers._mantle_gpt_prompt_cache_support import ( + ANTHROPIC_PATH, + AZURE, + BURST, + CLAUDE, + ENDPOINTS, + EPHEMERAL, + EXPLICIT, + GPT, + GPT_FLAGGED_ROW, + GPT_BARE, + GPT_REGION, + GPT_UNFLAGGED_ROW, + HOSTILE_OPTIONS, + IMPLICIT, + ITEMS, + MALFORMED_POINTS, + MAX_TOKENS, + MIXED_POINTS, + NO_CACHE, + ODD_FLAGS, + SYSTEM, + SYSTEM_POINT, + THIRD_PARTY, + Endpoint, + Outcome, + assert_answered, + assert_burst_landed, + assert_marked_sdk_cell, + assert_priced_row, + assert_wire, + async_sdk_call, + body_of, + breakpoint_count, + burst, + deployment, + expected_wire, + failing_peer, + failure_row, + fresh_marker, + mantle_deployment, + mantle_peer, + marked_cell, + observe, + only_received, + openai_shaped_peer, + prompt_text, + request_body, + row_key, + sdk_call, + send, + send_raw, + settled, + spend_rows, + success_row, + system_item, +) +from pydantic import JsonValue + +pytestmark = pytest.mark.timeout(600) + + +@pytest.mark.parametrize("stream", [False, True], ids=["json", "stream"]) +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_a_configured_system_point_reaches_mantle_as_a_breakpoint_over_httpx( + gateway: Gateway, endpoint: Endpoint, stream: bool +) -> None: + marked_cell(gateway, endpoint, stream=stream) + + +@pytest.mark.parametrize("stream", [False, True], ids=["json", "stream"]) +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_a_the_sync_sdks_see_the_breakpoint_and_the_mapped_cache_usage( + gateway: Gateway, endpoint: Endpoint, stream: bool +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + outcome: Final = sdk_call(gateway, endpoint, name, prompt_text(marker), stream=stream) + assert_marked_sdk_cell(wire, endpoint, name, marker, outcome, stream=stream) + + +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_a_the_async_sdks_see_the_breakpoint_and_the_mapped_cache_usage(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + outcome: Final = asyncio.run(async_sdk_call(gateway, endpoint, name, prompt_text(marker))) + assert_marked_sdk_cell(wire, endpoint, name, marker, outcome, stream=False) + + +@pytest.mark.parametrize("endpoint", ["responses", "chat"]) +def test_b1_b2_a_pinned_explicit_mode_reaches_the_wire_with_the_breakpoint( + gateway: Gateway, endpoint: Endpoint +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, prompt_cache_options=EXPLICIT) + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_wire( + received, + expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True, options=EXPLICIT), + streaming=False, + ) + assert_priced_row(success_row(name, marker), GPT) + + +@pytest.mark.parametrize("endpoint", ["responses", "chat"]) +def test_b3_b4_a_region_prefixed_deployment_reads_the_region_free_row(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, GPT_REGION) + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_wire(received, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=False) + assert_priced_row(success_row(name, marker), GPT) + + +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_b15_to_b17_a_bare_deployment_name_with_its_provider_reads_the_provider_keyed_row( + gateway: Gateway, endpoint: Endpoint +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, GPT_BARE, custom_llm_provider="bedrock_mantle") + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_wire(received, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=False) + assert_priced_row(success_row(name, marker), GPT) + + +def assert_anthropic_marked_wire(received: Request, marker: str) -> None: + body: Final = body_of(received) + assert urlsplit(received.target).path == ANTHROPIC_PATH, received.target + assert body["system"] == [{"type": "text", "text": SYSTEM, "cache_control": EPHEMERAL}], received.body + assert body["messages"] == [{"role": "user", "content": [{"type": "text", "text": prompt_text(marker)}]}], ( + received.body + ) + assert "prompt_cache_options" not in body, received.body + assert "prompt_cache_breakpoint" not in received.body.decode(), received.body + + +@pytest.mark.parametrize("endpoint", ["messages", "chat"]) +def test_b5_b6_claude_on_mantle_keeps_the_anthropic_dialect(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, CLAUDE) + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_anthropic_marked_wire(received, marker) + success_row(name, marker, outcome.response_id) + + +def test_b7_a_deployment_flag_true_opts_an_unflagged_mantle_row_in(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment( + gateway, scenario, wire, GPT_UNFLAGGED_ROW, model_info={"supports_prompt_cache_breakpoint": True} + ) + outcome, received = observe(gateway, wire, "responses", request_body("responses", name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_wire( + received, + expected_wire(GPT_UNFLAGGED_ROW, prompt_text(marker), endpoint="responses", marked=True), + streaming=False, + ) + success_row(name, marker) + + +def test_b8_a_deployment_flag_false_opts_a_flagged_mantle_row_out(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment( + gateway, scenario, wire, GPT_FLAGGED_ROW, model_info={"supports_prompt_cache_breakpoint": False} + ) + outcome, received = observe(gateway, wire, "responses", request_body("responses", name, prompt_text(marker))) + assert_answered(outcome, marker) + body: Final = body_of(received) + assert breakpoint_count(body) == 0, received.body + assert "prompt_cache_options" not in body, received.body + assert "cache_control" not in received.body.decode(), received.body + success_row(name, marker) + + +@pytest.mark.parametrize("endpoint", ["chat", "responses"]) +def test_b13_azure_openai_stays_ineligible_for_the_breakpoint_dialect(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(openai_shaped_peer()) as wire, gateway.scenario() as scenario: + name: Final = scenario.model( + model=AZURE, + api_base=wire.url, + api_key="synthetic-azure-key", + api_version="2025-04-01-preview", + cache_control_injection_points=SYSTEM_POINT, + ) + settled(gateway, name, wire) + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert urlsplit(received.target).path.startswith("/openai/"), received.target + assert "prompt_cache_breakpoint" not in received.body.decode(), received.body + assert "prompt_cache_options" not in body_of(received), received.body + assert SYSTEM in received.body.decode(), received.body + success_row(name, marker) + + +def test_b14_an_openai_entry_on_a_third_party_host_keeps_the_anthropic_dialect(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(openai_shaped_peer()) as wire, gateway.scenario() as scenario: + name: Final = scenario.model( + model=THIRD_PARTY, + api_base=wire.url, + api_key="synthetic-third-party-key", + cache_control_injection_points=SYSTEM_POINT, + ) + settled(gateway, name, wire) + outcome, received = observe(gateway, wire, "chat", request_body("chat", name, prompt_text(marker))) + assert_answered(outcome, marker) + body: Final = body_of(received) + assert urlsplit(received.target).path == "/chat/completions", received.target + assert body["messages"] == [ + {"role": "system", "content": SYSTEM, "cache_control": EPHEMERAL}, + {"role": "user", "content": prompt_text(marker)}, + ], received.body + assert "prompt_cache_options" not in body, received.body + success_row(name, marker) + + +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_c_a_response_cache_hit_serves_the_marked_request_again(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + body: Final = request_body(endpoint, name, prompt_text(marker)) + first, received = observe(gateway, wire, endpoint, body) + assert_answered(first, marker) + assert_wire(received, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=False) + sends: Final[list[Outcome]] = [] + + def resend() -> Outcome: + served: Final = send(gateway, endpoint, body) + sends.append(served) + return served + + hit: Final = eventually( + resend, lambda served: served.status == 200 and served.response_id == first.response_id, seconds=30 + ) + assert_answered(hit, marker) + misses: Final = wire.drain() + assert len(misses) == len(sends) - 1, (len(misses), len(sends)) + for miss in misses: + assert_wire(miss, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=False) + rows: Final = spend_rows(name, frozenset({marker}), expected=len(sends) + 1, seconds=70) + assert len(rows) == len(sends) + 1, rows + (hit_row,) = tuple(row for row in rows if row["cache_hit"] == "True") + assert row_key(str(hit_row["request_id"])) == first.response_id, (hit_row, first.response_id) + + +@pytest.mark.parametrize("endpoint", ["chat", "responses"]) +@pytest.mark.parametrize("shape", sorted(HOSTILE_OPTIONS)) +def test_d1_to_d4_hostile_prompt_cache_options_reach_the_wire_verbatim_and_the_providers_400_reaches_the_caller( + gateway: Gateway, endpoint: Endpoint, shape: str +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + hostile: Final = HOSTILE_OPTIONS[shape] + outcome, received = observe( + gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker), prompt_cache_options=hostile) + ) + assert outcome.status == 400, (outcome.status, outcome.raw) + assert "Invalid prompt_cache_options" in outcome.raw, outcome.raw + assert_wire( + received, + expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True, options=hostile), + streaming=False, + ) + failure_row(outcome.call_id) + + +@pytest.mark.parametrize("endpoint", ["chat", "responses"]) +def test_d5_a_duplicated_prompt_cache_options_key_resolves_to_the_last_value( + gateway: Gateway, endpoint: Endpoint +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + encoded: Final = json.dumps(request_body(endpoint, name, prompt_text(marker), prompt_cache_options=IMPLICIT)) + duplicated: Final = encoded[:-1] + ', "prompt_cache_options": {"mode": "explicit"}}' + wire.drain() + outcome: Final = send_raw(gateway, endpoint, duplicated) + assert_answered(outcome, marker) + assert_wire( + only_received(wire), + expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True, options=EXPLICIT), + streaming=False, + ) + success_row(name, marker) + + +@pytest.mark.parametrize("endpoint", ["chat", "responses"]) +def test_d6_a_null_prompt_cache_options_is_treated_as_unset(gateway: Gateway, endpoint: Endpoint) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + outcome, received = observe( + gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker), prompt_cache_options=None) + ) + assert_answered(outcome, marker) + assert_wire(received, expected_wire(GPT, prompt_text(marker), endpoint=endpoint, marked=True), streaming=False) + success_row(name, marker) + + +def test_d7_an_unauthenticated_hostile_request_never_reaches_the_wire(gateway: Gateway) -> None: + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + wire.drain() + outcome: Final = send( + gateway, + "responses", + request_body("responses", name, prompt_text(fresh_marker()), prompt_cache_options=5), + key="sk-not-a-key", + ) + assert outcome.status == 401, (outcome.status, outcome.raw) + assert wire.drain() == (), "the upstream saw an unauthenticated request" + + +@pytest.mark.parametrize("shape", sorted(MALFORMED_POINTS)) +def test_d8_to_d10_malformed_injection_points_never_crash_the_request(gateway: Gateway, shape: str) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = deployment(scenario, wire, GPT, points=MALFORMED_POINTS[shape]) + settled(gateway, name, wire) + outcome, received = observe(gateway, wire, "responses", request_body("responses", name, prompt_text(marker))) + assert_answered(outcome, marker) + body: Final = body_of(received) + assert breakpoint_count(body) == 0, received.body + assert "prompt_cache_options" not in body, received.body + success_row(name, marker) + + +@pytest.mark.parametrize( + ("status", "endpoint"), + [(400, "chat"), (400, "responses"), (401, "responses")], + ids=["400-chat", "400-responses", "401-responses"], +) +def test_d11_d12_a_provider_error_on_the_marked_request_reaches_the_caller_after_one_attempt( + gateway: Gateway, status: int, endpoint: Endpoint +) -> None: + marker: Final = fresh_marker() + message: Final = f"scripted provider failure {marker}" + with wire_server(failing_peer(status, message, "scripted_failure")) as wire, gateway.scenario() as scenario: + name: Final = deployment(scenario, wire, GPT) + settled(gateway, name, wire, accepted=frozenset({status})) + outcome: Final = send(gateway, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert outcome.status == status, (outcome.status, outcome.raw) + assert message in outcome.raw, outcome.raw + attempts: Final = wire.drain() + assert len(attempts) == 1, [attempt.target for attempt in attempts] + assert breakpoint_count(body_of(attempts[0])) == 1, attempts[0].body + failure_row(outcome.call_id) + + +@pytest.mark.parametrize(("label", "model", "flag"), ODD_FLAGS, ids=[label for label, _, _ in ODD_FLAGS]) +def test_d13_an_odd_typed_deployment_flag_never_opts_a_mantle_row_in( + gateway: Gateway, label: str, model: str, flag: JsonValue +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment( + gateway, scenario, wire, model, model_info={"supports_prompt_cache_breakpoint": flag} + ) + outcome, received = observe(gateway, wire, "responses", request_body("responses", name, prompt_text(marker))) + assert_answered(outcome, marker) + body: Final = body_of(received) + assert breakpoint_count(body) == 0, (label, received.body) + assert "prompt_cache_options" not in body, (label, received.body) + success_row(name, marker) + + +@pytest.mark.parametrize("endpoint", ENDPOINTS) +def test_d14_a_point_beside_junk_entries_still_marks_the_anthropic_dialect( + gateway: Gateway, endpoint: Endpoint +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, CLAUDE, points=MIXED_POINTS) + outcome, received = observe(gateway, wire, endpoint, request_body(endpoint, name, prompt_text(marker))) + assert_answered(outcome, marker) + assert_anthropic_marked_wire(received, marker) + success_row(name, marker, outcome.response_id) + + +def test_e1_client_breakpoint_prompt_cache_key_and_explicit_mode_pass_through_with_no_second_breakpoint( + gateway: Gateway, +) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + body: Final[dict[str, JsonValue]] = { + "model": name, + "input": [system_item(marked=True, endpoint="responses"), {"role": "user", "content": prompt_text(marker)}], + "max_output_tokens": MAX_TOKENS, + "prompt_cache_key": f"key-{marker}", + "prompt_cache_options": EXPLICIT, + } + outcome, received = observe(gateway, wire, "responses", body) + assert_answered(outcome, marker) + assert_wire( + received, + expected_wire( + GPT, + prompt_text(marker), + endpoint="responses", + marked=True, + options=EXPLICIT, + prompt_cache_key=f"key-{marker}", + ), + streaming=False, + ) + success_row(name, marker) + + +def test_e2_four_client_breakpoints_leave_no_slot_for_the_configured_point(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + turns: Final = tuple(f"Earlier turn {index} marker-{marker}." for index in range(3)) + body: Final[dict[str, JsonValue]] = { + "model": name, + "input": [ + {"role": "system", "content": SYSTEM}, + *( + { + "role": "user", + "content": [{"type": "input_text", "text": turn, "prompt_cache_breakpoint": EXPLICIT}], + } + for turn in turns + ), + { + "role": "user", + "content": [ + {"type": "input_text", "text": prompt_text(marker), "prompt_cache_breakpoint": EXPLICIT} + ], + }, + ], + "max_output_tokens": MAX_TOKENS, + } + outcome, received = observe(gateway, wire, "responses", body) + assert_answered(outcome, marker) + wire_body: Final = body_of(received) + items: Final = ITEMS.validate_python(wire_body["input"]) + assert breakpoint_count(wire_body) == 4, received.body + assert items[0]["role"] == "system" and "prompt_cache_breakpoint" not in json.dumps(items[0]), items[0] + assert "prompt_cache_options" not in wire_body, received.body + success_row(name, marker) + + +def test_e3_a_client_implicit_mode_on_a_pinned_explicit_deployment_wins_on_the_wire(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire, prompt_cache_options=EXPLICIT) + outcome, received = observe( + gateway, + wire, + "responses", + request_body("responses", name, prompt_text(marker), prompt_cache_options=IMPLICIT), + ) + assert_answered(outcome, marker) + assert_wire( + received, + expected_wire(GPT, prompt_text(marker), endpoint="responses", marked=True, options=IMPLICIT), + streaming=False, + ) + success_row(name, marker) + + +def test_e4_an_empty_prompt_cache_options_object_is_kept_as_the_clients_value(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + outcome, received = observe( + gateway, wire, "responses", request_body("responses", name, prompt_text(marker), prompt_cache_options={}) + ) + assert_answered(outcome, marker) + assert_wire( + received, + expected_wire(GPT, prompt_text(marker), endpoint="responses", marked=True, options={}), + streaming=False, + ) + success_row(name, marker) + + +def test_e5_the_same_request_twice_writes_two_rows_and_two_marked_upstream_requests(gateway: Gateway) -> None: + marker: Final = fresh_marker() + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + body: Final = request_body("chat", name, prompt_text(marker), cache=NO_CACHE) + first, first_received = observe(gateway, wire, "chat", body) + second, second_received = observe(gateway, wire, "chat", body) + assert_answered(first, marker) + assert_answered(second, marker) + assert first.response_id != second.response_id, (first.response_id, second.response_id) + for received in (first_received, second_received): + assert_wire( + received, expected_wire(GPT, prompt_text(marker), endpoint="chat", marked=True), streaming=False + ) + rows: Final = spend_rows(name, frozenset({marker}), expected=2) + assert {row_key(str(row["request_id"])) for row in rows} == {first.response_id, second.response_id}, rows + + +def test_f1_a_mixed_burst_marks_every_request_and_lands_every_id_once(gateway: Gateway) -> None: + with wire_server(mantle_peer()) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + wire.drain() + assert_burst_landed(wire, name, burst(gateway, name, BURST), marked=True) + + +def test_f4_a_slow_mantle_stream_during_a_burst_completes_with_every_id_landing_once(gateway: Gateway) -> None: + with wire_server(mantle_peer(pause=0.3)) as wire, gateway.scenario() as scenario: + name: Final = mantle_deployment(gateway, scenario, wire) + wire.drain() + assert_burst_landed(wire, name, burst(gateway, name, 12), marked=True) diff --git a/tests/integration/providers/test_cache_control_tool_call_marks_wire.py b/tests/integration/providers/test_cache_control_tool_call_marks_wire.py index de5d4170fc8..21fee0554ea 100644 --- a/tests/integration/providers/test_cache_control_tool_call_marks_wire.py +++ b/tests/integration/providers/test_cache_control_tool_call_marks_wire.py @@ -1,16 +1,26 @@ import asyncio import json -from collections.abc import Iterable +from collections.abc import Iterable, Iterator from datetime import datetime, timedelta, timezone +from pathlib import Path from typing import Final, cast import anthropic import httpx import openai import pytest +import yaml from anthropic.types import MessageParam, TextBlockParam, ToolParam -from integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from integration._support.client import ( + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, + string_value, +) from integration._support.database import read_rows +from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, Wire, wire_server from integration.providers._cache_control_marks_support import ( ANTHROPIC_MODEL, @@ -52,6 +62,7 @@ from pydantic import JsonValue, TypeAdapter from litellm.utils import get_prompt_cache_min_tokens _JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_LANE_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml" _CLIENT_MARKED: Final = client_marked() _GEMINI_REPLY: Final = json.dumps( { @@ -71,6 +82,25 @@ _OPENAI_REPLY: Final = json.dumps( ).encode() +def _uncapped_cache_config(directory: Path) -> Path: + lane: Final = object_value(yaml.safe_load(_LANE_CONFIG.read_text())) + settings: Final = object_value(lane["litellm_settings"]) + cache_params: Final = object_value(settings["cache_params"]) + config: Final = {**lane, "litellm_settings": {**settings, "cache_params": {**cache_params, "max_messages": None}}} + path: Final = directory / "proxy_config_uncapped_cache.yaml" + path.write_text(yaml.safe_dump(config, sort_keys=False)) + return path + + +@pytest.fixture(scope="module") +def uncapped_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + """A proxy whose response cache has no `max_messages` cap, so a seven-message tool conversation is cached""" + with gateway_from_environment() as lane: + directory: Final = tmp_path_factory.mktemp("uncapped_cache") + with owned_proxy(lane, directory, {}, config=_uncapped_cache_config(directory)) as owned: + yield owned + + def _anthropic_deployment(scenario: Scenario, wire: Wire, **fields: JsonValue) -> str: return scenario.model(model=f"anthropic/{ANTHROPIC_MODEL}", api_base=wire.url, api_key=PROVIDER_KEY, **fields) @@ -485,13 +515,14 @@ def test_responses_bridge_keeps_system_and_user_marks_within_the_cap(gateway: Ga assert anthropic_labels(_only_request(wire)) == [SYSTEM_LABEL, ASK_LABEL] -def test_response_cache_serves_the_capped_request_once(gateway: Gateway) -> None: +@pytest.mark.timeout(240) +def test_response_cache_serves_the_capped_request_once(uncapped_gateway: Gateway) -> None: marker: Final = new_marker() - with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + with wire_server(anthropic_peer) as wire, uncapped_gateway.scenario() as scenario: model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) body: Final = chat_body(model, conversation(marker, marked_calls())) - first: Final = post_chat(gateway, body) - second: Final = post_chat(gateway, body) + first: Final = post_chat(uncapped_gateway, body) + second: Final = post_chat(uncapped_gateway, body) received: Final = wire.drain() assert (first[0], second[0]) == (200, 200), (first[2], second[2]) assert first[1].startswith("chatcmpl-"), first[2] @@ -828,14 +859,15 @@ def _cache_hit_rows(response_id: str) -> list[dict[str, JsonValue]]: ) -def test_response_cache_hit_records_the_injection_on_the_first_row_only(gateway: Gateway) -> None: +@pytest.mark.timeout(240) +def test_response_cache_hit_records_the_injection_on_the_first_row_only(uncapped_gateway: Gateway) -> None: marker: Final = new_marker() unmarked: Final = [tool_call(city) for city in CITIES] - with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + with wire_server(anthropic_peer) as wire, uncapped_gateway.scenario() as scenario: model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) body: Final = chat_body(model, conversation(marker, unmarked, ask_marked=False)) - first: Final = post_chat(gateway, body) - second: Final = post_chat(gateway, body) + first: Final = post_chat(uncapped_gateway, body) + second: Final = post_chat(uncapped_gateway, body) received: Final = wire.drain() assert (first[0], second[0]) == (200, 200), (first[2], second[2]) assert first[1] == second[1], (first[2], second[2]) @@ -845,15 +877,18 @@ def test_response_cache_hit_records_the_injection_on_the_first_row_only(gateway: assert hit_rows[0]["injected"] is None, hit_rows +@pytest.mark.timeout(240) @pytest.mark.parametrize("path", ("/v1/messages", "/v1/responses")) -def test_response_cache_twins_on_messages_and_responses_stay_within_the_cap(gateway: Gateway, path: str) -> None: +def test_response_cache_twins_on_messages_and_responses_stay_within_the_cap( + uncapped_gateway: Gateway, path: str +) -> None: marker: Final = new_marker() - with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + with wire_server(anthropic_peer) as wire, uncapped_gateway.scenario() as scenario: model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) body: Final = ( messages_body(model, marker, stream=False) if path == "/v1/messages" else responses_body(model, marker) ) - responses: Final = tuple(gateway.request("POST", path, body) for _ in range(2)) + responses: Final = tuple(uncapped_gateway.request("POST", path, body) for _ in range(2)) received: Final = wire.drain() assert [response.status_code for response in responses] == [200, 200], [response.text for response in responses] expected: Final = _CLIENT_MARKED if path == "/v1/messages" else [SYSTEM_LABEL, ASK_LABEL] diff --git a/tests/integration/providers/test_decisions_wire.py b/tests/integration/providers/test_decisions_wire.py index 84b31088d60..db6629a43d2 100644 --- a/tests/integration/providers/test_decisions_wire.py +++ b/tests/integration/providers/test_decisions_wire.py @@ -107,7 +107,7 @@ _PROVIDERS: Final = ( _PERPLEXITY: Final = _PROVIDERS[0] _OPENROUTER: Final = _PROVIDERS[2] _OPENROUTER_CHAT_MODEL: Final = "openrouter/openai/gpt-5-mini" -_UNSUPPORTED_PROVIDER_MODEL: Final = "openai/gpt-6-luna" +_UNSUPPORTED_PROVIDER_MODEL: Final = "anthropic/claude-opus-5-5" _CONNECTION_ERROR: Final = "litellm.APIConnectionError" _GENERIC_API_ERROR: Final = "litellm.APIError" _INVALID_BODIES: Final[tuple[tuple[str, dict[str, JsonValue]], ...]] = ( diff --git a/tests/integration/providers/test_gemini_chat_video_metadata_wire.py b/tests/integration/providers/test_gemini_chat_video_metadata_wire.py new file mode 100644 index 00000000000..32914341483 --- /dev/null +++ b/tests/integration/providers/test_gemini_chat_video_metadata_wire.py @@ -0,0 +1,144 @@ +from __future__ import annotations + +import base64 +import json +from collections.abc import Callable +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, Scenario, object_value +from integration._support.vertex import service_account_json +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +BACKEND: Final = "gemini-3.7-flash" +PROJECT: Final = "scripted-project" +LOCATION: Final = "us-central1" +MODEL_PATH: Final = f"/v1/projects/{PROJECT}/locations/{LOCATION}/publishers/google/models/{BACKEND}" +CLIP: Final = base64.b64encode(b"\x00\x00\x00\x18ftypmp42" + bytes(24)).decode() +DATA_URI: Final = f"data:video/mp4;base64,{CLIP}" +PROMPT: Final = "describe the clip" +FILE_BLOCK: Final[dict[str, JsonValue]] = { + "type": "file", + "file": {"file_data": DATA_URI, "video_metadata": {"fps": 1.0, "start_offset": "0s", "end_offset": "3s"}}, +} +TEXT_PART: Final[dict[str, JsonValue]] = {"text": PROMPT} +BARE_VIDEO_PART: Final[dict[str, JsonValue]] = {"inline_data": {"mime_type": "video/mp4", "data": CLIP}} +VIDEO_PART: Final[dict[str, JsonValue]] = { + **BARE_VIDEO_PART, + "video_metadata": {"fps": 1.0, "startOffset": "0s", "endOffset": "3s"}, +} +REPLY: Final[dict[str, JsonValue]] = { + "candidates": [{"content": {"role": "model", "parts": [{"text": "a cat"}]}, "finishReason": "STOP"}], + "usageMetadata": {"promptTokenCount": 5, "candidatesTokenCount": 2, "totalTokenCount": 7}, +} + + +def chat_peer(path: str, authorization: str | None) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert request.method == "POST", request.method + assert authorization is None or request.headers["authorization"] == authorization + target, _, query = request.target.partition("?") + if target == f"{path}:streamGenerateContent": + assert "alt=sse" in query, request.target + return Reply(content_type="text/event-stream", chunks=(f"data: {json.dumps(REPLY)}\n\n".encode(),)) + assert target == f"{path}:generateContent", request.target + return Reply(body=json.dumps(REPLY).encode()) + + return respond + + +def wire_parts(wire: Wire) -> list[JsonValue]: + received: Final = wire.drain() + assert len(received) == 1, [item.target for item in received] + body: Final = object_value(json.loads(received[0].body)) + contents: Final = body["contents"] + assert isinstance(contents, list) and len(contents) == 1, body + parts: Final = object_value(contents[0])["parts"] + assert isinstance(parts, list), body + return parts + + +def gemini_model(scenario: Scenario, url: str) -> str: + return scenario.model(model=f"gemini/{BACKEND}", api_key="scripted-gemini-key", api_base=url) + + +def vertex_model(gateway: Gateway, scenario: Scenario, url: str) -> str: + return scenario.model( + model=f"vertex_ai/{BACKEND}", + api_key=None, + api_base=url, + vertex_project=PROJECT, + vertex_location=LOCATION, + vertex_credentials=service_account_json(PROJECT, gateway.upstream_url), + ) + + +def chat(gateway: Gateway, model: str, stream: bool) -> httpx.Response: + body: Final[dict[str, JsonValue]] = { + "model": model, + "messages": [{"role": "user", "content": [{"type": "text", "text": PROMPT}, FILE_BLOCK]}], + "stream": stream, + } + return gateway.request("POST", "/v1/chat/completions", body) + + +def answered(response: httpx.Response) -> None: + assert response.status_code == 200, response.text + assert "a cat" in response.text, response.text + + +@pytest.mark.parametrize("stream", (False, True), ids=("non-stream", "stream")) +def test_gemini_chat_file_block_video_metadata_reaches_the_wire(gateway: Gateway, stream: bool) -> None: + with wire_server(chat_peer(f"/models/{BACKEND}", None)) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + answered(chat(gateway, model, stream)) + assert wire_parts(wire) == [TEXT_PART, VIDEO_PART] + + +@pytest.mark.parametrize("stream", (False, True), ids=("non-stream", "stream")) +def test_vertex_chat_file_block_video_metadata_reaches_the_wire(gateway: Gateway, stream: bool) -> None: + with wire_server(chat_peer(MODEL_PATH, "Bearer scripted-token")) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + answered(chat(gateway, model, stream)) + assert wire_parts(wire) == [TEXT_PART, VIDEO_PART] + + +def test_responses_input_file_video_reaches_the_wire_as_inline_data(gateway: Gateway) -> None: + with wire_server(chat_peer(f"/models/{BACKEND}", None)) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + body: Final[dict[str, JsonValue]] = { + "model": model, + "input": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": PROMPT}, + {"type": "input_file", "file_data": DATA_URI, "filename": "clip.mp4"}, + ], + } + ], + } + answered(gateway.request("POST", "/v1/responses", body)) + assert wire_parts(wire) == [TEXT_PART, BARE_VIDEO_PART] + + +def test_messages_document_video_reaches_the_wire_as_inline_data(gateway: Gateway) -> None: + with wire_server(chat_peer(f"/models/{BACKEND}", None)) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + body: Final[dict[str, JsonValue]] = { + "model": model, + "max_tokens": 32, + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": PROMPT}, + {"type": "document", "source": {"type": "base64", "media_type": "video/mp4", "data": CLIP}}, + ], + } + ], + } + answered(gateway.request("POST", "/v1/messages", body)) + assert wire_parts(wire) == [TEXT_PART, BARE_VIDEO_PART] diff --git a/tests/integration/providers/test_gemini_embedding_file_block_chaos.py b/tests/integration/providers/test_gemini_embedding_file_block_chaos.py new file mode 100644 index 00000000000..cee0758b433 --- /dev/null +++ b/tests/integration/providers/test_gemini_embedding_file_block_chaos.py @@ -0,0 +1,223 @@ +from __future__ import annotations + +import json +import os +import signal +import threading +import uuid +import zlib +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import psutil +import pytest +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.process import group_members, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +BACKEND: Final = "gemini-embedding-001" +TARGET: Final = f"/models/{BACKEND}:batchEmbedContents" +END_OFFSETS: Final = {"fast": "3s", "slow": "5s", "malformed": "3s", "dropped": "9s"} +PLAN: Final = tuple(enumerate(("fast",) * 16 + ("slow",) * 8 + ("malformed",) * 8 + ("dropped",) * 8)) +BURST: Final = 24 +SPEND_SQL: Final = 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = %s' + + +@dataclass(frozen=True, slots=True) +class Attempt: + kind: str + index: int + status: int + call_id: str + + +def source(kind: str, index: int) -> str: + return f"gs://scripted-bucket/chaos/{kind}-{index}.mp4" + + +def body(model: str, kind: str, index: int) -> dict[str, JsonValue]: + fps: Final[JsonValue] = "x" if kind == "malformed" else 1.0 + metadata: Final[dict[str, JsonValue]] = {"fps": fps, "start_offset": "0s", "end_offset": END_OFFSETS[kind]} + return { + "model": model, + "input": [{"type": "file", "file": {"file_id": source(kind, index), "video_metadata": metadata}}], + } + + +def first_part(request: Request) -> dict[str, JsonValue]: + requests: Final = object_value(json.loads(request.body))["requests"] + assert isinstance(requests, list) and len(requests) == 1, requests + parts: Final = object_value(object_value(requests[0])["content"])["parts"] + assert isinstance(parts, list) and len(parts) == 1, parts + return object_value(parts[0]) + + +def source_on_the_wire(request: Request) -> str: + return string_value(object_value(first_part(request)["file_data"])["file_uri"]) + + +def embed_reply(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target.split("?")[0] == TARGET, request.target + values: Final = [zlib.crc32(source_on_the_wire(request).encode()) / 2**32, 0.5] + return Reply(body=json.dumps({"embeddings": [{"values": values}]}).encode()) + + +def chaos_peer(healed: threading.Event) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + end_offset: Final = object_value(first_part(request)["video_metadata"])["endOffset"] + if end_offset == END_OFFSETS["dropped"] and not healed.is_set(): + return Reply(drop_connection=True) + reply: Final = embed_reply(request) + if end_offset == END_OFFSETS["slow"]: + return Reply(chunks=(reply.body[:8], reply.body[8:]), pause_between_chunks=1.5) + return reply + + return respond + + +def wire_sources(wire: Wire) -> tuple[str, ...]: + return tuple(source_on_the_wire(request) for request in wire.drain()) + + +def spend_row_count(call_id: str) -> int: + return len(read_rows(SPEND_SQL, (call_id,))) + + +def spend_row_counts(call_ids: tuple[str, ...]) -> tuple[int, ...]: + return tuple(spend_row_count(call_id) for call_id in call_ids) + + +def attempt(gateway: Gateway, model: str, index: int, kind: str) -> Attempt: + response: Final = gateway.request("POST", "/v1/embeddings", body(model, kind, index)) + return Attempt(kind, index, response.status_code, response.headers.get("x-litellm-call-id", "")) + + +def of_kind(attempts: tuple[Attempt, ...], *kinds: str) -> tuple[Attempt, ...]: + return tuple(item for item in attempts if item.kind in kinds) + + +@pytest.mark.timeout(240) +def test_mixed_burst_keeps_every_answer_and_spend_row_honest(gateway: Gateway) -> None: + healed: Final = threading.Event() + with wire_server(chaos_peer(healed)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"gemini/{BACKEND}", api_key="scripted-gemini-key", api_base=wire.url) + with ThreadPoolExecutor(max_workers=len(PLAN)) as pool: + attempts: Final = tuple(pool.map(lambda plan: attempt(gateway, model, plan[0], plan[1]), PLAN)) + served: Final = of_kind(attempts, "fast", "slow") + assert all(item.status == 200 for item in served), attempts + assert all(item.status == 400 for item in of_kind(attempts, "malformed")), attempts + assert all(item.status >= 500 for item in of_kind(attempts, "dropped")), attempts + reached: Final = wire_sources(wire) + assert sorted(item for item in reached if "/malformed-" not in item and "/dropped-" not in item) == sorted( + source(item.kind, item.index) for item in served + ), reached + assert not any("/malformed-" in item for item in reached), reached + assert {item for item in reached if "/dropped-" in item} == { + source(item.kind, item.index) for item in of_kind(attempts, "dropped") + }, reached + call_ids: Final = tuple(item.call_id for item in served) + assert all(call_ids), attempts + counts: Final = eventually( + lambda: spend_row_counts(call_ids), lambda found: all(count >= 1 for count in found), seconds=90 + ) + assert counts == (1,) * len(call_ids), counts + healed.set() + with ThreadPoolExecutor(max_workers=len(PLAN)) as pool: + resent: Final = tuple( + pool.map(lambda item: attempt(gateway, model, item.index, item.kind), of_kind(attempts, "dropped")) + ) + assert all(item.status == 200 for item in resent), resent + assert sorted(wire_sources(wire)) == sorted(source(item.kind, item.index) for item in resent) + + +def write_config(directory: Path, url: str, name: str) -> Path: + config: Final = directory / f"gemini_embeddings_{uuid.uuid4().hex}.yaml" + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": name, + "litellm_params": { + "model": f"gemini/{BACKEND}", + "api_key": "scripted-gemini-key", + "api_base": url, + }, + } + ], + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + }, + } + ) + ) + return config + + +def worker_pids(root_pid: int) -> frozenset[int]: + def is_worker(process: psutil.Process) -> bool: + try: + return "spawn_main" in " ".join(process.cmdline()) + except (psutil.NoSuchProcess, psutil.AccessDenied, psutil.ZombieProcess): + return False + + return frozenset(process.pid for process in group_members(root_pid) if is_worker(process)) + + +def outcome(candidate: Gateway, model: str, index: int) -> str: + try: + response: Final = candidate.request("POST", "/v1/embeddings", body(model, "fast", index)) + except httpx.TransportError as error: + return f"transport:{type(error).__name__}" + assert response.status_code == 200, response.text + return f"ok:{response.headers['x-litellm-call-id']}" + + +@pytest.mark.timeout(300) +def test_block_embeddings_survive_a_worker_kill(gateway: Gateway, tmp_path: Path) -> None: + name: Final = f"gemini-embeddings-{uuid.uuid4().hex}" + with wire_server(embed_reply) as wire: + config: Final = write_config(tmp_path, wire.url, name) + with owned_proxy_process(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + workers: Final = eventually(lambda: worker_pids(owned.process.pid), lambda pids: len(pids) == 2, seconds=30) + victim: Final = min(workers) + assert outcome(candidate, name, 0).startswith("ok:") + + def attempt_around_the_kill(index: int) -> str: + if index == 2: + os.kill(victim, signal.SIGKILL) + return outcome(candidate, name, index) + + with ThreadPoolExecutor(max_workers=BURST) as pool: + outcomes: Final = tuple(pool.map(attempt_around_the_kill, range(1, BURST + 1))) + assert outcomes.count("ok") == 0 and any(item.startswith("ok:") for item in outcomes), outcomes + assert all(item.startswith("ok:") or item.startswith("transport:") for item in outcomes), outcomes + respawned: Final = eventually( + lambda: worker_pids(owned.process.pid), + lambda pids: len(pids) == 2 and victim not in pids, + seconds=60, + ) + assert victim not in respawned, respawned + + def settled_burst() -> tuple[str, ...]: + with ThreadPoolExecutor(max_workers=BURST) as pool: + return tuple(pool.map(lambda index: outcome(candidate, name, BURST + 1 + index), range(BURST))) + + final: Final = eventually( + settled_burst, lambda values: all(v.startswith("ok:") for v in values), seconds=40 + ) + call_ids: Final = tuple(item.removeprefix("ok:") for item in (*outcomes, *final) if item.startswith("ok:")) + counts: Final = eventually( + lambda: spend_row_counts(call_ids), lambda found: all(count >= 1 for count in found), seconds=90 + ) + assert counts == (1,) * len(call_ids), counts + assert len(wire.drain()) >= len(call_ids) diff --git a/tests/integration/providers/test_gemini_embedding_file_block_wire.py b/tests/integration/providers/test_gemini_embedding_file_block_wire.py new file mode 100644 index 00000000000..d6439d27cd5 --- /dev/null +++ b/tests/integration/providers/test_gemini_embedding_file_block_wire.py @@ -0,0 +1,294 @@ +from __future__ import annotations + +import asyncio +import base64 +import json +import uuid +import zlib +from pathlib import Path +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, Scenario, object_value, string_value +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from openai import AsyncOpenAI, OpenAI +from openai.types import CreateEmbeddingResponse +from pydantic import JsonValue + +BACKEND: Final = "gemini-embedding-001" +TARGET: Final = f"/models/{BACKEND}:batchEmbedContents" +CLIP: Final = base64.b64encode(b"\x00\x00\x00\x18ftypmp42" + bytes(24)).decode() +DATA_URI: Final = f"data:video/mp4;base64,{CLIP}" +GCS_URI: Final = "gs://scripted-bucket/clips/animals.mp4" +METADATA: Final[dict[str, JsonValue]] = {"fps": 1.0, "start_offset": "0s", "end_offset": "3s"} +WIRE_METADATA: Final[dict[str, JsonValue]] = {"fps": 1.0, "startOffset": "0s", "endOffset": "3s"} +INLINE_DATA: Final[dict[str, JsonValue]] = {"mime_type": "video/mp4", "data": CLIP} +INLINE_PART: Final[dict[str, JsonValue]] = {"inline_data": INLINE_DATA, "video_metadata": WIRE_METADATA} +TEXT_PART: Final[dict[str, JsonValue]] = {"text": "a red bus"} +ACCEPTED_FORMS: Final = "must be a data: URI, a gs:// URL, a files/ reference, or a Gemini Files API URI" +LONG_OFFSET: Final = "9" * 5000 + "s" + + +def block(**file: JsonValue) -> dict[str, JsonValue]: + return {"type": "file", "file": file} + + +def gcs_part(uri: str = GCS_URI, metadata: dict[str, JsonValue] | None = WIRE_METADATA) -> dict[str, JsonValue]: + file_data: Final[dict[str, JsonValue]] = {"file_data": {"mime_type": "video/mp4", "file_uri": uri}} + return file_data if metadata is None else {**file_data, "video_metadata": metadata} + + +GCS_PART: Final = gcs_part() +NOISY_BLOCK: Final[dict[str, JsonValue]] = { + **block(file_id=GCS_URI, mime_type="video/mp4", video_metadata={**METADATA, "frame_rate": 2}), + "caption": "unknown block key", +} + + +def list_value(value: JsonValue) -> list[JsonValue]: + assert isinstance(value, list), value + return value + + +def source_of(part: JsonValue) -> str: + item: Final = object_value(part) + if "text" in item: + return string_value(item["text"]) + if "file_data" in item: + return string_value(object_value(item["file_data"])["file_uri"]) + return string_value(object_value(item["inline_data"])["data"]) + + +def vector(parts: list[JsonValue]) -> list[float]: + return [zlib.crc32("|".join(source_of(part) for part in parts).encode()) / 2**32, 0.5] + + +def request_parts(request: Request) -> list[list[JsonValue]]: + body: Final = object_value(json.loads(request.body)) + return [list_value(object_value(object_value(item)["content"])["parts"]) for item in list_value(body["requests"])] + + +def embed_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target.split("?")[0] == TARGET, request.target + embeddings: Final = [{"values": vector(parts)} for parts in request_parts(request)] + return Reply(body=json.dumps({"embeddings": embeddings}).encode()) + + +def wire_parts(wire: Wire) -> list[list[JsonValue]]: + received: Final = wire.drain() + assert len(received) == 1, [item.target for item in received] + return request_parts(received[0]) + + +def embeddings(response: httpx.Response) -> list[JsonValue]: + assert response.status_code == 200, response.text + data: Final = [object_value(item) for item in list_value(object_value(response.json())["data"])] + assert [item["index"] for item in data] == list(range(len(data))), response.text + assert all(item["object"] == "embedding" for item in data), response.text + return [item["embedding"] for item in data] + + +def gemini_model(scenario: Scenario, url: str, **litellm_params: JsonValue) -> str: + return scenario.model(model=f"gemini/{BACKEND}", api_key="scripted-gemini-key", api_base=url, **litellm_params) + + +def embed(gateway: Gateway, model: str, elements: JsonValue, **extra: JsonValue) -> httpx.Response: + return gateway.request("POST", "/v1/embeddings", {"model": model, "input": elements, **extra}) + + +def proxy_root(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def write_config(directory: Path, url: str, name: str) -> Path: + config: Final = directory / f"gemini_embeddings_{uuid.uuid4().hex}.yaml" + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": name, + "litellm_params": { + "model": f"gemini/{BACKEND}", + "api_key": "scripted-gemini-key", + "api_base": url, + }, + } + ], + "litellm_settings": {"drop_params": True}, + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + }, + } + ) + ) + return config + + +def test_string_input_reaches_the_wire_as_a_text_part(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + assert embeddings(embed(gateway, model, ["a red bus"])) == [vector([TEXT_PART])] + assert wire_parts(wire) == [[TEXT_PART]] + + +@pytest.mark.parametrize( + ("file", "part"), + ( + pytest.param({"file_data": DATA_URI, "video_metadata": METADATA}, INLINE_PART, id="data-uri"), + pytest.param({"file_id": GCS_URI, "video_metadata": METADATA}, GCS_PART, id="gcs"), + pytest.param( + {"file_id": GCS_URI, "video_metadata": {"fps": 1, "start_offset": "", "end_offset": LONG_OFFSET}}, + gcs_part(metadata={"fps": 1, "startOffset": "", "endOffset": LONG_OFFSET}), + id="verbatim-strings-and-int-fps", + ), + pytest.param( + {"file_data": DATA_URI, "format": "video/quicktime"}, + {"inline_data": {"mime_type": "video/quicktime", "data": CLIP}}, + id="format-overrides-the-data-uri-mime", + ), + pytest.param( + {"file_id": "gs://scripted-bucket/clips/clip.bin", "format": "video/mp4"}, + gcs_part("gs://scripted-bucket/clips/clip.bin", None), + id="format-names-an-unlisted-extension", + ), + pytest.param({"file_id": GCS_URI}, gcs_part(metadata=None), id="no-metadata"), + pytest.param({"file_id": GCS_URI, "video_metadata": {}}, gcs_part(metadata=None), id="empty-metadata"), + ), +) +def test_file_block_reaches_the_wire_as_one_part( + gateway: Gateway, file: dict[str, JsonValue], part: dict[str, JsonValue] +) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + assert embeddings(embed(gateway, model, [block(**file)])) == [vector([part])] + assert wire_parts(wire) == [[part]] + + +def test_nested_text_and_block_become_one_request_with_two_parts(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + response: Final = embed(gateway, model, [["a red bus", block(file_id=GCS_URI, video_metadata=METADATA)]]) + assert embeddings(response) == [vector([TEXT_PART, GCS_PART])] + assert wire_parts(wire) == [[TEXT_PART, GCS_PART]] + + +def test_repeated_blocks_become_one_request_each(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + response: Final = embed(gateway, model, [block(file_id=GCS_URI, video_metadata=METADATA)] * 2) + assert embeddings(response) == [vector([GCS_PART])] * 2 + assert wire_parts(wire) == [[GCS_PART], [GCS_PART]] + + +@pytest.mark.parametrize( + ("element", "fragment"), + ( + pytest.param(block(file_id=GCS_URI, video_metadata={"fps": "1"}), "file.video_metadata.fps", id="fps-string"), + pytest.param( + block(file_id=GCS_URI, video_metadata={"start_offset": 0}), + "file.video_metadata.start_offset", + id="offset-int", + ), + pytest.param( + block(file_id=GCS_URI, video_metadata={"end_offset": ["3s"]}), + "file.video_metadata.end_offset", + id="offset-list", + ), + pytest.param(block(file_id=GCS_URI, video_metadata="1fps"), "file.video_metadata", id="metadata-string"), + pytest.param( + block(file_id=GCS_URI, video_metadata={**METADATA, "frame_rate": 2}), + "frame_rate", + id="unknown-metadata-key", + ), + pytest.param(block(file_id=GCS_URI, mime_type="video/mp4"), "mime_type", id="unknown-file-key"), + pytest.param({**block(file_id=GCS_URI), "caption": "x"}, "caption", id="unknown-block-key"), + pytest.param({"type": "video", "file": {"file_id": GCS_URI}}, "type", id="wrong-type"), + pytest.param( + block(file_id=GCS_URI, file_data=DATA_URI), + "takes file.file_id or file.file_data, not both", + id="both-sources", + ), + pytest.param(block(video_metadata=METADATA), "needs file.file_id or file.file_data", id="no-source"), + pytest.param(block(file_data=CLIP), ACCEPTED_FORMS, id="bare-base64"), + pytest.param(block(file_id="https://example.com/clip.mp4"), ACCEPTED_FORMS, id="http-url"), + pytest.param(block(file_id=""), ACCEPTED_FORMS, id="empty-file-id"), + pytest.param(block(file_id=GCS_URI, format=""), "file.format", id="empty-format"), + pytest.param(block(file_data=5), "file.file_data", id="file-data-int"), + pytest.param(block(file_data=[DATA_URI]), "file.file_data", id="file-data-list"), + pytest.param(1, "got int", id="int-element"), + ), +) +def test_malformed_blocks_answer_400_before_any_provider_call( + gateway: Gateway, element: JsonValue, fragment: str +) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + response: Final = embed(gateway, model, [element]) + assert response.status_code == 400, response.text + assert fragment in response.text, response.text + assert wire.drain() == (), "the rejected input reached the provider" + + +def test_bare_block_input_answers_400_before_any_provider_call(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + response: Final = embed(gateway, model, block(file_id=GCS_URI, video_metadata=METADATA)) + assert response.status_code == 400, response.text + assert "input must be a string or a list" in response.text, response.text + assert wire.drain() == (), "the rejected input reached the provider" + + +def test_request_drop_params_strips_unknown_keys_at_every_level(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + assert embeddings(embed(gateway, model, [NOISY_BLOCK], drop_params=True)) == [vector([GCS_PART])] + assert wire_parts(wire) == [[GCS_PART]] + + +def test_deployment_drop_params_strips_unknown_keys_at_every_level(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url, drop_params=True) + assert embeddings(embed(gateway, model, [NOISY_BLOCK])) == [vector([GCS_PART])] + assert wire_parts(wire) == [[GCS_PART]] + + +@pytest.mark.timeout(300) +def test_yaml_drop_params_strips_unknown_keys_at_every_level(gateway: Gateway, tmp_path: Path) -> None: + name: Final = f"gemini-embeddings-{uuid.uuid4().hex}" + with wire_server(embed_peer) as wire: + config: Final = write_config(tmp_path, wire.url, name) + with owned_proxy_process(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=config) as owned: + assert embeddings(embed(owned.gateway, name, [NOISY_BLOCK])) == [vector([GCS_PART])] + assert wire_parts(wire) == [[GCS_PART]] + + +def test_openai_sdk_clients_send_blocks_through_the_proxy(gateway: Gateway) -> None: + sync_uri: Final = "gs://scripted-bucket/clips/sync.mp4" + async_uri: Final = "gs://scripted-bucket/clips/async.mp4" + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = gemini_model(scenario, wire.url) + with OpenAI(base_url=f"{proxy_root(gateway)}/v1", api_key=gateway.key, max_retries=0) as client: + served: Final = client.post( + "/embeddings", + body={"model": model, "input": [block(file_id=sync_uri, video_metadata=METADATA)]}, + cast_to=CreateEmbeddingResponse, + ) + assert served.data[0].embedding == vector([gcs_part(sync_uri)]), served.model_dump_json() + assert wire_parts(wire) == [[gcs_part(sync_uri)]] + + async def drive() -> CreateEmbeddingResponse: + async with AsyncOpenAI(base_url=f"{proxy_root(gateway)}/v1", api_key=gateway.key, max_retries=0) as client: + return await client.post( + "/embeddings", + body={"model": model, "input": [block(file_id=async_uri, video_metadata=METADATA)]}, + cast_to=CreateEmbeddingResponse, + ) + + served_async: Final = asyncio.run(drive()) + assert served_async.data[0].embedding == vector([gcs_part(async_uri)]), served_async.model_dump_json() + assert wire_parts(wire) == [[gcs_part(async_uri)]] diff --git a/tests/integration/providers/test_gemini_thinking_replay_wire.py b/tests/integration/providers/test_gemini_thinking_replay_wire.py index 055aa2fee4e..c0fbe0c4fa3 100644 --- a/tests/integration/providers/test_gemini_thinking_replay_wire.py +++ b/tests/integration/providers/test_gemini_thinking_replay_wire.py @@ -6,17 +6,20 @@ import uuid from collections.abc import Iterator, Mapping, Sequence from collections.abc import Set as AbstractSet from dataclasses import dataclass +from pathlib import Path from typing import Final, Literal, TypeAlias import anthropic import httpx import openai import pytest +import yaml from _pytest.mark.structures import ParameterSet from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric import rsa from integration._support.client import Gateway, Scenario, eventually from integration._support.database import read_rows +from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, Wire, wire_server from openai.types.responses import ResponseCompletedEvent from pydantic import JsonValue, TypeAdapter @@ -463,6 +466,26 @@ def _spend_row(*request_ids: str) -> Mapping[str, JsonValue]: return rows[0] +def _spend_proxy_server_request(request_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT proxy_server_request FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (request_id,), + ), + lambda found: len(found) == 1, + seconds=70, + ) + return _JSON_OBJECT.validate_python(rows[0]["proxy_server_request"]) + + +def _config_storing_prompts(directory: Path) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"]["store_prompts_in_spend_logs"] = True + path: Final = directory / "store-prompts.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + def _assert_caller_id(endpoint: Endpoint, stream: bool, caller_id: str, response_id: str) -> None: if endpoint == "messages" and stream: assert caller_id.startswith("msg_"), caller_id @@ -654,6 +677,74 @@ def test_gemini_own_tool_call_signature_is_still_replayed_on_the_function_call_p assert _spend_row(rounds[1])["status"] == "success" +def test_responses_previous_response_id_replays_gemini_session_history( + gateway: Gateway, tmp_path: Path +) -> None: + first_prompt: Final = f"gemini-session-first-{uuid.uuid4().hex}" + second_prompt: Final = f"gemini-session-second-{uuid.uuid4().hex}" + first_answer: Final = f"gemini-answer-first-{uuid.uuid4().hex}" + second_answer: Final = f"gemini-answer-second-{uuid.uuid4().hex}" + first_provider_id: Final = f"gemini-response-first-{uuid.uuid4().hex}" + second_provider_id: Final = f"gemini-response-second-{uuid.uuid4().hex}" + calls: Final = itertools.count() + + def respond(request: Request) -> Reply: + if next(calls) == 0: + return _reply(first_provider_id, stream=False, parts=({"text": first_answer},)) + return _reply(second_provider_id, stream=False, parts=({"text": second_answer},)) + + with ( + wire_server(respond) as wire, + owned_proxy(gateway, tmp_path, {}, config=_config_storing_prompts(tmp_path), workers=1) as prompt_gateway, + prompt_gateway.scenario() as scenario, + ): + model: Final = _register(prompt_gateway, scenario, "gemini", wire) + first_response: Final = prompt_gateway.request( + "POST", + "/v1/responses", + {"model": model, "input": first_prompt, **_CACHE_BUST}, + ) + assert first_response.status_code == 200, first_response.text + first_body: Final = _JSON_OBJECT.validate_json(first_response.content) + first_response_id: Final = str(first_body["id"]) + assert _output_text(first_body) == first_answer + first_spend_id: Final = _upstream_response_id("responses", first_response_id) + assert _spend_row(first_spend_id)["status"] == "success" + first_proxy_request: Final = _spend_proxy_server_request(first_spend_id) + assert first_proxy_request["input"] == first_prompt, first_proxy_request + + second_response: Final = prompt_gateway.request( + "POST", + "/v1/responses", + { + "model": model, + "input": second_prompt, + "previous_response_id": first_response_id, + **_CACHE_BUST, + }, + ) + assert second_response.status_code == 200, second_response.text + second_body: Final = _JSON_OBJECT.validate_json(second_response.content) + second_response_id: Final = str(second_body["id"]) + assert _output_text(second_body) == second_answer + assert _spend_row(_upstream_response_id("responses", second_response_id))["status"] == "success" + + requests: Final = wire.drain() + assert [(request.method, request.target) for request in requests] == [ + ("POST", _target("gemini", False)), + ("POST", _target("gemini", False)), + ] + first_provider_body: Final = _JSON_OBJECT.validate_json(requests[0].body) + second_provider_body: Final = _JSON_OBJECT.validate_json(requests[1].body) + assert first_provider_body == _expected_body("responses", _user(first_prompt)) + assert second_provider_body == _expected_body( + "responses", + _user(first_prompt), + _model_turn({"text": first_answer}), + _user(second_prompt), + ) + + _FIVE_KB: Final = "s" * 5000 _ASSISTANT_SHAPES: Final = ( pytest.param({"reasoning_content": _REASONING, "thinking_blocks": 5}, _REPLAYED_TURN, id="blocks-int"), diff --git a/tests/integration/providers/test_ocr_router_wire.py b/tests/integration/providers/test_ocr_router_wire.py new file mode 100644 index 00000000000..cc6455aa5f1 --- /dev/null +++ b/tests/integration/providers/test_ocr_router_wire.py @@ -0,0 +1,69 @@ +import json +from typing import Final + +import pytest + +from integration._support.client import Gateway, eventually, string_value +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, wire_server + +_DOCUMENT_URL: Final = "https://example.com/doc.pdf" +_OCR_COST_PER_PAGE: Final = 0.0125 +_MISTRAL_OCR_BODY: Final = json.dumps( + { + "model": "mistral-ocr-latest", + "pages": [{"index": 0, "markdown": "Test PDF File"}], + "usage_info": {"pages_processed": 1, "doc_size_bytes": 1024}, + } +).encode() + + +def _mistral_ocr_peer(request: Request) -> Reply: + return Reply(body=_MISTRAL_OCR_BODY) + + +def test_router_aocr_routes_to_mistral_and_logs_spend(gateway: Gateway) -> None: + with wire_server(_mistral_ocr_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="mistral/mistral-ocr-latest", + api_base=wire.url, + api_key="fake-mistral-key", + ocr_cost_per_page=_OCR_COST_PER_PAGE, + ) + response: Final = gateway.request( + "POST", "/v1/ocr", {"model": model, "document": {"type": "document_url", "document_url": _DOCUMENT_URL}} + ) + assert response.status_code == 200, response.text + upstream: Final = wire.drain() + assert len(upstream) == 1, upstream + assert (upstream[0].method, upstream[0].target) == ("POST", "/v1/ocr"), upstream[0] + sent: Final = json.loads(upstream[0].body) + assert sent["model"] == "mistral-ocr-latest", sent + assert sent["document"]["type"] == "document_url", sent + assert sent["document"]["document_url"] == _DOCUMENT_URL, sent + payload: Final = response.json() + assert payload["object"] == "ocr", payload + assert payload["model"] == model, payload + assert [page["index"] for page in payload["pages"]] == [0], payload + assert payload["pages"][0]["markdown"] == "Test PDF File", payload + assert payload["usage_info"]["pages_processed"] == len(payload["pages"]), payload + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(_OCR_COST_PER_PAGE), response.headers + + request_id: Final = string_value(response.headers["x-litellm-call-id"]) + rows: Final = eventually( + lambda: read_rows( + "SELECT status, call_type, custom_llm_provider, model, model_group, spend, prompt_tokens, " + 'completion_tokens, total_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (request_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + row: Final = rows[0] + assert row["status"] == "success", row + assert row["call_type"] == "aocr", row + assert row["custom_llm_provider"] == "mistral", row + assert row["model"] == "mistral/mistral-ocr-latest", row + assert row["model_group"] == model, row + assert float(row["spend"]) == pytest.approx(_OCR_COST_PER_PAGE), row + assert (row["prompt_tokens"], row["completion_tokens"], row["total_tokens"]) == (0, 0, 0), row diff --git a/tests/integration/providers/test_openai_dialect_prompt_cache_breakpoint_owned_proxy.py b/tests/integration/providers/test_openai_dialect_prompt_cache_breakpoint_owned_proxy.py new file mode 100644 index 00000000000..cc7ab06da90 --- /dev/null +++ b/tests/integration/providers/test_openai_dialect_prompt_cache_breakpoint_owned_proxy.py @@ -0,0 +1,385 @@ +import asyncio +import re +import signal +import threading +from collections.abc import Callable, Iterator +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment +from integration._support.database import read_rows +from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process +from integration._support.responses_vendor import newest_marker +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.providers._cache_control_marks_support import owned_config +from integration.providers._mantle_gpt_prompt_cache_support import ( + EXPLICIT, + GPT, + IMPLICIT, + SYSTEM, + SYSTEM_POINT, + TOKEN, + Outcome, + assert_answered, + assert_burst_landed, + assert_wire, + body_of, + breakpoint_count, + expected_wire, + fresh_marker, + mantle_deployment, + mantle_peer, + observe, + openai_shaped_peer, + parse_outcome, + plan_burst, + prompt_text, + request_body, + row_key, + send, + settled, + spend_rows, + success_row, +) +from pydantic import JsonValue + +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_FOUNDRY_BASE: Final = "http://prompt-cache-breakpoint-audit.services.ai.azure.com" +_FOUNDRY_HOST: Final = urlsplit(_FOUNDRY_BASE).netloc +_FOUNDRY_MODEL: Final = "gpt-6-astra" +_FOUNDRY_KEY: Final = "synthetic-foundry-key" +_FOUNDRY_CHAT_PATH: Final = "/models/chat/completions" +_CELL_TIMEOUT: Final = int(2 * graceful_stop_seconds() + 120) +_RESTART_TIMEOUT: Final = int(4 * graceful_stop_seconds() + 240) +_HELD_BURST: Final = 20 +_SPREAD_BURST: Final = 12 +_SPREAD_ATTEMPTS: Final = 20 +_MANTLE_NAME: Final = "mantle-gpt-cache-owned" + + +def _peer(release: threading.Event, held: SimpleQueue[str]) -> Callable[[Request], Reply]: + mantle: Final = mantle_peer() + foundry: Final = openai_shaped_peer() + + def respond(request: Request) -> Reply: + if urlsplit(request.target).netloc == _FOUNDRY_HOST: + return foundry(request) + marker: Final = newest_marker(request.body.decode()) + assert marker is not None, request.body + held.put(marker) + assert release.wait(timeout=120), "The held burst was never released" + return mantle(request) + + return respond + + +@dataclass(frozen=True, slots=True) +class _Rig: + gateway: Gateway + wire: Wire + owned: OwnedProxy + release: threading.Event + held: SimpleQueue[str] + + +def _overrides(wire: Wire) -> MappingProxyType[str, str]: + return MappingProxyType({"HTTP_PROXY": wire.url, "NO_PROXY": "127.0.0.1,localhost", "AIOHTTP_TRUST_ENV": "True"}) + + +def _started_workers(log: Path) -> tuple[int, ...]: + return tuple(int(match.group(1)) for match in _STARTED_WORKER.finditer(log.read_text())) + + +def _live_workers(log: Path) -> tuple[int, ...]: + return tuple(pid for pid in _started_workers(log) if psutil.pid_exists(pid)) + + +def _open_peer_connections(pid: int, peer_url: str) -> int: + port: Final = urlsplit(peer_url).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +def _drain(queue: SimpleQueue[str]) -> None: + while not queue.empty(): + queue.get_nowait() + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]: + directory: Final = tmp_path_factory.mktemp("openai-dialect-prompt-cache-breakpoint") + release: Final = threading.Event() + release.set() + held: Final[SimpleQueue[str]] = SimpleQueue() + with gateway_from_environment() as environment, wire_server(_peer(release, held)) as wire: + with owned_proxy_process(environment, directory, _overrides(wire), workers=2) as owned: + eventually(lambda: len(_started_workers(owned.log)), lambda count: count == 2, seconds=60) + wire.drain() + yield _Rig(owned.gateway, wire, owned, release, held) + + +def _foundry_deployment(rig: _Rig, scenario: Scenario, *, bare: bool) -> str: + name: Final = ( + scenario.model( + model=_FOUNDRY_MODEL, + custom_llm_provider="azure_ai", + api_base=_FOUNDRY_BASE, + api_key=_FOUNDRY_KEY, + cache_control_injection_points=SYSTEM_POINT, + ) + if bare + else scenario.model( + model=f"azure_ai/{_FOUNDRY_MODEL}", + api_base=_FOUNDRY_BASE, + api_key=_FOUNDRY_KEY, + cache_control_injection_points=SYSTEM_POINT, + ) + ) + settled(rig.gateway, name, rig.wire) + return name + + +def _assert_foundry_chat_wire(received: Request, prompt: str, *, options: JsonValue | None) -> None: + body: Final = body_of(received) + assert urlsplit(received.target).netloc == _FOUNDRY_HOST, received.target + assert urlsplit(received.target).path == _FOUNDRY_CHAT_PATH, received.target + assert "prompt_cache_breakpoint" not in received.body.decode(), received.body + assert body["model"] == _FOUNDRY_MODEL, received.body + assert body["messages"] == [{"role": "system", "content": SYSTEM}, {"role": "user", "content": prompt}], ( + received.body + ) + assert body.get("prompt_cache_options") == options, received.body + + +@pytest.mark.timeout(_CELL_TIMEOUT) +@pytest.mark.parametrize("bare", [False, True], ids=["prefixed", "bare-provider-field"]) +def test_b9_b12_foundry_gpt6_responses_reach_the_openai_dialect_on_the_foundry_host(rig: _Rig, bare: bool) -> None: + marker: Final = fresh_marker() + with rig.gateway.scenario() as scenario: + name: Final = _foundry_deployment(rig, scenario, bare=bare) + outcome, received = observe( + rig.gateway, rig.wire, "responses", request_body("responses", name, prompt_text(marker)) + ) + assert_answered(outcome, marker) + assert urlsplit(received.target).netloc == _FOUNDRY_HOST, received.target + assert_wire( + received, + expected_wire(_FOUNDRY_MODEL, prompt_text(marker), endpoint="responses", marked=True), + streaming=False, + ) + success_row(name, marker) + + +@pytest.mark.timeout(_CELL_TIMEOUT) +def test_b10_foundry_gpt6_messages_bridge_to_chat_without_a_breakpoint(rig: _Rig) -> None: + marker: Final = fresh_marker() + with rig.gateway.scenario() as scenario: + name: Final = _foundry_deployment(rig, scenario, bare=False) + outcome, received = observe( + rig.gateway, rig.wire, "messages", request_body("messages", name, prompt_text(marker)) + ) + assert_answered(outcome, marker) + _assert_foundry_chat_wire(received, prompt_text(marker), options=IMPLICIT) + success_row(name, marker) + + +@pytest.mark.timeout(_CELL_TIMEOUT) +def test_b11_foundry_gpt6_chat_stays_on_the_plain_chat_wire(rig: _Rig) -> None: + marker: Final = fresh_marker() + with rig.gateway.scenario() as scenario: + name: Final = _foundry_deployment(rig, scenario, bare=False) + outcome, received = observe(rig.gateway, rig.wire, "chat", request_body("chat", name, prompt_text(marker))) + assert_answered(outcome, marker) + _assert_foundry_chat_wire(received, prompt_text(marker), options=IMPLICIT) + success_row(name, marker) + + +def _spread_burst(rig: _Rig, name: str, workers: tuple[int, ...]) -> dict[int, int]: + plan: Final = plan_burst(_SPREAD_BURST) + _drain(rig.held) + rig.wire.drain() + rig.release.clear() + with ThreadPoolExecutor(max_workers=_SPREAD_BURST) as pool: + futures: Final = tuple( + pool.submit(send, rig.gateway, endpoint, request_body(endpoint, name, prompt_text(marker), stream=stream)) + for endpoint, stream, marker in plan + ) + eventually(rig.held.qsize, lambda size: size == _SPREAD_BURST, seconds=60) + held_by: Final = {pid: _open_peer_connections(pid, rig.wire.url) for pid in workers} + rig.release.set() + outcomes: Final = tuple(future.result() for future in futures) + assert sum(held_by.values()) == _SPREAD_BURST, held_by + assert_burst_landed( + rig.wire, name, tuple((marker, outcome) for (_, _, marker), outcome in zip(plan, outcomes)), marked=True + ) + return held_by + + +@pytest.mark.timeout(_CELL_TIMEOUT) +def test_e6_every_worker_of_a_two_worker_proxy_marks_the_mantle_request(rig: _Rig) -> None: + with rig.gateway.scenario() as scenario: + name: Final = mantle_deployment(rig.gateway, scenario, rig.wire) + workers: Final = _live_workers(rig.owned.log) + assert len(workers) == 2, workers + for _ in range(_SPREAD_ATTEMPTS): + if all(count > 0 for count in _spread_burst(rig, name, workers).values()): + break + else: + raise AssertionError("One worker never took a marked request") + + +def _mantle_config(wire: Wire, directory: Path, **litellm_params: JsonValue) -> Path: + return owned_config( + directory, + [ + { + "model_name": _MANTLE_NAME, + "litellm_params": { + "model": GPT, + "api_base": wire.url, + "api_key": TOKEN, + "aws_region_name": "us-east-1", + "cache_control_injection_points": SYSTEM_POINT, + **litellm_params, + }, + } + ], + ) + + +async def _one(client: httpx.AsyncClient, key: str, body: dict[str, JsonValue]) -> Outcome: + response: Final = await client.post("/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {key}"}) + lines: Final = tuple(line for line in response.text.splitlines() if line) + return parse_outcome("chat", stream=False, status=response.status_code, headers=response.headers, lines=lines) + + +async def _held_burst(url: str, key: str, markers: tuple[str, ...]) -> tuple[tuple[str, Outcome], ...]: + async with httpx.AsyncClient(base_url=url, timeout=180, trust_env=False) as client: + results: Final = await asyncio.gather( + *(_one(client, key, request_body("chat", _MANTLE_NAME, prompt_text(marker))) for marker in markers), + return_exceptions=True, + ) + return tuple((marker, result) for marker, result in zip(markers, results) if isinstance(result, Outcome)) + + +def _assert_served_and_landed(wire: Wire, served: tuple[tuple[str, Outcome], ...], *, expected_received: int) -> None: + for marker, outcome in served: + assert_answered(outcome, marker) + received: Final = wire.drain() + assert len(received) == expected_received, (len(received), expected_received) + for request in received: + assert breakpoint_count(body_of(request)) == 1, request.body + identities: Final = frozenset(outcome.response_id for _, outcome in served) + assert len(identities) == len(served), identities + rows: Final = spend_rows(_MANTLE_NAME, frozenset(marker for marker, _ in served), expected=len(served), seconds=120) + assert sorted(row_key(str(row["request_id"])) for row in rows) == sorted(identities), rows + + +@pytest.mark.timeout(_CELL_TIMEOUT) +async def test_f2_a_worker_killed_mid_burst_leaves_the_survivor_marking_requests( + gateway: Gateway, tmp_path: Path +) -> None: + release: Final = threading.Event() + held: Final[SimpleQueue[str]] = SimpleQueue() + markers: Final = tuple(fresh_marker() for _ in range(_HELD_BURST)) + with wire_server(_peer(release, held)) as wire: + config: Final = _mantle_config(wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + url: Final = str(candidate.client.base_url).rstrip("/") + workers: Final = eventually(lambda: _live_workers(owned.log), lambda pids: len(pids) == 2, seconds=60) + burst: Final = asyncio.create_task(_held_burst(url, candidate.key, markers)) + await asyncio.to_thread(eventually, held.qsize, lambda size: size == _HELD_BURST, 60) + held_by: Final = MappingProxyType({pid: _open_peer_connections(pid, wire.url) for pid in workers}) + assert sum(held_by.values()) == _HELD_BURST, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + served: Final = await burst + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + follow_up: Final = fresh_marker() + (answered,) = await _held_burst(url, candidate.key, (follow_up,)) + _assert_served_and_landed(wire, (*served, answered), expected_received=_HELD_BURST + 1) + eventually(lambda: len(_started_workers(owned.log)), lambda count: count == 3, seconds=90) + + +def _model_id(name: str) -> str: + (row,) = read_rows('SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_name=%s', (name,)) + return str(row["model_id"]) + + +@pytest.mark.timeout(_RESTART_TIMEOUT) +async def test_f3_a_graceful_restart_mid_burst_drains_the_held_requests_and_keeps_the_stored_explicit_mode( + gateway: Gateway, tmp_path: Path +) -> None: + release: Final = threading.Event() + held: Final[SimpleQueue[str]] = SimpleQueue() + markers: Final = tuple(fresh_marker() for _ in range(_HELD_BURST)) + stored: Final = f"mantle-gpt-explicit-stored-{fresh_marker()}" + with wire_server(_peer(release, held)) as wire: + config: Final = _mantle_config(wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + url: Final = str(candidate.client.base_url).rstrip("/") + eventually(lambda: len(_live_workers(owned.log)), lambda count: count == 2, seconds=60) + candidate.post( + "/model/new", + { + "model_name": stored, + "litellm_params": { + "model": GPT, + "api_base": wire.url, + "api_key": TOKEN, + "aws_region_name": "us-east-1", + "cache_control_injection_points": SYSTEM_POINT, + "prompt_cache_options": EXPLICIT, + }, + }, + ) + release.set() + settled(candidate, stored, wire) + before, before_received = observe( + candidate, wire, "responses", request_body("responses", stored, prompt_text(markers[0])) + ) + assert_answered(before, markers[0]) + assert_wire( + before_received, + expected_wire(GPT, prompt_text(markers[0]), endpoint="responses", marked=True, options=EXPLICIT), + streaming=False, + ) + success_row(stored, markers[0]) + _drain(held) + release.clear() + burst: Final = asyncio.create_task(_held_burst(url, candidate.key, markers)) + await asyncio.to_thread(eventually, held.qsize, lambda size: size == _HELD_BURST, 60) + owned.process.terminate() + release.set() + served: Final = await burst + assert len(served) == _HELD_BURST, len(served) + _assert_served_and_landed(wire, served, expected_received=_HELD_BURST) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as restarted: + settled(restarted.gateway, stored, wire) + after, after_received = observe( + restarted.gateway, wire, "responses", request_body("responses", stored, prompt_text(markers[1])) + ) + assert_answered(after, markers[1]) + assert_wire( + after_received, + expected_wire(GPT, prompt_text(markers[1]), endpoint="responses", marked=True, options=EXPLICIT), + streaming=False, + ) + success_row(stored, markers[1]) + restarted.gateway.post("/model/delete", {"id": _model_id(stored)}) diff --git a/tests/integration/providers/test_openai_passthrough_files_wire.py b/tests/integration/providers/test_openai_passthrough_files_wire.py new file mode 100644 index 00000000000..bf62db8aec6 --- /dev/null +++ b/tests/integration/providers/test_openai_passthrough_files_wire.py @@ -0,0 +1,49 @@ +import uuid +from pathlib import Path +from typing import Final + +from integration._support.client import Gateway +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server + +_UPSTREAM_KEY: Final = "synthetic-openai-key" + + +def test_openai_passthrough_file_upload_and_delete(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "openai-file-" + uuid.uuid4().hex + file_id: Final = f"file-{marker}" + + def respond(request: Request) -> Reply: + assert request.headers["authorization"] == f"Bearer {_UPSTREAM_KEY}", request.headers + if request.method == "POST" and request.target == "/files": + assert request.headers["content-type"].startswith("multipart/form-data"), request.headers + assert b'name="purpose"\r\n\r\nassistants\r\n' in request.body, request.body[:400] + assert b'filename="notes.txt"' in request.body and marker.encode() in request.body, request.body[:400] + return Reply( + body=( + b'{"id": "' + file_id.encode() + b'", "object": "file", "bytes": 12, ' + b'"created_at": 1700000000, "purpose": "assistants", "filename": "notes.txt"}' + ), + ) + if request.method == "DELETE" and request.target == f"/files/{file_id}": + return Reply(body=b'{"id": "' + file_id.encode() + b'", "object": "file", "deleted": true}') + return Reply(status=404) + + with wire_server(respond) as wire: + with owned_proxy( + gateway, + tmp_path, + {"OPENAI_API_BASE": wire.url, "OPENAI_API_KEY": _UPSTREAM_KEY}, + ) as candidate: + upload: Final = candidate.request_multipart( + "/openai/files", + {"purpose": "assistants"}, + {"file": ("notes.txt", f"contents {marker}".encode(), "text/plain")}, + ) + assert upload.status_code == 200, upload.text + assert upload.json()["id"] == file_id + delete: Final = candidate.request("DELETE", f"/openai/files/{file_id}") + assert delete.status_code == 200, delete.text + assert delete.json()["deleted"] is True + forwarded: Final = tuple((request.method, request.target) for request in wire.drain()) + assert forwarded == (("POST", "/files"), ("DELETE", f"/files/{file_id}")), forwarded diff --git a/tests/integration/providers/test_responses_bridge_input_audio_wire.py b/tests/integration/providers/test_responses_bridge_input_audio_wire.py new file mode 100644 index 00000000000..389d31dae3e --- /dev/null +++ b/tests/integration/providers/test_responses_bridge_input_audio_wire.py @@ -0,0 +1,536 @@ +import asyncio +import uuid +from collections.abc import Iterator, Mapping, Sequence +from dataclasses import dataclass +from typing import Final, Literal, TypeAlias + +import anthropic +import httpx +import openai +import pytest +from integration._support import prompt_cache_breakpoint as pcb +from integration._support import responses_vendor as rv +from integration._support.client import Gateway, eventually, gateway_from_environment, string_value +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +pytestmark: Final = pytest.mark.timeout(120) + +Mode: TypeAlias = Literal["on", "off"] + +_MODES: Final[tuple[Mode, ...]] = ("on", "off") +_CAPABLE_MODEL: Final = "openai/responses/gpt-audio-mini" +_CAPABLE_WIRE_MODEL: Final = "gpt-audio-mini" +_HOSTILE: Final[tuple[tuple[str, JsonValue], ...]] = ( + ("int", 7), + ("list", ["Zm9v"]), + ("empty-string", ""), + ("5kb-string", "x" * 5000), +) +_HOSTILE_IDS: Final = tuple(name for name, _ in _HOSTILE) +_HOSTILE_VALUES: Final = tuple(value for _, value in _HOSTILE) + + +@dataclass(frozen=True, slots=True) +class _Bridge: + gateway: Gateway + wire: Wire + on: str + off: str + null: str + capable_on: str + base_model_param_on: str + base_model_info_on: str + injecting_off: str + spend: pcb.SpendLogs + + def model(self, mode: Mode) -> str: + return self.on if mode == "on" else self.off + + +@pytest.fixture(scope="module") +def bridge() -> Iterator[_Bridge]: + with ( + wire_server(pcb.respond) as wire, + gateway_from_environment() as gateway, + gateway.scenario() as scenario, + pcb.spend_logs() as spend, + ): + api_base: Final = f"{wire.url}/v1" + yield _Bridge( + gateway, + wire, + scenario.model(model=pcb.MODEL, api_base=api_base, drop_params=True), + scenario.model(model=pcb.MODEL, api_base=api_base), + scenario.model(model=pcb.MODEL, api_base=api_base, drop_params=None), + scenario.model(model=_CAPABLE_MODEL, api_base=api_base, drop_params=True), + scenario.model(model=pcb.MODEL, api_base=api_base, drop_params=True, base_model=_CAPABLE_WIRE_MODEL), + scenario.model( + model=pcb.MODEL, + api_base=api_base, + drop_params=True, + model_info={"base_model": _CAPABLE_WIRE_MODEL}, + ), + scenario.model(model=pcb.MODEL, api_base=api_base, **pcb.INJECTION), + spend, + ) + + +def _v1(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + "/v1" + + +def _text(marker: str) -> dict[str, JsonValue]: + return {"type": "input_text", "text": pcb.prompt(marker)} + + +def _audio_on_wire(payload: JsonValue) -> dict[str, JsonValue]: + return {"type": "input_audio", "input_audio": payload} + + +def _user(marker: str, *parts: JsonValue) -> dict[str, JsonValue]: + return {"role": "user", "content": [pcb.text(pcb.prompt(marker)), *parts]} + + +def _chat(bridge: _Bridge, model: str, messages: Sequence[JsonValue], *, stream: bool = False) -> httpx.Response: + return bridge.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": list(messages), "stream": stream, **pcb.NO_CACHE}, + ) + + +def _completion(response: httpx.Response, marker: str) -> str: + assert response.status_code == 200, response.text + body: Final = rv.JSON_OBJECT.validate_json(response.text) + assert pcb.answers(string_value(body["id"]), marker), body + (choice,) = rv.ITEMS.validate_python(body["choices"]) + assert rv.JSON_OBJECT.validate_python(choice["message"])["content"] == rv.answer(marker), body + return response.headers["x-litellm-call-id"] + + +def _wire_request(bridge: _Bridge, marker: str, *, model: str = "gpt-6.1-sol", stream: bool = False) -> Request: + request: Final = pcb.posted(bridge.wire, marker) + body: Final = pcb.body_of(request) + assert body["model"] == model, body + assert (body.get("stream") is True) is stream, body + return request + + +def _user_content_on_wire( + bridge: _Bridge, marker: str, *, model: str = "gpt-6.1-sol", stream: bool = False +) -> list[dict[str, JsonValue]]: + return pcb.content_of(pcb.input_items(_wire_request(bridge, marker, model=model, stream=stream)), "user") + + +def _expected(mode: Mode, marker: str, *forwarded: JsonValue) -> list[JsonValue]: + return [_text(marker)] if mode == "on" else [_text(marker), *forwarded] + + +def test_audio_part_is_dropped_under_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.on, [_user(marker, pcb.audio())]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker)] + bridge.spend.landed(bridge.on, call_id, marker) + + +def test_audio_part_is_forwarded_as_input_audio_without_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.off, [_user(marker, pcb.audio())]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker), _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + bridge.spend.landed(bridge.off, call_id, marker) + + +@pytest.mark.parametrize("mode", _MODES) +def test_openai_sdk_stream_shapes_the_audio_part_by_mode(bridge: _Bridge, mode: Mode) -> None: + marker: Final = uuid.uuid4().hex + with openai.OpenAI(api_key=bridge.gateway.key, base_url=_v1(bridge.gateway), max_retries=0) as client: + raw: Final = client.chat.completions.with_raw_response.create( + model=bridge.model(mode), + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": pcb.prompt(marker)}, + {"type": "input_audio", "input_audio": {"data": "Zm9v", "format": "wav"}}, + ], + } + ], + stream=True, + extra_body=dict(pcb.NO_CACHE), + ) + chunks: Final = tuple(raw.parse()) + assert chunks and pcb.answers(chunks[0].id, marker), chunks + streamed: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) + assert streamed == rv.answer(marker), chunks + content: Final = _user_content_on_wire(bridge, marker, stream=True) + assert content == _expected(mode, marker, _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))), content + bridge.spend.landed(bridge.model(mode), raw.headers["x-litellm-call-id"], marker) + + +@pytest.mark.parametrize("mode", _MODES) +async def test_async_openai_sdk_shapes_the_audio_part_by_mode(bridge: _Bridge, mode: Mode) -> None: + marker: Final = uuid.uuid4().hex + async with openai.AsyncOpenAI(api_key=bridge.gateway.key, base_url=_v1(bridge.gateway), max_retries=0) as client: + raw: Final = await client.chat.completions.with_raw_response.create( + model=bridge.model(mode), + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": pcb.prompt(marker)}, + {"type": "input_audio", "input_audio": {"data": "Zm9v", "format": "wav"}}, + ], + } + ], + extra_body=dict(pcb.NO_CACHE), + ) + completion: Final = raw.parse() + assert pcb.answers(completion.id, marker), completion + assert completion.choices[0].message.content == rv.answer(marker), completion + content: Final = _user_content_on_wire(bridge, marker) + assert content == _expected(mode, marker, _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))), content + bridge.spend.landed(bridge.model(mode), raw.headers["x-litellm-call-id"], marker) + + +def test_audio_capable_model_keeps_the_audio_part_under_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.capable_on, [_user(marker, pcb.audio())]), marker) + content: Final = _user_content_on_wire(bridge, marker, model=_CAPABLE_WIRE_MODEL) + assert content == [_text(marker), _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))], content + bridge.spend.landed(bridge.capable_on, call_id, marker) + + +def test_base_model_in_litellm_params_keeps_the_audio_part_under_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.base_model_param_on, [_user(marker, pcb.audio())]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker), _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + bridge.spend.landed(bridge.base_model_param_on, call_id, marker) + + +def test_base_model_in_model_info_keeps_the_audio_part_under_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.base_model_info_on, [_user(marker, pcb.audio())]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker), _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + bridge.spend.landed(bridge.base_model_info_on, call_id, marker) + + +@pytest.mark.parametrize("mode", _MODES) +def test_tool_output_audio_part_follows_the_mode(bridge: _Bridge, mode: Mode) -> None: + marker: Final = uuid.uuid4().hex + messages: Final[list[JsonValue]] = [ + {"role": "user", "content": pcb.prompt(marker)}, + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "record", "arguments": "{}"}}], + }, + {"role": "tool", "tool_call_id": "call_1", "content": [pcb.text("heard it"), pcb.audio()]}, + ] + call_id: Final = _completion(_chat(bridge, bridge.model(mode), messages), marker) + output: Final = pcb.function_output(pcb.input_items(_wire_request(bridge, marker)), "call_1") + heard: Final[dict[str, JsonValue]] = {"type": "input_text", "text": "heard it"} + expected: Final[list[JsonValue]] = [heard] if mode == "on" else [heard, _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + assert output == expected, output + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +def test_injected_system_marker_lands_on_a_trailing_audio_part_without_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + system: Final[dict[str, JsonValue]] = {"role": "system", "content": [pcb.text("sys"), pcb.audio()]} + call_id: Final = _completion(_chat(bridge, bridge.injecting_off, [system, _user(marker)]), marker) + request: Final = _wire_request(bridge, marker) + assert pcb.body_of(request)["prompt_cache_options"] == {"mode": "explicit"}, request.body + system_content: Final = pcb.content_of(pcb.input_items(request), "system") + assert system_content == [ + {"type": "input_text", "text": "sys"}, + pcb.marked(_audio_on_wire(dict(pcb.AUDIO_PAYLOAD)), pcb.EXPLICIT), + ], system_content + bridge.spend.landed(bridge.injecting_off, call_id, marker) + + +@pytest.mark.parametrize("mode", _MODES) +def test_assistant_audio_part_follows_the_mode(bridge: _Bridge, mode: Mode) -> None: + marker: Final = uuid.uuid4().hex + messages: Final[list[JsonValue]] = [ + {"role": "user", "content": pcb.prompt(marker)}, + {"role": "assistant", "content": [pcb.text("earlier answer"), pcb.audio()]}, + {"role": "user", "content": "and again"}, + ] + call_id: Final = _completion(_chat(bridge, bridge.model(mode), messages), marker) + earlier: Final = pcb.content_of(pcb.input_items(_wire_request(bridge, marker)), "assistant") + spoken: Final[dict[str, JsonValue]] = {"type": "output_text", "text": "earlier answer"} + expected: Final[list[JsonValue]] = [spoken] if mode == "on" else [spoken, _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + assert earlier == expected, earlier + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +def test_anthropic_sdk_request_carries_no_audio_part_into_the_bridge(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + client: Final = anthropic.Anthropic( + base_url=str(bridge.gateway.client.base_url), api_key=bridge.gateway.key, max_retries=0 + ) + with client: + raw: Final = client.messages.with_raw_response.create( + model=bridge.on, + max_tokens=64, + messages=[{"role": "user", "content": [pcb.text(pcb.prompt(marker)), pcb.audio()]}], + extra_body=dict(pcb.NO_CACHE), + ) + (content,) = raw.parse().content + assert content.type == "text" and content.text == rv.answer(marker), content + assert _user_content_on_wire(bridge, marker) == [_text(marker)] + bridge.spend.landed(bridge.on, raw.headers["x-litellm-call-id"], None) + + +def test_native_responses_audio_part_never_enters_the_bridge(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + audio: Final = _audio_on_wire(dict(pcb.AUDIO_PAYLOAD)) + response: Final = bridge.gateway.request( + "POST", + "/v1/responses", + { + "model": bridge.on, + "input": [{"type": "message", "role": "user", "content": [_text(marker), audio]}], + **pcb.NO_CACHE, + }, + ) + assert response.status_code == 200, response.text + body: Final = rv.JSON_OBJECT.validate_json(response.text) + assert pcb.answers(string_value(body["id"]), marker), body + assert rv.answer(marker) in response.text, response.text + assert _user_content_on_wire(bridge, marker) == [_text(marker), audio] + bridge.spend.landed(bridge.on, response.headers["x-litellm-call-id"], marker) + + +@pytest.mark.parametrize("value", _HOSTILE_VALUES, ids=_HOSTILE_IDS) +def test_hostile_audio_value_is_forwarded_verbatim_without_drop_params(bridge: _Bridge, value: JsonValue) -> None: + marker: Final = uuid.uuid4().hex + hostile: Final[dict[str, JsonValue]] = {"type": "input_audio", "input_audio": value} + call_id: Final = _completion(_chat(bridge, bridge.off, [_user(marker, hostile)]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker), _audio_on_wire(value)] + bridge.spend.landed(bridge.off, call_id, marker) + + +@pytest.mark.parametrize("value", _HOSTILE_VALUES, ids=_HOSTILE_IDS) +def test_hostile_audio_value_is_dropped_under_drop_params(bridge: _Bridge, value: JsonValue) -> None: + marker: Final = uuid.uuid4().hex + hostile: Final[dict[str, JsonValue]] = {"type": "input_audio", "input_audio": value} + call_id: Final = _completion(_chat(bridge, bridge.on, [_user(marker, hostile)]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker)] + bridge.spend.landed(bridge.on, call_id, marker) + + +@pytest.mark.parametrize("mode", _MODES) +def test_audio_part_without_a_payload_follows_the_mode(bridge: _Bridge, mode: Mode) -> None: + marker: Final = uuid.uuid4().hex + bare: Final[dict[str, JsonValue]] = {"type": "input_audio"} + call_id: Final = _completion(_chat(bridge, bridge.model(mode), [_user(marker, bare)]), marker) + content: Final = _user_content_on_wire(bridge, marker) + assert content == _expected(mode, marker, _audio_on_wire(None)), content + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +@pytest.mark.parametrize("mode", _MODES) +def test_two_identical_audio_parts_follow_the_mode(bridge: _Bridge, mode: Mode) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.model(mode), [_user(marker, pcb.audio(), pcb.audio())]), marker) + audio: Final = _audio_on_wire(dict(pcb.AUDIO_PAYLOAD)) + content: Final = _user_content_on_wire(bridge, marker) + assert content == _expected(mode, marker, audio, audio), content + bridge.spend.landed(bridge.model(mode), call_id, marker) + + +def test_audio_only_message_is_forwarded_with_empty_content_under_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + messages: Final[list[JsonValue]] = [ + {"role": "user", "content": [pcb.audio()]}, + {"role": "assistant", "content": "I could not hear that"}, + {"role": "user", "content": pcb.prompt(marker)}, + ] + call_id: Final = _completion(_chat(bridge, bridge.on, messages), marker) + items: Final = pcb.input_items(_wire_request(bridge, marker)) + users: Final = [item for item in items if item.get("role") == "user"] + assert [item["content"] for item in users] == [[], [_text(marker)]], items + bridge.spend.landed(bridge.on, call_id, marker) + + +def test_leading_audio_marker_has_no_preceding_part_to_carry_to(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [pcb.marked(pcb.audio(), pcb.EXPLICIT), pcb.text(pcb.prompt(marker))] + call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker)] + bridge.spend.landed(bridge.on, call_id, marker) + + +def test_each_dropped_audio_marker_moves_to_its_own_preceding_text(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [ + pcb.text(pcb.prompt(marker)), + pcb.marked(pcb.audio(), pcb.EXPLICIT), + pcb.text("second clip follows"), + pcb.marked(pcb.audio(), pcb.EXPLICIT_30M), + ] + call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker) + assert _user_content_on_wire(bridge, marker) == [ + pcb.marked(_text(marker), pcb.EXPLICIT), + {"type": "input_text", "text": "second clip follows", "prompt_cache_breakpoint": pcb.EXPLICIT_30M}, + ] + bridge.spend.landed(bridge.on, call_id, marker) + + +def test_text_marker_wins_over_the_dropped_audio_marker(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [ + pcb.marked(pcb.text(pcb.prompt(marker)), pcb.EXPLICIT_30M), + pcb.marked(pcb.audio(), pcb.EXPLICIT), + ] + call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker) + assert _user_content_on_wire(bridge, marker) == [pcb.marked(_text(marker), pcb.EXPLICIT_30M)] + bridge.spend.landed(bridge.on, call_id, marker) + + +def test_null_drop_params_forwards_the_audio_part(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + call_id: Final = _completion(_chat(bridge, bridge.null, [_user(marker, pcb.audio())]), marker) + assert _user_content_on_wire(bridge, marker) == [_text(marker), _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))] + bridge.spend.landed(bridge.null, call_id, marker) + + +@dataclass(frozen=True, slots=True) +class _Call: + mode: Mode + stream: bool + marker: str + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + text: str + call_id: str + + +def _calls(count: int) -> tuple[_Call, ...]: + return tuple(_Call(_MODES[index % 2], index % 4 >= 2, uuid.uuid4().hex) for index in range(count)) + + +async def _send(client: httpx.AsyncClient, bridge: _Bridge, model: str, call: _Call) -> _Served: + body: Final[Mapping[str, JsonValue]] = { + "model": model, + "messages": [_user(call.marker, pcb.audio())], + "stream": call.stream, + **pcb.NO_CACHE, + } + async with client.stream( + "POST", "/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {bridge.gateway.key}"} + ) as response: + raw: Final = await response.aread() + return _Served(call, response.status_code, raw.decode(), response.headers["x-litellm-call-id"]) + + +async def _burst(bridge: _Bridge, calls: Sequence[_Call], model_for: Mapping[Mode, str]) -> tuple[_Served, ...]: + async with httpx.AsyncClient(base_url=str(bridge.gateway.client.base_url), timeout=60, trust_env=False) as client: + return tuple(await asyncio.gather(*(_send(client, bridge, model_for[call.mode], call) for call in calls))) + + +def _answered_in_its_own_shape(served: _Served) -> None: + assert served.status == 200, served.text + assert set(rv.MARKER.findall(served.text)) == {served.call.marker}, served.text + assert served.text.startswith("data:") == served.call.stream, served.text + assert rv.answer(served.call.marker) in served.text, served.text + + +def _attempts_by_marker(bridge: _Bridge, calls: Sequence[_Call]) -> Mapping[str, tuple[Request, ...]]: + posts: Final = pcb.drained_posts(bridge.wire) + marked: Final = tuple((rv.newest_marker(request.body.decode()), request) for request in posts) + attempts: Final = { + call.marker: tuple(request for marker, request in marked if marker == call.marker) for call in calls + } + assert sum(len(group) for group in attempts.values()) == len(posts), [request.body for request in posts] + assert all(attempts.values()), sorted(marker for marker, group in attempts.items() if not group) + return attempts + + +def _posts_by_marker(bridge: _Bridge, calls: Sequence[_Call]) -> Mapping[str, Request]: + attempts: Final = _attempts_by_marker(bridge, calls) + assert all(len(group) == 1 for group in attempts.values()), { + marker: len(group) for marker, group in attempts.items() + } + return {marker: group[0] for marker, group in attempts.items()} + + +def _landed_once(bridge: _Bridge, model: str, served: Sequence[_Served], *, status: str = "success") -> None: + rows: Final = eventually( + lambda: bridge.spend.rows_for(model), + lambda found: {string_value(row["litellm_call_id"]) for row in found} >= {item.call_id for item in served}, + seconds=70, + ) + by_call: Final = {string_value(row["litellm_call_id"]): row for row in rows} + assert len(by_call) == len(rows), rows + for item in served: + assert by_call[item.call_id]["status"] == status, (item.call_id, by_call[item.call_id]) + + +async def test_mixed_audio_burst_shapes_every_upstream_request_by_its_mode(bridge: _Bridge) -> None: + calls: Final = _calls(24) + served: Final = await _burst(bridge, calls, {"on": bridge.on, "off": bridge.off}) + assert len(served) == 24 + for item in served: + _answered_in_its_own_shape(item) + by_marker: Final = _posts_by_marker(bridge, calls) + for call in calls: + body: Final = pcb.body_of(by_marker[call.marker]) + assert (body.get("stream") is True) is call.stream, body + content: Final = pcb.content_of(pcb.input_items(by_marker[call.marker]), "user") + assert content == _expected(call.mode, call.marker, _audio_on_wire(dict(pcb.AUDIO_PAYLOAD))), content + _landed_once(bridge, bridge.on, tuple(item for item in served if item.call.mode == "on")) + _landed_once(bridge, bridge.off, tuple(item for item in served if item.call.mode == "off")) + + +@dataclass(frozen=True, slots=True) +class _Doomed: + markers: frozenset[str] + + def respond(self, request: Request) -> Reply: + marker: Final = rv.newest_marker(request.body.decode()) if request.method == "POST" else None + if marker in self.markers: + return Reply(drop_connection=True) + return pcb.respond(request) + + +async def test_dropped_upstream_connections_fail_only_their_own_calls(bridge: _Bridge) -> None: + calls: Final = tuple(_Call("on", False, uuid.uuid4().hex) for _ in range(16)) + doomed: Final = _Doomed(frozenset(call.marker for index, call in enumerate(calls) if index % 4 == 0)) + with wire_server(doomed.respond) as wire, bridge.gateway.scenario() as scenario: + model: Final = scenario.model(model=pcb.MODEL, api_base=f"{wire.url}/v1", drop_params=True) + rig: Final = _Bridge( + bridge.gateway, + wire, + model, + model, + model, + model, + model, + model, + model, + bridge.spend, + ) + served: Final = await _burst(rig, calls, {"on": model, "off": model}) + failed: Final = tuple(item for item in served if item.call.marker in doomed.markers) + answered: Final = tuple(item for item in served if item.call.marker not in doomed.markers) + assert (len(failed), len(answered)) == (4, 12), [(item.call.marker, item.status) for item in served] + for item in failed: + assert item.status >= 500, (item.status, item.text) + assert "answer marker" not in item.text, item.text + for item in answered: + _answered_in_its_own_shape(item) + attempts: Final = _attempts_by_marker(rig, calls) + for item in answered: + assert len(attempts[item.call.marker]) == 1, attempts[item.call.marker] + for call in calls: + for request in attempts[call.marker]: + assert pcb.content_of(pcb.input_items(request), "user") == [_text(call.marker)], request.body + _landed_once(rig, model, answered) + _landed_once(rig, model, failed, status="failure") diff --git a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py index 75bd9da7795..9fad3a2b191 100644 --- a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py +++ b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire.py @@ -234,6 +234,28 @@ def test_malformed_marker_is_dropped_on_every_block_kind(bridge: _Bridge, kind: bridge.spend.landed(bridge.on, call_id, marker) +def test_input_audio_block_is_forwarded_with_its_marker_without_drop_params(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [pcb.text(pcb.prompt(marker)), pcb.marked(pcb.audio(), pcb.EXPLICIT)] + call_id: Final = _completion(_chat(bridge, bridge.off, [{"role": "user", "content": content}]), marker) + second: Final = _second_block_on_wire(bridge, marker) + assert second["type"] == "input_audio" and second["input_audio"] == pcb.AUDIO_PAYLOAD, second + pcb.assert_marker(second, pcb.EXPLICIT) + bridge.spend.landed(bridge.off, call_id, marker) + + +def test_input_audio_block_is_dropped_under_drop_params_and_its_marker_moves_to_the_text(bridge: _Bridge) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [pcb.text(pcb.prompt(marker)), pcb.marked(pcb.audio(), pcb.EXPLICIT)] + call_id: Final = _completion(_chat(bridge, bridge.on, [{"role": "user", "content": content}]), marker) + assert _user_block_on_wire(bridge, marker) == { + "type": "input_text", + "text": pcb.prompt(marker), + "prompt_cache_breakpoint": pcb.EXPLICIT, + } + bridge.spend.landed(bridge.on, call_id, marker) + + @pytest.mark.parametrize(("mode", "breakpoint", "expected"), _CASE_VALUES, ids=_CASE_IDS) def test_tool_output_marker(bridge: _Bridge, mode: Mode, breakpoint: JsonValue, expected: JsonValue) -> None: marker: Final = uuid.uuid4().hex @@ -274,17 +296,14 @@ def test_assistant_list_marker(bridge: _Bridge, mode: Mode, breakpoint: JsonValu def test_injected_system_marker_survives_a_trailing_audio_block(bridge: _Bridge) -> None: marker: Final = uuid.uuid4().hex - system: Final[dict[str, JsonValue]] = {"role": "system", "content": [pcb.text("sys"), pcb.block("input_audio", "")]} + system: Final[dict[str, JsonValue]] = {"role": "system", "content": [pcb.text("sys"), pcb.audio()]} messages: Final[list[JsonValue]] = [system, {"role": "user", "content": pcb.prompt(marker)}] call_id: Final = _completion(_chat(bridge, bridge.injecting_on, messages), marker) request: Final = pcb.posted(bridge.wire, marker) body: Final = _wire_body(request) assert body["prompt_cache_options"] == {"mode": "explicit"}, body - first, audio = pcb.content_of(pcb.input_items(request), "system") - assert first == {"type": "input_text", "text": "sys"}, first - assert audio["type"] == "input_text", audio - assert string_value(audio["text"]).startswith("{'type': 'input_audio'"), audio - pcb.assert_marker(audio, pcb.EXPLICIT) + (system_block,) = pcb.content_of(pcb.input_items(request), "system") + assert system_block == {"type": "input_text", "text": "sys", "prompt_cache_breakpoint": pcb.EXPLICIT}, system_block bridge.spend.landed(bridge.injecting_on, call_id, marker) diff --git a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire_chaos.py b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire_chaos.py index 7742d4b9821..6f9a05d8dfa 100644 --- a/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire_chaos.py +++ b/tests/integration/providers/test_responses_bridge_prompt_cache_breakpoint_wire_chaos.py @@ -252,6 +252,26 @@ def test_global_drop_params_drops_a_malformed_marker(global_rig: _GlobalRig, mod spend.landed(model, control_id, control) +@pytest.mark.parametrize("model", (_GLOBAL_UNSET, _GLOBAL_FALSE), ids=("deployment-unset", "deployment-false")) +def test_global_drop_params_drops_the_audio_part_and_carries_its_marker( + global_rig: _GlobalRig, model: str, spend: pcb.SpendLogs +) -> None: + marker: Final = uuid.uuid4().hex + content: Final[list[JsonValue]] = [pcb.text(pcb.prompt(marker)), pcb.marked(pcb.audio(), pcb.EXPLICIT)] + response: Final = global_rig.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": content}], **pcb.NO_CACHE}, + ) + call_id: Final = _completion(response, marker) + assert _user_block_on_wire(global_rig.wire, marker) == { + "type": "input_text", + "text": pcb.prompt(marker), + "prompt_cache_breakpoint": pcb.EXPLICIT, + } + spend.landed(model, call_id, marker) + + async def test_mixed_burst_carries_every_marker_once(gateway: Gateway, spend: pcb.SpendLogs) -> None: calls: Final = _calls(24, _ENDPOINTS) with wire_server(pcb.respond) as wire, gateway.scenario() as scenario: diff --git a/tests/integration/providers/test_responses_error_status_wire.py b/tests/integration/providers/test_responses_error_status_wire.py new file mode 100644 index 00000000000..01372c8ae7a --- /dev/null +++ b/tests/integration/providers/test_responses_error_status_wire.py @@ -0,0 +1,71 @@ +import json +from typing import Final + +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + + +def test_unknown_model_provider_404_surfaces_to_client_as_404(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/responses", request.target + body: Final = json.loads(request.body) + assert body["model"] == "non-existent-model" + return Reply( + status=404, + body=json.dumps( + {"error": {"message": "model not found", "type": "invalid_request_error", "code": "404"}} + ).encode(), + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/responses/non-existent-model", api_base=wire.url, api_key="synthetic-openai-key" + ) + response: Final = gateway.request("POST", "/v1/responses", {"model": model, "input": "say hi"}) + assert response.status_code == 404, response.text + assert len(wire.drain()) == 1 + + +def test_provider_400_for_bad_temperature_surfaces_to_client_as_400(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/responses", request.target + body: Final = json.loads(request.body) + assert body["model"] == "gpt-4o" + assert body["temperature"] == 2000 + return Reply( + status=400, + body=json.dumps( + {"error": {"message": "temperature out of range", "type": "invalid_request_error", "code": "400"}} + ).encode(), + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/responses/gpt-4o", api_base=wire.url, api_key="synthetic-openai-key" + ) + response: Final = gateway.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "temperature": 2000} + ) + assert response.status_code == 400, response.text + assert len(wire.drain()) == 1 + + +def test_cancel_invalid_response_id_surfaces_error_status(gateway: Gateway) -> None: + response_id: Final = "invalid_response_id_12345" + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == f"/responses/{response_id}/cancel", request.target + return Reply( + status=404, + body=json.dumps( + {"error": {"message": "No such response", "type": "invalid_request_error", "code": "404"}} + ).encode(), + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/responses/gpt-4o", api_base=wire.url, api_key="synthetic-openai-key" + ) + response: Final = gateway.request("POST", f"/v1/responses/{response_id}/cancel", {"model": model}) + assert response.status_code == 404, response.text + assert len(wire.drain()) == 1 diff --git a/tests/integration/providers/test_vertex_anthropic_inline_tools_beta_wire.py b/tests/integration/providers/test_vertex_anthropic_inline_tools_beta_wire.py new file mode 100644 index 00000000000..5061f0b34a4 --- /dev/null +++ b/tests/integration/providers/test_vertex_anthropic_inline_tools_beta_wire.py @@ -0,0 +1,589 @@ +import asyncio +import json +import re +import signal +import threading +import uuid +from collections.abc import Callable, Mapping +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from typing import Final + +import anthropic +import httpx +import openai +import psutil +import pytest +import yaml +from integration._support.client import Gateway, Scenario, eventually +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.vertex import service_account_json +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +_INLINE_TOOLS: Final = "inline-tools-2026-09-15" +_THINKING: Final = "interleaved-thinking-2025-05-14" +_VERTEX_UNSUPPORTED: Final = "effort-2025-11-24" +_UNKNOWN: Final = "not-a-beta-2099-01-01" +_BACKEND: Final = "claude-opus-5-5" +_PROJECT: Final = "scripted-project" +_LOCATION: Final = "us-east5" +_MODEL_PATH: Final = f"/v1/projects/{_PROJECT}/locations/{_LOCATION}/publishers/anthropic/models/{_BACKEND}" +_VERTEX_TARGET: Final = f"{_MODEL_PATH}:rawPredict" +_VERTEX_STREAM_TARGET: Final = f"{_MODEL_PATH}:streamRawPredict?alt=sse" +_ANTHROPIC_TARGET: Final = "/v1/messages" +_ANTHROPIC_KEY: Final = "synthetic-anthropic-key" +_REPLY_TEXT: Final = "inline tools beta control" +_OWNED_MODEL: Final = "partner-claude" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_OWNED_PROXY_CELL_SECONDS: Final = 2 * graceful_stop_seconds() + 120 + +_REPLY: Final[dict[str, JsonValue]] = { + "id": "msg_inline_tools_beta", + "type": "message", + "role": "assistant", + "model": _BACKEND, + "content": [{"type": "text", "text": _REPLY_TEXT}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 5, "output_tokens": 3}, +} +_EVENTS: Final[tuple[tuple[str, dict[str, JsonValue]], ...]] = ( + ( + "message_start", + { + "type": "message_start", + "message": {**_REPLY, "content": [], "stop_reason": None, "usage": {"input_tokens": 5, "output_tokens": 1}}, + }, + ), + ("content_block_start", {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}), + ( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _REPLY_TEXT}}, + ), + ("content_block_stop", {"type": "content_block_stop", "index": 0}), + ( + "message_delta", + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 3}, + }, + ), + ("message_stop", {"type": "message_stop"}), +) +_SSE: Final = tuple(f"event: {name}\ndata: {json.dumps(data)}\n\n".encode() for name, data in _EVENTS) +_REPLY_BODY: Final = json.dumps(_REPLY).encode() + + +@dataclass(frozen=True, slots=True) +class _Delivery: + target: str + beta: str | None + body: str + + +def _streams(request: Request) -> bool: + return request.target == _VERTEX_STREAM_TARGET or _JSON_OBJECT.validate_json(request.body).get("stream") is True + + +def _peer(request: Request) -> Reply: + if request.target in (_VERTEX_TARGET, _VERTEX_STREAM_TARGET): + assert request.headers["authorization"] == "Bearer scripted-token", request.headers + elif request.target == _ANTHROPIC_TARGET: + assert request.headers["x-api-key"] == _ANTHROPIC_KEY, request.headers + else: + return Reply(status=404, body=json.dumps({"error": f"unscripted target {request.target}"}).encode()) + return Reply(content_type="text/event-stream", chunks=_SSE) if _streams(request) else Reply(body=_REPLY_BODY) + + +def _holding(held: SimpleQueue[str], release: threading.Event) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + held.put(request.target) + assert release.wait(timeout=60), "Held request was never released" + return _peer(request) + + return respond + + +def _delivered(wire: Wire) -> tuple[_Delivery, ...]: + return tuple( + _Delivery(request.target, request.headers.get("anthropic-beta"), request.body.decode()) + for request in wire.drain() + ) + + +def _one_delivery(wire: Wire, target: str, marker: str) -> str | None: + (sent,) = _delivered(wire) + assert sent.target == target, sent.target + assert marker in sent.body, sent.body + return sent.beta + + +def _vertex_deployment(gateway: Gateway, scenario: Scenario, api_base: str) -> str: + return scenario.model( + model=f"vertex_ai/{_BACKEND}", + api_base=api_base, + api_key=None, + vertex_project=_PROJECT, + vertex_location=_LOCATION, + vertex_credentials=service_account_json(_PROJECT, gateway.upstream_url.rstrip("/")), + ) + + +def _anthropic_deployment(scenario: Scenario, api_base: str) -> str: + return scenario.model(model=f"anthropic/{_BACKEND}", api_base=api_base, api_key=_ANTHROPIC_KEY) + + +def _marker(label: str) -> str: + return f"{label} {uuid.uuid4().hex}" + + +def _messages_body(model: str, marker: str, stream: bool = False) -> dict[str, JsonValue]: + return {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": marker}], "stream": stream} + + +def _chat_body(model: str, marker: str, stream: bool = False) -> dict[str, JsonValue]: + return {"model": model, "messages": [{"role": "user", "content": marker}], "stream": stream} + + +def _responses_body(model: str, marker: str, stream: bool = False) -> dict[str, JsonValue]: + return {"model": model, "input": marker, "stream": stream} + + +def _beta_headers(value: str | None) -> dict[str, str]: + return {} if value is None else {"anthropic-beta": value} + + +def _generated(gateway: Gateway, path: str, body: Mapping[str, JsonValue], beta: str | None) -> httpx.Response: + response: Final = gateway.request("POST", path, body, headers=_beta_headers(beta)) + assert response.status_code == 200, response.text + assert _REPLY_TEXT in response.text, response.text + return response + + +def _vertex_target(stream: bool) -> str: + return _VERTEX_STREAM_TARGET if stream else _VERTEX_TARGET + + +def _clients(stack: ExitStack, base_url: str, count: int) -> tuple[httpx.Client, ...]: + return tuple( + stack.enter_context(httpx.Client(base_url=base_url, timeout=60, trust_env=False)) for _ in range(count) + ) + + +def _post(client: httpx.Client, key: str, path: str, body: Mapping[str, JsonValue], beta: str) -> int: + return client.post( + path, json=dict(body), headers={"Authorization": f"Bearer {key}", **_beta_headers(beta)} + ).status_code + + +def _post_or_dropped(client: httpx.Client, key: str, path: str, body: Mapping[str, JsonValue], beta: str) -> int | None: + try: + return _post(client, key, path, body, beta) + except httpx.TransportError: + return None + + +def _local_port(client: httpx.Client) -> int: + with client.stream("GET", "/health/liveliness") as response: + port: Final = int(response.extensions["network_stream"].get_extra_info("client_addr")[1]) + response.read() + assert response.status_code == 200, response.text + return port + + +def _accepted_client_ports(pid: int, proxy_port: int) -> frozenset[int]: + return frozenset( + connection.raddr.port + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.raddr and connection.laddr.port == proxy_port + ) + + +def _burst_bodies(model: str, label: str, count: int) -> tuple[tuple[str, dict[str, JsonValue], str], ...]: + markers: Final = tuple(_marker(label) for _ in range(count)) + return tuple( + ( + "/v1/messages" if index % 2 == 0 else "/v1/chat/completions", + _messages_body(model, marker, stream=index % 4 == 2) + if index % 2 == 0 + else _chat_body(model, marker, stream=index % 4 == 3), + marker, + ) + for index, marker in enumerate(markers) + ) + + +def _owned_config(path: Path, gateway: Gateway, api_base: str) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + path.write_text( + yaml.safe_dump( + { + **config, + "model_list": [ + { + "model_name": _OWNED_MODEL, + "litellm_params": { + "model": f"vertex_ai/{_BACKEND}", + "api_base": api_base, + "vertex_project": _PROJECT, + "vertex_location": _LOCATION, + "vertex_credentials": service_account_json(_PROJECT, gateway.upstream_url.rstrip("/")), + }, + } + ], + } + ) + ) + return path + + +@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"]) +def test_vertex_messages_forwards_the_inline_tools_beta(gateway: Gateway, stream: bool) -> None: + marker: Final = _marker("vertex messages") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + response: Final = _generated(gateway, "/v1/messages", _messages_body(model, marker, stream), _INLINE_TOOLS) + assert not stream or "event: message_stop" in response.text, response.text + assert _one_delivery(wire, _vertex_target(stream), marker) == _INLINE_TOOLS + + +@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"]) +def test_vertex_chat_completions_forwards_the_inline_tools_beta(gateway: Gateway, stream: bool) -> None: + marker: Final = _marker("vertex chat") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + response: Final = _generated(gateway, "/v1/chat/completions", _chat_body(model, marker, stream), _INLINE_TOOLS) + assert not stream or response.text.rstrip().endswith("data: [DONE]"), response.text + assert _one_delivery(wire, _vertex_target(stream), marker) == _INLINE_TOOLS + + +@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"]) +def test_vertex_responses_forwards_the_inline_tools_beta(gateway: Gateway, stream: bool) -> None: + marker: Final = _marker("vertex responses") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + _generated(gateway, "/v1/responses", _responses_body(model, marker, stream), _INLINE_TOOLS) + assert _one_delivery(wire, _vertex_target(stream), marker) == _INLINE_TOOLS + + +def test_vertex_messages_through_the_anthropic_sdk_forwards_the_inline_tools_beta(gateway: Gateway) -> None: + marker: Final = _marker("anthropic sdk sync") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + with anthropic.Anthropic( + base_url=str(gateway.client.base_url), + api_key=gateway.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=60), + ) as client: + reply: Final = client.beta.messages.create( + model=model, max_tokens=16, messages=[{"role": "user", "content": marker}], betas=[_INLINE_TOOLS] + ) + assert _REPLY_TEXT in reply.model_dump_json(), reply + assert _one_delivery(wire, _VERTEX_TARGET, marker) == _INLINE_TOOLS + + +def test_vertex_messages_through_the_async_anthropic_sdk_forwards_the_inline_tools_beta(gateway: Gateway) -> None: + marker: Final = _marker("anthropic sdk async") + + async def generate(model: str) -> str: + async with anthropic.AsyncAnthropic( + base_url=str(gateway.client.base_url), + api_key=gateway.key, + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=60), + ) as client: + reply: Final = await client.beta.messages.create( + model=model, max_tokens=16, messages=[{"role": "user", "content": marker}], betas=[_INLINE_TOOLS] + ) + return reply.model_dump_json() + + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + assert _REPLY_TEXT in asyncio.run(generate(model)) + assert _one_delivery(wire, _VERTEX_TARGET, marker) == _INLINE_TOOLS + + +def test_vertex_chat_through_the_openai_sdk_forwards_the_inline_tools_beta(gateway: Gateway) -> None: + marker: Final = _marker("openai sdk sync") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + with openai.OpenAI( + base_url=f"{gateway.client.base_url}/v1", + api_key=gateway.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=60), + ) as client: + reply: Final = client.chat.completions.create( + model=model, messages=[{"role": "user", "content": marker}], extra_headers=_beta_headers(_INLINE_TOOLS) + ) + assert reply.choices[0].message.content == _REPLY_TEXT, reply + assert _one_delivery(wire, _VERTEX_TARGET, marker) == _INLINE_TOOLS + + +def test_vertex_chat_through_the_async_openai_sdk_forwards_the_inline_tools_beta(gateway: Gateway) -> None: + marker: Final = _marker("openai sdk async") + + async def generate(model: str) -> str | None: + async with openai.AsyncOpenAI( + base_url=f"{gateway.client.base_url}/v1", + api_key=gateway.key, + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=60), + ) as client: + reply: Final = await client.chat.completions.create( + model=model, messages=[{"role": "user", "content": marker}], extra_headers=_beta_headers(_INLINE_TOOLS) + ) + return reply.choices[0].message.content + + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + assert asyncio.run(generate(model)) == _REPLY_TEXT + assert _one_delivery(wire, _VERTEX_TARGET, marker) == _INLINE_TOOLS + + +@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"]) +def test_anthropic_chat_completions_forwards_the_inline_tools_beta(gateway: Gateway, stream: bool) -> None: + marker: Final = _marker("anthropic chat") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire.url) + response: Final = _generated(gateway, "/v1/chat/completions", _chat_body(model, marker, stream), _INLINE_TOOLS) + assert not stream or response.text.rstrip().endswith("data: [DONE]"), response.text + assert _one_delivery(wire, _ANTHROPIC_TARGET, marker) == _INLINE_TOOLS + + +def test_anthropic_messages_passes_the_inline_tools_beta_through_unfiltered(gateway: Gateway) -> None: + marker: Final = _marker("anthropic messages") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire.url) + _generated(gateway, "/v1/messages", _messages_body(model, marker), f"{_INLINE_TOOLS},{_UNKNOWN}") + assert _one_delivery(wire, _ANTHROPIC_TARGET, marker) == f"{_INLINE_TOOLS},{_UNKNOWN}" + + +def test_vertex_messages_drops_an_unsupported_beta_sent_next_to_inline_tools(gateway: Gateway) -> None: + marker: Final = _marker("vertex unsupported sibling") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + _generated(gateway, "/v1/messages", _messages_body(model, marker), f"{_VERTEX_UNSUPPORTED},{_INLINE_TOOLS}") + assert _one_delivery(wire, _VERTEX_TARGET, marker) == _INLINE_TOOLS + + +def test_vertex_messages_forwards_inline_tools_with_another_supported_beta(gateway: Gateway) -> None: + marker: Final = _marker("vertex supported sibling") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + _generated(gateway, "/v1/messages", _messages_body(model, marker), f"{_THINKING},{_INLINE_TOOLS}") + assert _one_delivery(wire, _VERTEX_TARGET, marker) == f"{_INLINE_TOOLS},{_THINKING}" + + +@pytest.mark.parametrize( + "beta", + [None, _UNKNOWN, "", "5", f"{_UNKNOWN}," * 240, _INLINE_TOOLS.upper()], + ids=["absent", "unknown", "empty", "int", "5kb_junk", "upper_case"], +) +def test_vertex_messages_sends_no_beta_header_when_nothing_survives_the_filter( + gateway: Gateway, beta: str | None +) -> None: + marker: Final = _marker("vertex nothing survives") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + _generated(gateway, "/v1/messages", _messages_body(model, marker), beta) + assert _one_delivery(wire, _VERTEX_TARGET, marker) is None + + +@pytest.mark.parametrize( + "beta", + [f"{_INLINE_TOOLS} , {_INLINE_TOOLS}", f"{_UNKNOWN}," * 240 + _INLINE_TOOLS], + ids=["duplicated_with_spaces", "inside_5kb_junk"], +) +def test_vertex_messages_forwards_inline_tools_once_from_a_hostile_header_value(gateway: Gateway, beta: str) -> None: + marker: Final = _marker("vertex hostile value") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + _generated(gateway, "/v1/messages", _messages_body(model, marker), beta) + assert _one_delivery(wire, _VERTEX_TARGET, marker) == _INLINE_TOOLS + + +def test_vertex_messages_forwards_inline_tools_once_from_a_repeated_header_line(gateway: Gateway) -> None: + marker: Final = _marker("vertex repeated line") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + response: Final = gateway.client.post( + "/v1/messages", + json=_messages_body(model, marker), + headers=httpx.Headers( + [ + ("Authorization", f"Bearer {gateway.key}"), + ("anthropic-beta", _INLINE_TOOLS), + ("anthropic-beta", _INLINE_TOOLS), + ] + ), + ) + assert response.status_code == 200, response.text + assert _one_delivery(wire, _VERTEX_TARGET, marker) == _INLINE_TOOLS + + +def test_vertex_messages_unauthenticated_request_with_the_beta_never_reaches_the_peer(gateway: Gateway) -> None: + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + response: Final = gateway.request( + "POST", + "/v1/messages", + _messages_body(model, _marker("unauthenticated")), + key="sk-not-a-key-this-proxy-issued", + headers=_beta_headers(_INLINE_TOOLS), + ) + assert response.status_code == 401, response.text + assert wire.drain() == () + + +def test_reloading_the_allowlist_keeps_forwarding_inline_tools(gateway: Gateway) -> None: + marker: Final = _marker("after reload") + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + reloaded: Final = gateway.request("POST", "/reload/anthropic_beta_headers") + assert reloaded.status_code == 200, reloaded.text + status: Final = gateway.request("GET", "/schedule/anthropic_beta_headers_reload/status") + assert status.status_code == 200, status.text + assert _JSON_OBJECT.validate_json(status.content)["scheduled"] is False, status.text + _generated(gateway, "/v1/messages", _messages_body(model, marker), _INLINE_TOOLS) + assert _one_delivery(wire, _VERTEX_TARGET, marker) == _INLINE_TOOLS + + +def test_vertex_messages_repeated_requests_each_carry_the_inline_tools_beta(gateway: Gateway) -> None: + markers: Final = tuple(_marker("repeated") for _ in range(2)) + with wire_server(_peer) as wire, gateway.scenario() as scenario: + model: Final = _vertex_deployment(gateway, scenario, wire.url) + for marker in markers: + _generated(gateway, "/v1/messages", _messages_body(model, marker), _INLINE_TOOLS) + delivered: Final = _delivered(wire) + assert tuple((sent.target, sent.beta) for sent in delivered) == ((_VERTEX_TARGET, _INLINE_TOOLS),) * 2 + assert tuple(marker in sent.body for marker, sent in zip(markers, delivered, strict=True)) == (True, True) + + +def test_concurrent_burst_across_both_workers_forwards_inline_tools_on_every_request(gateway: Gateway) -> None: + with ExitStack() as stack: + wire: Final = stack.enter_context(wire_server(_peer)) + scenario: Final = stack.enter_context(gateway.scenario()) + model: Final = _vertex_deployment(gateway, scenario, wire.url) + burst: Final = _burst_bodies(model, "burst", 12) + clients: Final = _clients(stack, str(gateway.client.base_url), len(burst)) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(burst))) + statuses: Final = tuple( + pool.map( + lambda pair: _post(pair[0], gateway.key, pair[1][0], pair[1][1], _INLINE_TOOLS), + zip(clients, burst, strict=True), + ) + ) + assert statuses == (200,) * len(burst), statuses + delivered: Final = _delivered(wire) + assert tuple(sent.beta for sent in delivered) == (_INLINE_TOOLS,) * len(burst), delivered + seen: Final = tuple(sum(marker in sent.body for sent in delivered) for _, _, marker in burst) + assert seen == (1,) * len(burst), seen + + +def test_upstream_outage_mid_run_recovers_with_inline_tools_still_forwarded(gateway: Gateway) -> None: + with ExitStack() as stack: + clients: Final = _clients(stack, str(gateway.client.base_url), 8) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients))) + scenario: Final = stack.enter_context(gateway.scenario()) + + def wave(label: str) -> tuple[tuple[int | None, str], ...]: + markers: Final = tuple(_marker(label) for _ in clients) + statuses: Final = tuple( + pool.map( + lambda pair: _post_or_dropped( + pair[0], gateway.key, "/v1/messages", _messages_body(model, pair[1]), _INLINE_TOOLS + ), + zip(clients, markers, strict=True), + ) + ) + return tuple(zip(statuses, markers, strict=True)) + + def forwarded(wire: Wire, results: tuple[tuple[int | None, str], ...]) -> None: + assert tuple(status for status, _ in results) == (200,) * len(clients), results + delivered: Final = _delivered(wire) + assert tuple(sent.beta for sent in delivered) == (_INLINE_TOOLS,) * len(clients), delivered + seen: Final = tuple(sum(marker in sent.body for sent in delivered) for _, marker in results) + assert seen == (1,) * len(clients), seen + + with wire_server(_peer) as wire: + port: Final = int(wire.url.rsplit(":", 1)[1]) + model: Final = _vertex_deployment(gateway, scenario, wire.url) + forwarded(wire, wave("before outage")) + outage: Final = wave("during outage") + assert all(status is not None and status >= 500 for status, _ in outage), outage + assert gateway.request("GET", "/health/liveliness").status_code == 200 + with wire_server(_peer, port=port) as revived: + forwarded(revived, wave("after outage")) + + +@pytest.mark.timeout(_OWNED_PROXY_CELL_SECONDS) +def test_worker_sigkill_mid_burst_leaves_the_replacement_forwarding_inline_tools( + gateway: Gateway, tmp_path: Path +) -> None: + held: Final[SimpleQueue[str]] = SimpleQueue() + release: Final = threading.Event() + with ExitStack() as stack: + wire: Final = stack.enter_context(wire_server(_holding(held, release))) + stack.callback(release.set) + config: Final = _owned_config(tmp_path / "inline-tools-worker-kill.yaml", gateway, wire.url) + owned: Final = stack.enter_context( + owned_proxy_process( + gateway, tmp_path, {"LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS": "true"}, config=config, workers=2 + ) + ) + proxy_url: Final = owned.gateway.client.base_url + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + clients: Final = _clients(stack, str(proxy_url), 12) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients))) + ports: Final = tuple(_local_port(client) for client in clients) + markers: Final = tuple(_marker("worker kill") for _ in clients) + futures: Final = tuple( + pool.submit( + _post_or_dropped, + client, + gateway.key, + "/v1/messages", + _messages_body(_OWNED_MODEL, marker), + _INLINE_TOOLS, + ) + for client, marker in zip(clients, markers, strict=True) + ) + eventually(held.qsize, lambda size: size == len(clients), seconds=30) + shares: Final = {pid: _accepted_client_ports(pid, proxy_url.port or 0) & frozenset(ports) for pid in workers} + assert sum(map(len, shares.values())) == len(clients), shares + victim: Final = min((pid for pid in workers if shares[pid]), key=lambda pid: len(shares[pid])) + psutil.Process(victim).send_signal(signal.SIGKILL) + release.set() + results: Final = tuple(future.result(timeout=60) for future in futures) + for port, result in zip(ports, results, strict=True): + assert result == (None if port in shares[victim] else 200), (port, result, shares) + eventually(lambda: len(_STARTED_WORKER.findall(owned.log.read_text())), lambda started: started >= 3, 120) + second_wave: Final = tuple(_marker("after kill") for _ in range(6)) + for marker in second_wave: + assert ( + _post( + owned.gateway.client, + gateway.key, + "/v1/messages", + _messages_body(_OWNED_MODEL, marker), + _INLINE_TOOLS, + ) + == 200 + ) + delivered: Final = _delivered(wire) + assert tuple(sent.beta for sent in delivered) == (_INLINE_TOOLS,) * (len(clients) + len(second_wave)), delivered + seen: Final = tuple(sum(marker in sent.body for sent in delivered) for marker in markers + second_wave) + assert seen == (1,) * len(seen), seen + assert owned.process.poll() is None diff --git a/tests/integration/providers/test_vertex_embedding_batch_file_block_wire.py b/tests/integration/providers/test_vertex_embedding_batch_file_block_wire.py new file mode 100644 index 00000000000..042ca8b0889 --- /dev/null +++ b/tests/integration/providers/test_vertex_embedding_batch_file_block_wire.py @@ -0,0 +1,109 @@ +from __future__ import annotations + +import json +from typing import Final +from urllib.parse import unquote + +import httpx +from integration._support.client import Gateway, Scenario +from integration._support.vertex import service_account_json +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue + +BACKEND: Final = "gemini-embedding-2-preview" +PROJECT: Final = "scripted-project" +LOCATION: Final = "us-central1" +BUCKET: Final = "scripted-bucket" +UPLOAD_PREFIX: Final = f"/upload/storage/v1/b/{BUCKET}/o?uploadType=media&name=" +GCS_URI: Final = "gs://scripted-bucket/clips/animals.mp4" +BLOCK: Final[dict[str, JsonValue]] = { + "type": "file", + "file": {"file_id": GCS_URI, "video_metadata": {"fps": 1.0, "start_offset": "0s", "end_offset": "3s"}}, +} +GCS_PART: Final[dict[str, JsonValue]] = { + "file_data": {"mime_type": "video/mp4", "file_uri": GCS_URI}, + "video_metadata": {"fps": 1.0, "startOffset": "0s", "endOffset": "3s"}, +} +TEXT_PART: Final[dict[str, JsonValue]] = {"text": "a red bus"} + + +def row(custom_id: str, model: str, elements: JsonValue) -> dict[str, JsonValue]: + return { + "custom_id": custom_id, + "method": "POST", + "url": "/v1/embeddings", + "body": {"model": model, "input": elements}, + } + + +def stored_row(key: str, parts: list[JsonValue]) -> dict[str, JsonValue]: + return {"key": key, "request": {"content": {"parts": parts}}} + + +def gcs_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target.startswith(UPLOAD_PREFIX), request.target + assert request.headers["authorization"] == "Bearer scripted-token" + name: Final = unquote(request.target.removeprefix(UPLOAD_PREFIX)) + stored: Final = { + "kind": "storage#object", + "id": f"{BUCKET}/{name}/1759950000000000", + "name": name, + "bucket": BUCKET, + "size": str(len(request.body)), + "timeCreated": "2026-10-08T19:00:00.000Z", + "contentType": "application/octet-stream", + } + return Reply(body=json.dumps(stored).encode()) + + +def vertex_batch_model(gateway: Gateway, scenario: Scenario, url: str) -> str: + return scenario.model( + model=f"vertex_ai/{BACKEND}", + api_key=None, + api_base=url, + vertex_project=PROJECT, + vertex_location=LOCATION, + vertex_credentials=service_account_json(PROJECT, gateway.upstream_url), + gcs_bucket_name=BUCKET, + ) + + +def upload(gateway: Gateway, model: str, rows: tuple[dict[str, JsonValue], ...]) -> httpx.Response: + jsonl: Final = "\n".join(json.dumps(line) for line in rows) + "\n" + return gateway.request_multipart( + "/v1/files", + {"purpose": "batch", "target_model_names": model}, + {"file": ("in.jsonl", jsonl.encode(), "application/jsonl")}, + ) + + +def test_block_rows_upload_as_embed_content_requests(gateway: Gateway) -> None: + with wire_server(gcs_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_batch_model(gateway, scenario, wire.url) + rows: Final = ( + row("mixed", model, [BLOCK, "a red bus"]), + row("bare", model, "a red bus"), + row("nested", model, [[BLOCK, "a red bus"]]), + ) + response: Final = upload(gateway, model, rows) + assert response.status_code == 200, response.text + assert response.json()["object"] == "file" and response.json()["purpose"] == "batch", response.text + uploads: Final = wire.drain() + assert len(uploads) == 1, [upload.target for upload in uploads] + stored: Final = tuple(json.loads(line) for line in uploads[0].body.decode().splitlines() if line.strip()) + assert stored == ( + stored_row("mixed#0/2", [GCS_PART]), + stored_row("mixed#1/2", [TEXT_PART]), + stored_row("bare", [TEXT_PART]), + stored_row("nested", [GCS_PART, TEXT_PART]), + ) + + +def test_bare_block_row_fails_the_upload_before_any_storage_write(gateway: Gateway) -> None: + with wire_server(gcs_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_batch_model(gateway, scenario, wire.url) + response: Final = upload(gateway, model, (row("mixed", model, [BLOCK, "a red bus"]), row("bare", model, BLOCK))) + assert response.status_code >= 400, response.text + assert "got dict" in response.text, response.text + assert wire.drain() == (), "a rejected batch file reached the bucket" diff --git a/tests/integration/providers/test_vertex_embedding_file_block_wire.py b/tests/integration/providers/test_vertex_embedding_file_block_wire.py new file mode 100644 index 00000000000..ad66412a27a --- /dev/null +++ b/tests/integration/providers/test_vertex_embedding_file_block_wire.py @@ -0,0 +1,118 @@ +from __future__ import annotations + +import json +from typing import Final + +import httpx +from integration._support.client import Gateway, Scenario, object_value +from integration._support.vertex import service_account_json +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +BACKEND: Final = "gemini-embedding-2-preview" +PROJECT: Final = "scripted-project" +LOCATION: Final = "us-central1" +TARGET: Final = f"/v1/projects/{PROJECT}/locations/{LOCATION}/publishers/google/models/{BACKEND}:embedContent" +GCS_URI: Final = "gs://scripted-bucket/clips/animals.mp4" +METADATA: Final[dict[str, JsonValue]] = {"fps": 1.0, "start_offset": "0s", "end_offset": "3s"} +GCS_PART: Final[dict[str, JsonValue]] = { + "file_data": {"mime_type": "video/mp4", "file_uri": GCS_URI}, + "video_metadata": {"fps": 1.0, "startOffset": "0s", "endOffset": "3s"}, +} +TEXT_PART: Final[dict[str, JsonValue]] = {"text": "a red bus"} +VALUES: Final = [0.75, 0.25] + + +def block(**file: JsonValue) -> dict[str, JsonValue]: + return {"type": "file", "file": file} + + +def embed_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target == TARGET, request.target + assert request.headers["authorization"] == "Bearer scripted-token" + return Reply(body=json.dumps({"embedding": {"values": VALUES}}).encode()) + + +def wire_parts(wire: Wire) -> list[JsonValue]: + received: Final = wire.drain() + assert len(received) == 1, [item.target for item in received] + body: Final = object_value(json.loads(received[0].body)) + parts: Final = object_value(body["content"])["parts"] + assert isinstance(parts, list), body + return parts + + +def vertex_model(gateway: Gateway, scenario: Scenario, url: str, **litellm_params: JsonValue) -> str: + return scenario.model( + model=f"vertex_ai/{BACKEND}", + api_key=None, + api_base=url, + vertex_project=PROJECT, + vertex_location=LOCATION, + vertex_credentials=service_account_json(PROJECT, gateway.upstream_url), + **litellm_params, + ) + + +def embed(gateway: Gateway, model: str, elements: JsonValue, **extra: JsonValue) -> httpx.Response: + return gateway.request("POST", "/v1/embeddings", {"model": model, "input": elements, **extra}) + + +def single_embedding(response: httpx.Response) -> JsonValue: + assert response.status_code == 200, response.text + data: Final = object_value(response.json())["data"] + assert isinstance(data, list) and len(data) == 1, response.text + item: Final = object_value(data[0]) + assert item["index"] == 0 and item["object"] == "embedding", response.text + return item["embedding"] + + +def rejected(gateway: Gateway, wire: Wire, model: str, elements: JsonValue, fragment: str) -> None: + response: Final = embed(gateway, model, elements) + assert response.status_code == 400, response.text + assert fragment in response.text, response.text + assert wire.drain() == (), "the rejected input reached the provider" + + +def test_string_input_reaches_embed_content_as_a_text_part(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + assert single_embedding(embed(gateway, model, ["a red bus"])) == VALUES + assert wire_parts(wire) == [TEXT_PART] + + +def test_file_block_with_video_metadata_reaches_embed_content(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + assert single_embedding(embed(gateway, model, [block(file_id=GCS_URI, video_metadata=METADATA)])) == VALUES + assert wire_parts(wire) == [GCS_PART] + + +def test_request_drop_params_strips_unknown_keys_at_every_level(gateway: Gateway) -> None: + noisy: Final[dict[str, JsonValue]] = { + **block(file_id=GCS_URI, mime_type="video/mp4", video_metadata={**METADATA, "frame_rate": 2}), + "caption": "unknown block key", + } + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + assert single_embedding(embed(gateway, model, [noisy], drop_params=True)) == VALUES + assert wire_parts(wire) == [GCS_PART] + + +def test_bad_fps_answers_400_before_any_provider_call(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + rejected(gateway, wire, model, [block(file_id=GCS_URI, video_metadata={"fps": "1"})], "file.video_metadata.fps") + + +def test_files_api_reference_answers_400_naming_the_gemini_provider(gateway: Gateway) -> None: + with wire_server(embed_peer) as wire, gateway.scenario() as scenario: + model: Final = vertex_model(gateway, scenario, wire.url) + rejected( + gateway, + wire, + model, + [block(file_id="files/abc", video_metadata=METADATA)], + "Gemini Files API references are only supported through the gemini/ provider", + ) diff --git a/tests/integration/providers/test_voyage_mongodb_host_wire.py b/tests/integration/providers/test_voyage_mongodb_host_wire.py new file mode 100644 index 00000000000..1a6ed907c9b --- /dev/null +++ b/tests/integration/providers/test_voyage_mongodb_host_wire.py @@ -0,0 +1,688 @@ +import asyncio +import json +import os +import re +import signal +import socket +import uuid +from collections.abc import Callable, Iterator +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import JSON_OBJECT, Gateway, eventually, gateway_from_environment, object_value +from integration._support.database import read_rows +from integration._support.forward_proxy import ForwardProxy, refusing_forward_proxy +from integration._support.process import owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request, wire_server +from openai import APIStatusError, AsyncOpenAI, OpenAI +from pydantic import JsonValue, TypeAdapter + +EMBEDDING_MODEL: Final = "voyage/voyage-3.5" +CONTEXTUAL_MODEL: Final = "voyage/voyage-context-3" +MULTIMODAL_MODEL: Final = "voyage/voyage-multimodal-3" +RERANK_MODEL: Final = "voyage/rerank-2.5" +MONGODB_KEY: Final = "al-synthetic-mongodb-issued-key" +VOYAGE_KEY: Final = "pa-synthetic-voyage-issued-key" +ENV_TOKEN: Final = "al-synthetic-env-token" +ENV_PRIMARY: Final = "pa-synthetic-env-primary" +ENV_SECONDARY: Final = "al-synthetic-env-secondary" +MONGODB_HOST: Final = "ai.mongodb.com:443" +VOYAGE_HOST: Final = "api.voyageai.com:443" +NO_CACHE: Final[dict[str, JsonValue]] = {"cache": {"no-cache": True}} +REFUSED_STATUS: Final = 500 +REFUSED_MARKER: Final = "403" +OUTAGE_MARKER: Final = "APIConnectionError" +MONGODB_EMBEDDINGS: Final = "voyage-audit-mongodb-embeddings" +MONGODB_CONTEXTUAL: Final = "voyage-audit-mongodb-contextual" +MONGODB_MULTIMODAL: Final = "voyage-audit-mongodb-multimodal" +MONGODB_RERANK: Final = "voyage-audit-mongodb-rerank" +VOYAGE_EMBEDDINGS: Final = "voyage-audit-voyage-embeddings" +VOYAGE_RERANK: Final = "voyage-audit-voyage-rerank" +BLANK_EMBEDDINGS: Final = "voyage-audit-blank-key-embeddings" +BLANK_RERANK: Final = "voyage-audit-blank-key-rerank" +NULL_RERANK: Final = "voyage-audit-null-key-rerank" +SCRIPTED_CHAT: Final = "voyage-audit-scripted-chat" +TOKEN_EMBEDDINGS: Final = "voyage-audit-token-embeddings" +TOKEN_RERANK: Final = "voyage-audit-token-rerank" +TOKEN_YAML_RERANK: Final = "voyage-audit-token-yaml-rerank" +TOKEN_BLANK_EMBEDDINGS: Final = "voyage-audit-token-blank-embeddings" +TOKEN_BLANK_RERANK: Final = "voyage-audit-token-blank-rerank" +TOKEN_NULL_RERANK: Final = "voyage-audit-token-null-rerank" +PRECEDENCE_EMBEDDINGS: Final = "voyage-audit-precedence-embeddings" +PRECEDENCE_RERANK: Final = "voyage-audit-precedence-rerank" +PRECEDENCE_EXPLICIT_EMBEDDINGS: Final = "voyage-audit-precedence-explicit-embeddings" +PRECEDENCE_EXPLICIT_RERANK: Final = "voyage-audit-precedence-explicit-rerank" +VECTOR: Final = [0.1, 0.2, 0.3] +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_PROXY_VARIABLES: Final = frozenset({"HTTPS_PROXY", "HTTP_PROXY", "ALL_PROXY", "NO_PROXY"}) +_JSON_OBJECTS: Final = TypeAdapter(list[dict[str, JsonValue]]) +_WORKER_PIDS: Final = TypeAdapter(list[int]) +_load_yaml: Final[Callable[[str], object]] = yaml.safe_load +_OWNED_PROXY_CELL_SECONDS: Final = 2 * max(30.0, float(os.environ.get("INTEGRATION_PROXY_READY_SECONDS") or 70)) + 120 +pytestmark: Final = pytest.mark.timeout(_OWNED_PROXY_CELL_SECONDS) + + +def _deployment(name: str, model: str, mode: str, **litellm_params: JsonValue) -> dict[str, JsonValue]: + return { + "model_name": name, + "litellm_params": {"model": model, **litellm_params}, + "model_info": {"id": name, "mode": mode}, + } + + +def _isolated_deployments(upstream_url: str) -> list[dict[str, JsonValue]]: + return [ + _deployment(MONGODB_EMBEDDINGS, EMBEDDING_MODEL, "embedding", api_key=MONGODB_KEY), + _deployment(MONGODB_CONTEXTUAL, CONTEXTUAL_MODEL, "embedding", api_key=MONGODB_KEY), + _deployment(MONGODB_MULTIMODAL, MULTIMODAL_MODEL, "embedding", api_key=MONGODB_KEY), + _deployment(MONGODB_RERANK, RERANK_MODEL, "rerank", api_key=MONGODB_KEY), + _deployment(VOYAGE_EMBEDDINGS, EMBEDDING_MODEL, "embedding", api_key=VOYAGE_KEY), + _deployment(VOYAGE_RERANK, RERANK_MODEL, "rerank", api_key=VOYAGE_KEY), + _deployment(BLANK_EMBEDDINGS, EMBEDDING_MODEL, "embedding", api_key=""), + _deployment(BLANK_RERANK, RERANK_MODEL, "rerank", api_key=""), + _deployment(NULL_RERANK, RERANK_MODEL, "rerank", api_key=None), + _deployment( + SCRIPTED_CHAT, + "openai/gpt-4o-mini", + "chat", + api_key="integration-provider-key", + api_base=f"{upstream_url}/v1", + ), + ] + + +def _token_only_deployments() -> list[dict[str, JsonValue]]: + return [ + _deployment(TOKEN_EMBEDDINGS, EMBEDDING_MODEL, "embedding"), + _deployment(TOKEN_RERANK, RERANK_MODEL, "rerank"), + _deployment(TOKEN_YAML_RERANK, RERANK_MODEL, "rerank", api_key="os.environ/VOYAGE_AI_TOKEN"), + _deployment(TOKEN_BLANK_EMBEDDINGS, EMBEDDING_MODEL, "embedding", api_key=""), + _deployment(TOKEN_BLANK_RERANK, RERANK_MODEL, "rerank", api_key=""), + _deployment(TOKEN_NULL_RERANK, RERANK_MODEL, "rerank", api_key=None), + ] + + +def _precedence_deployments() -> list[dict[str, JsonValue]]: + return [ + _deployment(PRECEDENCE_EMBEDDINGS, EMBEDDING_MODEL, "embedding"), + _deployment(PRECEDENCE_RERANK, RERANK_MODEL, "rerank"), + _deployment(PRECEDENCE_EXPLICIT_EMBEDDINGS, EMBEDDING_MODEL, "embedding", api_key=MONGODB_KEY), + _deployment(PRECEDENCE_EXPLICIT_RERANK, RERANK_MODEL, "rerank", api_key=MONGODB_KEY), + ] + + +def _write_config(directory: Path, name: str, model_list: list[dict[str, JsonValue]]) -> Path: + base: Final = JSON_OBJECT.validate_python(_load_yaml(Path("tests/integration/proxy_config.yaml").read_text())) + configuration: Final = { + **base, + "model_list": model_list, + "router_settings": {"disable_cooldowns": True, "num_retries": 0}, + } + path: Final = directory / f"{name}.yaml" + path.write_text(yaml.safe_dump(configuration)) + return path + + +def _scrubbed() -> tuple[str, ...]: + return tuple(name for name in os.environ if name.upper().startswith("VOYAGE_") or name.upper() in _PROXY_VARIABLES) + + +def _forward_environment(forward_url: str, **extra: str) -> dict[str, str]: + return {"HTTPS_PROXY": forward_url, "NO_PROXY": "127.0.0.1,localhost", **extra} + + +@pytest.fixture(scope="module") +def forward() -> Iterator[ForwardProxy]: + with refusing_forward_proxy() as proxy: + yield proxy + + +@pytest.fixture(scope="module") +def module_gateway() -> Iterator[Gateway]: + with gateway_from_environment() as value: + yield value + + +@pytest.fixture(scope="module") +def isolated( + module_gateway: Gateway, forward: ForwardProxy, tmp_path_factory: pytest.TempPathFactory +) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("voyage-isolated") + config: Final = _write_config(directory, "isolated", _isolated_deployments(module_gateway.upstream_url)) + with owned_proxy( + module_gateway, + directory, + _forward_environment(forward.url), + config=config, + remove_environment=_scrubbed(), + workers=2, + ) as candidate: + yield candidate + + +@pytest.fixture(scope="module") +def token_only( + module_gateway: Gateway, forward: ForwardProxy, tmp_path_factory: pytest.TempPathFactory +) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("voyage-token-only") + config: Final = _write_config(directory, "token-only", _token_only_deployments()) + with owned_proxy( + module_gateway, + directory, + _forward_environment(forward.url, VOYAGE_AI_TOKEN=ENV_TOKEN), + config=config, + remove_environment=_scrubbed(), + ) as candidate: + yield candidate + + +@pytest.fixture(scope="module") +def precedence( + module_gateway: Gateway, forward: ForwardProxy, tmp_path_factory: pytest.TempPathFactory +) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("voyage-precedence") + config: Final = _write_config(directory, "precedence", _precedence_deployments()) + with owned_proxy( + module_gateway, + directory, + _forward_environment( + forward.url, + VOYAGE_API_KEY=ENV_PRIMARY, + VOYAGE_AI_API_KEY=ENV_SECONDARY, + VOYAGE_AI_TOKEN=ENV_TOKEN, + ), + config=config, + remove_environment=_scrubbed(), + ) as candidate: + yield candidate + + +def _embedding_body(model: str, **extra: JsonValue) -> dict[str, JsonValue]: + return {"model": model, "input": f"host selection {uuid.uuid4().hex}", **NO_CACHE, **extra} + + +def _rerank_body(model: str, **extra: JsonValue) -> dict[str, JsonValue]: + return { + "model": model, + "query": f"host selection {uuid.uuid4().hex}", + "documents": ["first document", "second document"], + **extra, + } + + +def _embed(candidate: Gateway, model: str, **extra: JsonValue) -> httpx.Response: + return candidate.request("POST", "/v1/embeddings", _embedding_body(model, **extra)) + + +def _rerank(candidate: Gateway, model: str, path: str = "/v1/rerank", **extra: JsonValue) -> httpx.Response: + return candidate.request("POST", path, _rerank_body(model, **extra)) + + +def _error_message(call_id: str) -> str: + rows: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)), + lambda values: len(values) == 1, + seconds=70, + ) + metadata: Final = rows[0]["metadata"] + parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata) + return str(object_value(parsed["error_information"])["error_message"]) + + +def _assert_refused(response: httpx.Response) -> None: + assert response.status_code == REFUSED_STATUS, response.text + assert REFUSED_MARKER in response.text, response.text + assert REFUSED_MARKER in _error_message(response.headers["x-litellm-call-id"]), response.text + + +def _assert_dialed(forward: ForwardProxy, response: httpx.Response, host: str) -> None: + _assert_refused(response) + assert forward.targets() == (host,), response.text + + +def test_scripted_chat_keeps_serving_behind_the_forward_proxy(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + reply: Final = isolated.chat(SCRIPTED_CHAT, text=f"forward proxy control {uuid.uuid4().hex}") + rows: Final = eventually( + lambda: read_rows('SELECT model FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (str(reply["id"]),)), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["model"] == "openai/gpt-4o-mini", rows + assert forward.targets() == () + + +def test_mongodb_key_embeddings_dial_ai_mongodb_openai_sdk_sync(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + with OpenAI( + api_key=isolated.key, + base_url=str(isolated.client.base_url).rstrip("/") + "/v1", + max_retries=0, + http_client=httpx.Client(timeout=15, trust_env=False), + ) as client: + with pytest.raises(APIStatusError) as caught: + client.embeddings.create( + model=MONGODB_EMBEDDINGS, input=f"sdk sync {uuid.uuid4().hex}", extra_body=NO_CACHE + ) + assert caught.value.status_code == REFUSED_STATUS, caught.value.message + assert REFUSED_MARKER in caught.value.message, caught.value.message + assert forward.targets() == (MONGODB_HOST,) + + +async def test_mongodb_key_embeddings_dial_ai_mongodb_openai_sdk_async( + isolated: Gateway, forward: ForwardProxy +) -> None: + forward.drain() + async with AsyncOpenAI( + api_key=isolated.key, + base_url=str(isolated.client.base_url).rstrip("/") + "/v1", + max_retries=0, + http_client=httpx.AsyncClient(timeout=15, trust_env=False), + ) as client: + with pytest.raises(APIStatusError) as caught: + await client.embeddings.create( + model=MONGODB_EMBEDDINGS, input=f"sdk async {uuid.uuid4().hex}", extra_body=NO_CACHE + ) + assert caught.value.status_code == REFUSED_STATUS, caught.value.message + assert REFUSED_MARKER in caught.value.message, caught.value.message + assert forward.targets() == (MONGODB_HOST,) + + +@pytest.mark.parametrize("deployment", (MONGODB_CONTEXTUAL, MONGODB_MULTIMODAL)) +def test_mongodb_key_other_embedding_shapes_dial_ai_mongodb( + isolated: Gateway, forward: ForwardProxy, deployment: str +) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, deployment), MONGODB_HOST) + + +@pytest.mark.parametrize("path", ("/v1/rerank", "/rerank", "/v2/rerank")) +def test_mongodb_key_rerank_dials_ai_mongodb(isolated: Gateway, forward: ForwardProxy, path: str) -> None: + forward.drain() + _assert_dialed(forward, _rerank(isolated, MONGODB_RERANK, path), MONGODB_HOST) + + +def test_voyage_key_embeddings_still_dial_api_voyageai(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, VOYAGE_EMBEDDINGS), VOYAGE_HOST) + + +def test_voyage_key_rerank_still_dials_api_voyageai(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _rerank(isolated, VOYAGE_RERANK), VOYAGE_HOST) + + +def test_request_body_mongodb_key_overrides_voyage_deployment_host(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, VOYAGE_EMBEDDINGS, api_key=MONGODB_KEY), MONGODB_HOST) + + +@pytest.mark.parametrize("deployment", (MONGODB_EMBEDDINGS, MONGODB_RERANK)) +def test_health_check_dials_ai_mongodb_for_mongodb_key( + isolated: Gateway, forward: ForwardProxy, deployment: str +) -> None: + forward.drain() + response: Final = isolated.request("GET", "/health", params={"model": deployment}) + assert response.status_code == 503, response.text + report: Final = JSON_OBJECT.validate_json(response.content) + assert report["unhealthy_count"] == 1 and report["healthy_count"] == 0, report + unhealthy: Final = _JSON_OBJECTS.validate_python(report["unhealthy_endpoints"]) + assert len(unhealthy) == 1 and REFUSED_MARKER in str(unhealthy[0]["error"]), report + assert forward.targets() == (MONGODB_HOST,), report + + +@pytest.mark.parametrize( + ("mode", "model"), (("embedding", EMBEDDING_MODEL), ("rerank", RERANK_MODEL)), ids=("embedding", "rerank") +) +def test_test_connection_dials_ai_mongodb_for_mongodb_key( + isolated: Gateway, forward: ForwardProxy, mode: str, model: str +) -> None: + forward.drain() + report: Final = isolated.post( + "/health/test_connection", {"litellm_params": {"model": model, "api_key": MONGODB_KEY}, "mode": mode} + ) + assert report["status"] == "error", report + assert REFUSED_MARKER in json.dumps(report), report + assert forward.targets() == (MONGODB_HOST,), report + + +@pytest.mark.parametrize("api_key", (5, ["al-list-member"]), ids=("int", "list")) +def test_non_string_request_body_key_is_rejected_by_validation_before_any_dial( + isolated: Gateway, forward: ForwardProxy, api_key: JsonValue +) -> None: + forward.drain() + response: Final = _embed(isolated, MONGODB_EMBEDDINGS, api_key=api_key) + assert response.status_code == 500, response.text + assert "LiteLLM_Params" in response.text and "api_key" in response.text, response.text + assert forward.targets() == () + + +def test_empty_request_body_key_falls_back_to_api_voyageai(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, MONGODB_EMBEDDINGS, api_key=""), VOYAGE_HOST) + + +def test_five_kilobyte_mongodb_request_body_key_dials_ai_mongodb(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, VOYAGE_EMBEDDINGS, api_key=MONGODB_KEY + "x" * 5000), MONGODB_HOST) + + +def test_duplicate_request_body_key_counts_once(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + raw: Final = ( + '{"model": "%s", "input": "duplicate key %s", "cache": {"no-cache": true}, "api_key": "%s", "api_key": "%s"}' + % (VOYAGE_EMBEDDINGS, uuid.uuid4().hex, MONGODB_KEY, MONGODB_KEY) + ) + response: Final = isolated.client.post( + "/v1/embeddings", + content=raw.encode(), + headers={"Authorization": f"Bearer {isolated.key}", "content-type": "application/json"}, + ) + _assert_dialed(forward, response, MONGODB_HOST) + + +def test_blank_deployment_key_without_env_dials_api_voyageai_for_embeddings( + isolated: Gateway, forward: ForwardProxy +) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, BLANK_EMBEDDINGS), VOYAGE_HOST) + + +@pytest.mark.parametrize("deployment", (BLANK_RERANK, NULL_RERANK), ids=("blank", "null")) +def test_rerank_without_any_key_fails_before_dialing(isolated: Gateway, forward: ForwardProxy, deployment: str) -> None: + forward.drain() + response: Final = _rerank(isolated, deployment) + assert response.status_code == 500, response.text + assert "Voyage AI API key is required" in response.text, response.text + assert forward.targets() == () + + +@dataclass(frozen=True, slots=True) +class _Probe: + host: str + response: httpx.Response + + +def _burst_plan(count: int) -> tuple[tuple[str, str, str], ...]: + shapes: Final = ( + ("/v1/embeddings", MONGODB_EMBEDDINGS, MONGODB_HOST), + ("/v1/rerank", MONGODB_RERANK, MONGODB_HOST), + ("/v1/embeddings", VOYAGE_EMBEDDINGS, VOYAGE_HOST), + ("/v1/rerank", VOYAGE_RERANK, VOYAGE_HOST), + ) + return tuple(shapes[index % len(shapes)] for index in range(count)) + + +async def _burst(base_url: str, key: str, count: int, *, tolerate_transport_errors: bool = False) -> tuple[_Probe, ...]: + async def one(client: httpx.AsyncClient, path: str, deployment: str, host: str) -> _Probe: + body: Final = _embedding_body(deployment) if path == "/v1/embeddings" else _rerank_body(deployment) + response: Final = await client.post(path, json=body, headers={"Authorization": f"Bearer {key}"}) + return _Probe(host, response) + + async with httpx.AsyncClient(base_url=base_url, timeout=30, trust_env=False) as client: + results: Final = await asyncio.gather( + *(one(client, path, deployment, host) for path, deployment, host in _burst_plan(count)), + return_exceptions=tolerate_transport_errors, + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Probe)) + + +def _host_counts(targets: tuple[str, ...]) -> dict[str, int]: + return {host: targets.count(host) for host in sorted(set(targets))} + + +async def test_concurrent_mixed_keys_each_dial_their_own_host(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + served: Final = await _burst(str(isolated.client.base_url), isolated.key, 24) + assert len(served) == 24 + for probe in served: + _assert_refused(probe.response) + assert _host_counts(forward.targets()) == {MONGODB_HOST: 12, VOYAGE_HOST: 12} + + +def test_voyage_ai_token_alone_routes_embeddings_to_ai_mongodb(token_only: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _embed(token_only, TOKEN_EMBEDDINGS), MONGODB_HOST) + + +@pytest.mark.parametrize( + "deployment", (TOKEN_RERANK, TOKEN_YAML_RERANK, TOKEN_NULL_RERANK), ids=("missing", "yaml-env", "null") +) +def test_voyage_ai_token_alone_routes_rerank_to_ai_mongodb( + token_only: Gateway, forward: ForwardProxy, deployment: str +) -> None: + forward.drain() + _assert_dialed(forward, _rerank(token_only, deployment), MONGODB_HOST) + + +@pytest.mark.parametrize( + ("path", "deployment"), + (("/v1/embeddings", TOKEN_BLANK_EMBEDDINGS), ("/v1/rerank", TOKEN_BLANK_RERANK)), + ids=("embeddings", "rerank"), +) +def test_blank_deployment_key_falls_through_to_voyage_ai_token( + token_only: Gateway, forward: ForwardProxy, path: str, deployment: str +) -> None: + forward.drain() + response: Final = ( + _embed(token_only, deployment) if path == "/v1/embeddings" else _rerank(token_only, deployment, path) + ) + _assert_dialed(forward, response, MONGODB_HOST) + + +@pytest.mark.parametrize( + ("path", "deployment"), + (("/v1/embeddings", PRECEDENCE_EMBEDDINGS), ("/v1/rerank", PRECEDENCE_RERANK)), + ids=("embeddings", "rerank"), +) +def test_voyage_api_key_wins_over_mongodb_fallback_env( + precedence: Gateway, forward: ForwardProxy, path: str, deployment: str +) -> None: + forward.drain() + response: Final = ( + _embed(precedence, deployment) if path == "/v1/embeddings" else _rerank(precedence, deployment, path) + ) + _assert_dialed(forward, response, VOYAGE_HOST) + + +@pytest.mark.parametrize( + ("path", "deployment"), + (("/v1/embeddings", PRECEDENCE_EXPLICIT_EMBEDDINGS), ("/v1/rerank", PRECEDENCE_EXPLICIT_RERANK)), + ids=("embeddings", "rerank"), +) +def test_explicit_mongodb_deployment_key_wins_over_voyage_env( + precedence: Gateway, forward: ForwardProxy, path: str, deployment: str +) -> None: + forward.drain() + response: Final = ( + _embed(precedence, deployment) if path == "/v1/embeddings" else _rerank(precedence, deployment, path) + ) + _assert_dialed(forward, response, MONGODB_HOST) + + +def _embedding_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.headers["authorization"] == f"Bearer {MONGODB_KEY}", request.headers + body: Final = JSON_OBJECT.validate_json(request.body) + return Reply( + body=json.dumps( + { + "object": "list", + "data": [{"object": "embedding", "embedding": VECTOR, "index": 0}], + "model": str(body["model"]), + "usage": {"total_tokens": 7}, + } + ).encode() + ) + + +def _rerank_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.headers["authorization"] == f"Bearer {MONGODB_KEY}", request.headers + return Reply( + body=json.dumps( + { + "object": "list", + "data": [{"index": 1, "relevance_score": 0.9}, {"index": 0, "relevance_score": 0.1}], + "model": "rerank-2.5", + "usage": {"total_tokens": 11}, + } + ).encode() + ) + + +def _spend_api_base(call_id: str) -> str: + rows: Final = eventually( + lambda: read_rows('SELECT api_base FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)), + lambda values: len(values) == 1, + seconds=70, + ) + return str(rows[0]["api_base"]) + + +@pytest.mark.parametrize( + ("model", "target"), + ( + (EMBEDDING_MODEL, "/embeddings"), + (CONTEXTUAL_MODEL, "/contextualizedembeddings"), + (MULTIMODAL_MODEL, "/multimodalembeddings"), + ), + ids=("embeddings", "contextual", "multimodal"), +) +def test_explicit_api_base_keeps_mongodb_key_embeddings_on_that_base(gateway: Gateway, model: str, target: str) -> None: + with wire_server(_embedding_peer) as wire, gateway.scenario() as scenario: + deployment: Final = scenario.model(model=model, api_key=MONGODB_KEY, api_base=wire.url) + response: Final = _embed(gateway, deployment) + assert response.status_code == 200, response.text + data: Final = _JSON_OBJECTS.validate_python(JSON_OBJECT.validate_json(response.content)["data"]) + assert data[0]["embedding"] == VECTOR, response.text + received: Final = wire.drain() + assert tuple(request.target for request in received) == (target,), received + assert _spend_api_base(response.headers["x-litellm-call-id"]).startswith(wire.url), response.text + + +@pytest.mark.parametrize("suffix", ("", "/v1"), ids=("bare", "v1")) +def test_explicit_api_base_keeps_mongodb_key_rerank_on_that_base(gateway: Gateway, suffix: str) -> None: + with wire_server(_rerank_peer) as wire, gateway.scenario() as scenario: + deployment: Final = scenario.model(model=RERANK_MODEL, api_key=MONGODB_KEY, api_base=wire.url + suffix) + response: Final = _rerank(gateway, deployment) + assert response.status_code == 200, response.text + payload: Final = JSON_OBJECT.validate_json(response.content) + results: Final = _JSON_OBJECTS.validate_python(payload["results"]) + assert [result["index"] for result in results] == [1, 0], response.text + received: Final = wire.drain() + assert tuple(request.target for request in received) == ("/v1/rerank",), received + assert _spend_api_base(str(payload["id"])).startswith(wire.url), response.text + + +def test_public_provider_fields_name_voyage_as_mongodb(gateway: Gateway) -> None: + response: Final = gateway.request("GET", "/public/providers/fields") + assert response.status_code == 200, response.text + voyage: Final = [ + entry for entry in _JSON_OBJECTS.validate_json(response.content) if entry["litellm_provider"] == "voyage" + ] + assert [entry["provider_display_name"] for entry in voyage] == ["VoyageAI by MongoDB"], response.text + + +def test_public_endpoints_name_voyage_as_mongodb(gateway: Gateway) -> None: + response: Final = gateway.request("GET", "/public/endpoints") + assert response.status_code == 200, response.text + endpoints: Final = _JSON_OBJECTS.validate_python(JSON_OBJECT.validate_json(response.content)["endpoints"]) + names: Final = tuple(_voyage_display_names(endpoints)) + assert names and set(names) == {"VoyageAI by MongoDB"}, response.text + + +def _voyage_display_names(endpoints: list[dict[str, JsonValue]]) -> Iterator[str]: + for endpoint in endpoints: + for provider in _JSON_OBJECTS.validate_python(endpoint["providers"]): + if provider["slug"] == "voyage": + yield str(provider["display_name"]) + + +def _served_call_ids(served: tuple[_Probe, ...]) -> tuple[str, ...]: + return tuple(probe.response.headers["x-litellm-call-id"] for probe in served) + + +def _spend_rows(call_id: str) -> list[dict[str, JsonValue]]: + return read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)) + + +def _reserve_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + address: Final[Callable[[], tuple[str, int]]] = reserve.getsockname + return address()[1] + + +async def test_forward_proxy_outage_mid_traffic_recovers_with_the_right_hosts(gateway: Gateway, tmp_path: Path) -> None: + port: Final = _reserve_port() + config: Final = _write_config(tmp_path, "outage", _isolated_deployments(gateway.upstream_url)) + with owned_proxy( + gateway, + tmp_path, + _forward_environment(f"http://127.0.0.1:{port}"), + config=config, + remove_environment=_scrubbed(), + workers=2, + ) as candidate: + with refusing_forward_proxy(port=port) as before: + first: Final = await _burst(str(candidate.client.base_url), candidate.key, 12) + before_counts: Final = _host_counts(before.targets()) + during: Final = await _burst(str(candidate.client.base_url), candidate.key, 12) + liveliness: Final = candidate.request("GET", "/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + with refusing_forward_proxy(port=port) as after: + third: Final = await _burst(str(candidate.client.base_url), candidate.key, 12) + after_counts: Final = _host_counts(after.targets()) + assert len(first) == len(during) == len(third) == 12 + for probe in (*first, *third): + _assert_refused(probe.response) + for probe in during: + assert probe.response.status_code == REFUSED_STATUS, probe.response.text + assert OUTAGE_MARKER in probe.response.text, probe.response.text + assert OUTAGE_MARKER in _error_message(probe.response.headers["x-litellm-call-id"]) + assert before_counts == {MONGODB_HOST: 6, VOYAGE_HOST: 6} + assert after_counts == {MONGODB_HOST: 6, VOYAGE_HOST: 6} + + +async def test_worker_sigkill_mid_traffic_leaves_the_sibling_routing_by_key(gateway: Gateway, tmp_path: Path) -> None: + config: Final = _write_config(tmp_path, "sigkill", _isolated_deployments(gateway.upstream_url)) + with ( + refusing_forward_proxy() as forward, + owned_proxy_process( + gateway, + tmp_path, + _forward_environment(forward.url), + config=config, + remove_environment=_scrubbed(), + workers=2, + ) as owned, + ): + candidate: Final = owned.gateway + workers: Final = eventually( + lambda: _WORKER_PIDS.validate_python(_STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=60, + ) + burst: Final = asyncio.create_task( + _burst(str(candidate.client.base_url), candidate.key, 20, tolerate_transport_errors=True) + ) + await asyncio.to_thread(eventually, lambda: forward.received.qsize(), lambda size: size >= 4, 60) + victim: Final = psutil.Process(workers[0]) + victim.suspend() + victim.send_signal(signal.SIGKILL) + served: Final = await burst + for probe in served: + assert probe.response.status_code == REFUSED_STATUS, probe.response.text + assert set(forward.targets()) <= {MONGODB_HOST, VOYAGE_HOST} + follow_up: Final = _rerank(candidate, MONGODB_RERANK) + _assert_dialed(forward, follow_up, MONGODB_HOST) + duplicates: Final = tuple(call_id for call_id in _served_call_ids(served) if len(_spend_rows(call_id)) > 1) + assert duplicates == (), duplicates diff --git a/tests/integration/routing/test_fallback_hop_foreign_encrypted_reasoning.py b/tests/integration/routing/test_fallback_hop_foreign_encrypted_reasoning.py new file mode 100644 index 00000000000..bea36ee6bf8 --- /dev/null +++ b/tests/integration/routing/test_fallback_hop_foreign_encrypted_reasoning.py @@ -0,0 +1,1299 @@ +"""Fallback hops to a deployment that cannot decrypt the previous deployment's encrypted reasoning. + +A model group lists an order-1 and an order-2 deployment on distinct encryption boundaries (a +distinct ``api_base`` each, both played by the scripted upstream). The integration proxy keeps +``disable_cooldowns: true`` and ``num_retries: 0``, so every request tries order 1 first and a +failure there hops to order 2 through the router's order-based fallback. The hop must strip the +``encrypted_content`` order 2 cannot decrypt from Responses ``input`` items (summary kept) and the +bridge-tagged thinking blocks from ``messages``, the encrypted-content affinity pin must yield to +the hop's ``_target_order``, and a client cannot forge hop state (``fallback_depth``, +``_target_order``, ``attempted_targets``) through the request body while ``max_fallbacks`` stays +client-settable. Every request carries a unique marker so the response cache never serves it, and +a hop is proven by the upstream's own record: the order-1 attempt with the full history, then the +order-2 request. Every call goes through a lane: one keep-alive connection pinned to a single +worker, used only once ``/model/info`` on that same connection lists every deployment the call +needs, because the peer worker learns a ``/model/new`` row through the config-sync resync one to +sixteen seconds later, and a resync that lands between a group's two rows leaves that worker +serving the group with one deployment until the next resync. +""" + +import asyncio +import json +import re +import signal +import time +import uuid +from collections.abc import AsyncIterator, Callable, Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack, asynccontextmanager, contextmanager +from dataclasses import dataclass +from pathlib import Path +from types import MappingProxyType +from typing import Final + +import anthropic +import httpx +import openai +import psutil +import pytest +from integration._support.client import ( + JSON_OBJECT, + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, + string_value, +) +from integration._support.database import read_rows +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request, wire_server +from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, SseResponse, StoredResponse +from pydantic import JsonValue + +ORDER_ONE_BLOB: Final = "gAAAAA-minted-by-order-1" +ORDER_TWO_BLOB: Final = "gAAAAA-minted-by-order-2" +SUMMARY_TEXT: Final = "multiply 17 by 23" +TAGGED_SIGNATURE: Final = f"litellm_encrypted_reasoning:{ORDER_ONE_BLOB}" +PROVIDER_KEY: Final = "integration-provider-key" +RESPONSES_MODEL: Final = "openai/gpt-5" +BRIDGED_MODEL: Final = "openai/gpt-5-codex" +AFFINITY_CHECK: Final = "encrypted_content_affinity" +HTTPX: Final = "httpx" +OPENAI_SYNC: Final = "openai-sync" +OPENAI_ASYNC: Final = "openai-async" +ANTHROPIC_SYNC: Final = "anthropic-sync" +ANTHROPIC_ASYNC: Final = "anthropic-async" +RESPONSES_CLIENTS: Final = (HTTPX, OPENAI_SYNC, OPENAI_ASYNC) +MESSAGES_CLIENTS: Final = (HTTPX, ANTHROPIC_SYNC, ANTHROPIC_ASYNC) + + +def _failure(message: str) -> dict[str, JsonValue]: + return {"error": {"message": message, "type": "server_error", "code": "server_error", "param": None}} + + +def _responses_body(blob: str) -> dict[str, JsonValue]: + return { + "id": "resp_$UNIQUE_ID", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-5-scripted", + "output": [ + { + "type": "reasoning", + "id": "rs_$UNIQUE_ID", + "summary": [{"type": "summary_text", "text": "nineteen times twenty-one"}], + "encrypted_content": blob, + }, + { + "type": "message", + "id": "msg_$UNIQUE_ID", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "399", "annotations": []}], + }, + ], + "usage": {"input_tokens": 5, "output_tokens": 7, "total_tokens": 12}, + } + + +CHAT_BODY: Final[dict[str, JsonValue]] = { + "id": "chatcmpl-$UNIQUE_ID", + "object": "chat.completion", + "created": 1, + "model": "gpt-5-scripted", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "399"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 5, "completion_tokens": 7, "total_tokens": 12}, +} + + +def failing(message: str = "order one is down") -> StoredResponse: + return JsonResponse(content_type="application/json", status=500, body=_failure(message)) + + +def healthy_json(blob: str = ORDER_TWO_BLOB) -> StoredResponse: + return RoutedResponse( + content_type="application/x-routed", + routes={ + "POST /responses": JsonResponse(content_type="application/json", body=_responses_body(blob)), + "POST /chat/completions": JsonResponse(content_type="application/json", body=CHAT_BODY), + }, + ) + + +def _responses_frame(event: Mapping[str, JsonValue]) -> str: + return f"event: {event['type']}\ndata: {json.dumps(event)}" + + +def broken_responses_stream() -> StoredResponse: + opened: Final = {**_responses_body(ORDER_ONE_BLOB), "status": "in_progress", "output": [], "usage": None} + return SseResponse( + content_type="text/event-stream", + frames=( + _responses_frame({"type": "response.created", "sequence_number": 0, "response": opened}), + _responses_frame({"type": "response.in_progress", "sequence_number": 1, "response": opened}), + _responses_frame({"type": "error", "sequence_number": 2, **_failure("order one broke mid-stream")}), + ), + ) + + +def healthy_responses_stream(blob: str = ORDER_TWO_BLOB) -> StoredResponse: + body: Final = _responses_body(blob) + opened: Final = {**body, "status": "in_progress", "output": [], "usage": None} + output: Final = body["output"] + assert isinstance(output, list) + reasoning, message = output + assert isinstance(reasoning, dict) and isinstance(message, dict) + part: Final = {"type": "output_text", "text": "", "annotations": []} + events: Final = ( + {"type": "response.created", "response": opened}, + {"type": "response.in_progress", "response": opened}, + {"type": "response.output_item.added", "output_index": 0, "item": {**reasoning, "encrypted_content": None}}, + {"type": "response.output_item.done", "output_index": 0, "item": reasoning}, + { + "type": "response.output_item.added", + "output_index": 1, + "item": {**message, "status": "in_progress", "content": []}, + }, + { + "type": "response.content_part.added", + "output_index": 1, + "content_index": 0, + "item_id": "msg_$UNIQUE_ID", + "part": part, + }, + { + "type": "response.output_text.delta", + "output_index": 1, + "content_index": 0, + "item_id": "msg_$UNIQUE_ID", + "delta": "399", + }, + { + "type": "response.output_text.done", + "output_index": 1, + "content_index": 0, + "item_id": "msg_$UNIQUE_ID", + "text": "399", + }, + { + "type": "response.content_part.done", + "output_index": 1, + "content_index": 0, + "item_id": "msg_$UNIQUE_ID", + "part": {**part, "text": "399"}, + }, + {"type": "response.output_item.done", "output_index": 1, "item": message}, + {"type": "response.completed", "response": body}, + ) + return SseResponse( + content_type="text/event-stream", + frames=tuple(_responses_frame({**event, "sequence_number": number}) for number, event in enumerate(events)), + ) + + +def _chat_chunk(delta: Mapping[str, JsonValue], finish_reason: str | None) -> str: + chunk: Final = { + "id": "chatcmpl-$UNIQUE_ID", + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-5-scripted", + "choices": [{"index": 0, "delta": dict(delta), "finish_reason": finish_reason}], + } + return f"data: {json.dumps(chunk)}" + + +def healthy_chat_stream() -> StoredResponse: + return SseResponse( + content_type="text/event-stream", + frames=(_chat_chunk({"role": "assistant", "content": "399"}, None), _chat_chunk({}, "stop"), "data: [DONE]"), + ) + + +@dataclass(frozen=True, slots=True) +class Target: + name: str + deployments: frozenset[str] + + +@dataclass(frozen=True, slots=True) +class OrderedGroup: + name: str + order_one: str + order_two: str + order_one_scenario: str + order_two_scenario: str + + @property + def target(self) -> Target: + return Target(self.name, frozenset({self.order_one, self.order_two})) + + +@dataclass(frozen=True, slots=True) +class Lane(Gateway): + pass + + +LANE_LIMITS: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1, keepalive_expiry=30) +CONVERGENCE_SECONDS: Final = 60 + + +def deployment_ids(entries: JsonValue) -> frozenset[str]: + assert isinstance(entries, list), entries + return frozenset(string_value(object_value(object_value(entry)["model_info"])["id"]) for entry in entries) + + +def knows(lane: Gateway, deployments: frozenset[str]) -> bool: + return deployments <= deployment_ids(lane.get("/model/info")["data"]) + + +@contextmanager +def pinned_client(gateway: Gateway) -> Iterator[Lane]: + with httpx.Client(base_url=base_url(gateway), timeout=15, trust_env=False, limits=LANE_LIMITS) as client: + yield Lane(client, gateway.key, gateway.upstream_url) + + +@contextmanager +def lane_for(gateway: Gateway, target: Target) -> Iterator[Gateway]: + if isinstance(gateway, Lane): + yield gateway + return + with pinned_client(gateway) as lane: + eventually(lambda: knows(lane, target.deployments), lambda converged: converged, seconds=CONVERGENCE_SECONDS) + yield lane + + +@contextmanager +def lanes(gateway: Gateway, targets: Sequence[Target], count: int) -> Iterator[tuple[Lane, ...]]: + wanted: Final = frozenset[str]().union(*(target.deployments for target in targets)) + with ExitStack() as stack: + opened: Final = tuple(stack.enter_context(pinned_client(gateway)) for _ in range(count)) + with ThreadPoolExecutor(max_workers=count) as pool: + _ = tuple(pool.map(lambda lane: knows(lane, wanted), opened)) + _ = eventually(lambda: tuple(knows(lane, wanted) for lane in opened), all, seconds=CONVERGENCE_SECONDS) + yield opened + + +async def async_knows(transport: httpx.AsyncClient, key: str, deployments: frozenset[str]) -> bool: + response: Final = await transport.get("/model/info", headers={"Authorization": f"Bearer {key}"}) + assert response.status_code == 200, response.text + return deployments <= deployment_ids(JSON_OBJECT.validate_json(response.content)["data"]) + + +@asynccontextmanager +async def async_lane(gateway: Gateway, target: Target) -> AsyncIterator[httpx.AsyncClient]: + async with httpx.AsyncClient( + base_url=base_url(gateway), timeout=15, trust_env=False, limits=LANE_LIMITS + ) as transport: + deadline: Final = time.monotonic() + CONVERGENCE_SECONDS + while not await async_knows(transport, gateway.key, target.deployments): + assert time.monotonic() < deadline, f"{target} never reached this worker" + await asyncio.sleep(0.1) + yield transport + + +def scenario_api_base(scenario: Scenario, scenario_id: str, response: StoredResponse) -> str: + handle: Final = register_scenario(scenario_id, response) + scenario.cleanups.callback(delete_scenario, handle) + return handle.api_base() + + +def deployment( + scenario: Scenario, + group: str, + api_base: str, + *, + order: int | None, + model: str = RESPONSES_MODEL, + extra_params: Mapping[str, JsonValue] = MappingProxyType({}), +) -> str: + created: Final = scenario.gateway.post( + "/model/new", + { + "model_name": group, + "litellm_params": { + "model": model, + "api_key": PROVIDER_KEY, + "api_base": api_base, + **({} if order is None else {"order": order}), + **extra_params, + }, + "model_info": {}, + }, + ) + model_id: Final = string_value(object_value(created["model_info"])["id"]) + scenario.cleanups.callback(scenario.delete_model, model_id) + return model_id + + +def ordered_group( + scenario: Scenario, + *, + order_one: StoredResponse, + order_two: StoredResponse, + model: str = RESPONSES_MODEL, +) -> OrderedGroup: + name: Final = f"hop-{uuid.uuid4().hex[:12]}" + one: Final = f"{name}-o1" + two: Final = f"{name}-o2" + one_base: Final = scenario_api_base(scenario, one, order_one) + two_base: Final = scenario_api_base(scenario, two, order_two) + return OrderedGroup( + name, + deployment(scenario, name, one_base, order=1, model=model), + deployment(scenario, name, two_base, order=2, model=model), + one, + two, + ) + + +def observed(gateway: Gateway) -> tuple[tuple[str, dict[str, JsonValue]], ...]: + with httpx.Client(base_url=gateway.upstream_url, trust_env=False, timeout=15) as upstream: + payload: Final = object_value(upstream.get("/__observations").json()) + requests: Final = payload.get("requests") + assert isinstance(requests, list), payload + return tuple( + (string_value(object_value(request)["path"]), object_value(object_value(request)["body"])) + for request in requests + ) + + +def bodies_for( + records: Sequence[tuple[str, dict[str, JsonValue]]], scenario_id: str +) -> tuple[dict[str, JsonValue], ...]: + return tuple(body for path, body in records if path.startswith(f"/{scenario_id}/")) + + +@dataclass(frozen=True, slots=True) +class HopRecord: + order_one: tuple[dict[str, JsonValue], ...] + order_two: tuple[dict[str, JsonValue], ...] + + +def hop_record(gateway: Gateway, group: OrderedGroup) -> HopRecord: + records: Final = observed(gateway) + return HopRecord(bodies_for(records, group.order_one_scenario), bodies_for(records, group.order_two_scenario)) + + +def user_item(text: str) -> dict[str, JsonValue]: + return {"type": "message", "role": "user", "content": text} + + +def reasoning_item(encrypted_content: JsonValue, *, summary: bool = True) -> dict[str, JsonValue]: + return { + "type": "reasoning", + "id": "rs_order1", + "encrypted_content": encrypted_content, + **({"summary": [{"type": "summary_text", "text": SUMMARY_TEXT}]} if summary else {}), + } + + +ASSISTANT_ITEM: Final[dict[str, JsonValue]] = { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "391"}], +} +STRIPPED_REASONING: Final[dict[str, JsonValue]] = { + "type": "reasoning", + "summary": [{"type": "summary_text", "text": SUMMARY_TEXT}], +} + + +def history(marker: str, *reasoning: dict[str, JsonValue]) -> list[JsonValue]: + replayed: Final = reasoning or (reasoning_item(ORDER_ONE_BLOB),) + return [user_item(f"What is 17*23? {marker}"), *replayed, ASSISTANT_ITEM, user_item("And 19*21?")] + + +def stripped_history(marker: str, *reasoning: dict[str, JsonValue]) -> list[JsonValue]: + replayed: Final = reasoning or (STRIPPED_REASONING,) + return [user_item(f"What is 17*23? {marker}"), *replayed, ASSISTANT_ITEM, user_item("And 19*21?")] + + +def chat_messages(marker: str) -> list[JsonValue]: + return [ + {"role": "user", "content": f"What is 17*23? {marker}"}, + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": SUMMARY_TEXT, "signature": TAGGED_SIGNATURE}, + {"type": "text", "text": "391"}, + ], + }, + {"role": "user", "content": "And 19*21?"}, + ] + + +def stripped_chat_messages(marker: str) -> list[JsonValue]: + return [ + {"role": "user", "content": f"What is 17*23? {marker}"}, + {"role": "assistant", "content": [{"type": "text", "text": "391"}]}, + {"role": "user", "content": "And 19*21?"}, + ] + + +def marker() -> str: + return uuid.uuid4().hex + + +def base_url(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def data_frames(text: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + object_value(json.loads(line.removeprefix("data: "))) + for line in text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + + +@dataclass(frozen=True, slots=True) +class Answer: + status: int + response_id: str | None + text: str + headers: Mapping[str, str] + + +def _httpx_answer(response: httpx.Response, response_id: Callable[[httpx.Response], str | None]) -> Answer: + return Answer( + response.status_code, + response_id(response) if response.status_code == 200 else None, + response.text, + MappingProxyType(dict(response.headers)), + ) + + +def _json_id(response: httpx.Response) -> str | None: + identity: Final = object_value(response.json()).get("id") + return identity if isinstance(identity, str) else None + + +def _completed_stream_id(response: httpx.Response) -> str | None: + frames: Final = data_frames(response.text) + assert frames and frames[-1].get("type") == "response.completed", response.text + identity: Final = object_value(frames[-1]["response"]).get("id") + return identity if isinstance(identity, str) else None + + +def responses_answer( + gateway: Gateway, + client: str, + target: Target, + request_input: JsonValue, + *, + stream: bool = False, + extra: Mapping[str, JsonValue] = MappingProxyType({}), +) -> Answer: + body: Final = {"model": target.name, "input": request_input, "store": False, **extra} + if client == HTTPX: + with lane_for(gateway, target) as lane: + response: Final = lane.request("POST", "/v1/responses", {**body, "stream": stream}) + return _httpx_answer(response, _completed_stream_id if stream else _json_id) + if client == OPENAI_SYNC: + with lane_for(gateway, target) as lane: + sdk: Final = openai.OpenAI( + base_url=f"{base_url(gateway)}/v1", api_key=gateway.key, max_retries=0, http_client=lane.client + ) + try: + if stream: + streamed: Final = sdk.responses.with_raw_response.create(**body, stream=True) + events: Final = list(streamed.parse()) + assert events and events[-1].type == "response.completed", events + return Answer( + streamed.status_code, + events[-1].response.id, + json.dumps([event.type for event in events]), + MappingProxyType(dict(streamed.headers)), + ) + raw: Final = sdk.responses.with_raw_response.create(**body) + return Answer(raw.status_code, raw.parse().id, raw.text, MappingProxyType(dict(raw.headers))) + except openai.APIStatusError as error: + return Answer( + error.status_code, None, error.response.text, MappingProxyType(dict(error.response.headers)) + ) + assert client == OPENAI_ASYNC, client + + async def call() -> Answer: + async with async_lane(gateway, target) as transport: + sdk: Final = openai.AsyncOpenAI( + base_url=f"{base_url(gateway)}/v1", api_key=gateway.key, max_retries=0, http_client=transport + ) + try: + if stream: + streamed: Final = await sdk.responses.with_raw_response.create(**body, stream=True) + events: Final = [event async for event in streamed.parse()] + assert events and events[-1].type == "response.completed", events + return Answer( + streamed.status_code, + events[-1].response.id, + json.dumps([event.type for event in events]), + MappingProxyType(dict(streamed.headers)), + ) + raw: Final = await sdk.responses.with_raw_response.create(**body) + return Answer(raw.status_code, raw.parse().id, raw.text, MappingProxyType(dict(raw.headers))) + except openai.APIStatusError as error: + return Answer( + error.status_code, None, error.response.text, MappingProxyType(dict(error.response.headers)) + ) + + return asyncio.run(call()) + + +def chat_answer( + gateway: Gateway, client: str, target: Target, messages: list[JsonValue], *, stream: bool = False +) -> Answer: + body: Final = {"model": target.name, "messages": messages} + if client == HTTPX: + with lane_for(gateway, target) as lane: + response: Final = lane.request("POST", "/v1/chat/completions", {**body, "stream": stream}) + if stream: + return _httpx_answer(response, lambda served: string_value(data_frames(served.text)[-1]["id"])) + return _httpx_answer(response, _json_id) + if client == OPENAI_SYNC: + with lane_for(gateway, target) as lane: + sdk: Final = openai.OpenAI( + base_url=f"{base_url(gateway)}/v1", api_key=gateway.key, max_retries=0, http_client=lane.client + ) + try: + raw: Final = sdk.chat.completions.with_raw_response.create(**body) + return Answer(raw.status_code, raw.parse().id, raw.text, MappingProxyType(dict(raw.headers))) + except openai.APIStatusError as error: + return Answer( + error.status_code, None, error.response.text, MappingProxyType(dict(error.response.headers)) + ) + assert client == OPENAI_ASYNC, client + + async def call() -> Answer: + async with async_lane(gateway, target) as transport: + sdk: Final = openai.AsyncOpenAI( + base_url=f"{base_url(gateway)}/v1", api_key=gateway.key, max_retries=0, http_client=transport + ) + try: + raw: Final = await sdk.chat.completions.with_raw_response.create(**body) + return Answer(raw.status_code, raw.parse().id, raw.text, MappingProxyType(dict(raw.headers))) + except openai.APIStatusError as error: + return Answer( + error.status_code, None, error.response.text, MappingProxyType(dict(error.response.headers)) + ) + + return asyncio.run(call()) + + +def messages_answer( + gateway: Gateway, client: str, target: Target, messages: list[JsonValue], *, stream: bool = False +) -> Answer: + body: Final = {"model": target.name, "max_tokens": 64, "messages": messages} + if client == HTTPX: + with lane_for(gateway, target) as lane: + response: Final = lane.request("POST", "/v1/messages", {**body, "stream": stream}) + if stream: + return _httpx_answer( + response, lambda served: string_value(object_value(data_frames(served.text)[0]["message"])["id"]) + ) + return _httpx_answer(response, _json_id) + if client == ANTHROPIC_SYNC: + with lane_for(gateway, target) as lane: + sdk: Final = anthropic.Anthropic( + base_url=base_url(gateway), api_key=gateway.key, max_retries=0, http_client=lane.client + ) + try: + raw: Final = sdk.messages.with_raw_response.create(**body) + return Answer(raw.status_code, raw.parse().id, raw.text, MappingProxyType(dict(raw.headers))) + except anthropic.APIStatusError as error: + return Answer( + error.status_code, None, error.response.text, MappingProxyType(dict(error.response.headers)) + ) + assert client == ANTHROPIC_ASYNC, client + + async def call() -> Answer: + async with async_lane(gateway, target) as transport: + sdk: Final = anthropic.AsyncAnthropic( + base_url=base_url(gateway), api_key=gateway.key, max_retries=0, http_client=transport + ) + try: + raw: Final = await sdk.messages.with_raw_response.create(**body) + return Answer(raw.status_code, raw.parse().id, raw.text, MappingProxyType(dict(raw.headers))) + except anthropic.APIStatusError as error: + return Answer( + error.status_code, None, error.response.text, MappingProxyType(dict(error.response.headers)) + ) + + return asyncio.run(call()) + + +def assert_served_by(answer: Answer, model_id: str) -> None: + assert answer.status == 200, answer.text + assert answer.response_id is not None, answer.text + assert answer.headers.get("x-litellm-model-id") == model_id, (model_id, answer.headers, answer.text) + + +def assert_served_by_order_two(answer: Answer, group: OrderedGroup) -> None: + assert_served_by(answer, group.order_two) + + +def assert_hop_stripped( + record: HopRecord, expected_order_one: list[JsonValue], expected_order_two: list[JsonValue] +) -> None: + assert len(record.order_one) == 1, record + assert record.order_one[0].get("input") == expected_order_one, record.order_one[0] + assert len(record.order_two) == 1, record + assert record.order_two[0].get("input") == expected_order_two, record.order_two[0] + + +def success_rows(call_ids: Sequence[str]) -> list[dict[str, JsonValue]]: + placeholders: Final = ", ".join("%s" for _ in call_ids) + rows: Final = eventually( + lambda: read_rows( + f'SELECT litellm_call_id, status, model_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id IN ({placeholders})', + tuple(call_ids), + ), + lambda found: len({row["litellm_call_id"] for row in found if row["status"] == "success"}) == len(call_ids), + seconds=90, + ) + return [row for row in rows if row["status"] == "success"] + + +def spend_rows(call_id: str) -> list[dict[str, JsonValue]]: + return eventually( + lambda: read_rows( + 'SELECT litellm_call_id, status, model_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id = %s', + (call_id,), + ), + lambda rows: any(row["status"] == "success" for row in rows), + seconds=70, + ) + + +def assert_logged_once_on_order_two(call_id: str, group: OrderedGroup) -> None: + successes: Final = [row for row in spend_rows(call_id) if row["status"] == "success"] + assert [row["model_id"] for row in successes] == [group.order_two], successes + + +def set_pre_call_checks(gateway: Gateway, checks: Sequence[str]) -> None: + gateway.post("/config/update", {"router_settings": {"optional_pre_call_checks": list(checks)}}) + assert router_setting(gateway, "optional_pre_call_checks") == list(checks), router_setting( + gateway, "optional_pre_call_checks" + ) + + +@pytest.fixture(scope="module", autouse=True) +def affinity_check_off() -> Iterator[None]: + with gateway_from_environment() as gateway: + original: Final = router_setting(gateway, "optional_pre_call_checks") + set_pre_call_checks(gateway, ()) + try: + yield + finally: + set_pre_call_checks(gateway, [str(check) for check in original] if isinstance(original, list) else ()) + + +def enable_affinity_check(scenario: Scenario) -> None: + scenario.cleanups.callback(set_pre_call_checks, scenario.gateway, ()) + set_pre_call_checks(scenario.gateway, (AFFINITY_CHECK,)) + + +def router_setting(gateway: Gateway, name: str) -> JsonValue: + return object_value(gateway.get("/router/settings")["current_values"]).get(name) + + +@pytest.mark.parametrize("client", RESPONSES_CLIENTS) +def test_responses_hop_drops_the_order_one_reasoning_order_two_cannot_decrypt(gateway: Gateway, client: str) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_json()) + mark: Final = marker() + answer: Final = responses_answer(gateway, client, group.target, history(mark)) + assert_served_by_order_two(answer, group) + assert_hop_stripped(hop_record(gateway, group), history(mark), stripped_history(mark)) + assert_logged_once_on_order_two(answer.headers["x-litellm-call-id"], group) + + +@pytest.mark.parametrize("client", RESPONSES_CLIENTS) +def test_responses_mid_stream_hop_drops_the_order_one_reasoning(gateway: Gateway, client: str) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group( + scenario, order_one=broken_responses_stream(), order_two=healthy_responses_stream() + ) + mark: Final = marker() + answer: Final = responses_answer(gateway, client, group.target, history(mark), stream=True) + assert_served_by_order_two(answer, group) + assert_hop_stripped(hop_record(gateway, group), history(mark), stripped_history(mark)) + + +def test_responses_stream_refused_before_any_frame_hops_stripped(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_responses_stream()) + mark: Final = marker() + answer: Final = responses_answer(gateway, HTTPX, group.target, history(mark), stream=True) + assert_served_by_order_two(answer, group) + assert_hop_stripped(hop_record(gateway, group), history(mark), stripped_history(mark)) + + +@pytest.mark.parametrize("client", RESPONSES_CLIENTS) +def test_chat_hop_drops_the_bridge_tagged_thinking_block(gateway: Gateway, client: str) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_json()) + mark: Final = marker() + answer: Final = chat_answer(gateway, client, group.target, chat_messages(mark)) + assert answer.status == 200, answer.text + assert answer.response_id is not None and answer.response_id.startswith( + f"chatcmpl-{group.order_two_scenario}-" + ), answer.text + record: Final = hop_record(gateway, group) + assert len(record.order_one) == 1 and record.order_one[0].get("messages") == chat_messages(mark), record + assert len(record.order_two) == 1 and record.order_two[0].get("messages") == stripped_chat_messages(mark), ( + record + ) + + +def test_chat_stream_hop_drops_the_bridge_tagged_thinking_block(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_chat_stream()) + mark: Final = marker() + answer: Final = chat_answer(gateway, HTTPX, group.target, chat_messages(mark), stream=True) + assert answer.status == 200, answer.text + assert answer.response_id is not None and answer.response_id.startswith( + f"chatcmpl-{group.order_two_scenario}-" + ), answer.text + record: Final = hop_record(gateway, group) + assert len(record.order_one) == 1 and record.order_one[0].get("messages") == chat_messages(mark), record + assert len(record.order_two) == 1 and record.order_two[0].get("messages") == stripped_chat_messages(mark), ( + record + ) + + +def assert_bridged_hop_stripped(record: HopRecord, mark: str) -> None: + assert len(record.order_one) == 1, record + assert ORDER_ONE_BLOB in json.dumps(record.order_one[0]), record.order_one[0] + assert len(record.order_two) == 1, record + order_two: Final = record.order_two[0] + assert ORDER_ONE_BLOB not in json.dumps(order_two), order_two + items: Final = order_two.get("input") + assert isinstance(items, list) and mark in json.dumps(items), order_two + assert not any(isinstance(item, dict) and "encrypted_content" in item for item in items), order_two + + +@pytest.mark.parametrize("client", MESSAGES_CLIENTS) +def test_messages_hop_drops_the_bridge_tagged_thinking_block(gateway: Gateway, client: str) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_json(), model=BRIDGED_MODEL) + mark: Final = marker() + answer: Final = messages_answer(gateway, client, group.target, chat_messages(mark)) + assert answer.status == 200, answer.text + assert "399" in answer.text, answer.text + assert_bridged_hop_stripped(hop_record(gateway, group), mark) + + +def test_messages_mid_stream_hop_drops_the_bridge_tagged_thinking_block(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group( + scenario, order_one=broken_responses_stream(), order_two=healthy_responses_stream(), model=BRIDGED_MODEL + ) + mark: Final = marker() + answer: Final = messages_answer(gateway, HTTPX, group.target, chat_messages(mark), stream=True) + assert answer.status == 200, answer.text + assert "399" in answer.text, answer.text + assert_bridged_hop_stripped(hop_record(gateway, group), mark) + + +def affinity_turn_one(gateway: Gateway, group: OrderedGroup) -> list[JsonValue]: + answer: Final = responses_answer( + gateway, HTTPX, group.target, f"hello affinity {marker()}", extra={"include": ["reasoning.encrypted_content"]} + ) + assert answer.status == 200, answer.text + assert answer.headers.get("x-litellm-model-id") == group.order_one, answer.headers + output: Final = object_value(json.loads(answer.text)).get("output") + assert isinstance(output, list), answer.text + reasoning: Final = next( + (object_value(item) for item in output if object_value(item).get("type") == "reasoning"), None + ) + assert reasoning is not None and str(reasoning["id"]).startswith("encitem_"), answer.text + message: Final = next((object_value(item) for item in output if object_value(item).get("type") == "message"), None) + assert message is not None, answer.text + return [reasoning, message] + + +def test_affinity_pin_yields_to_the_hop_and_strips_the_origins_reasoning(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + enable_affinity_check(scenario) + group: Final = ordered_group(scenario, order_one=healthy_json(ORDER_ONE_BLOB), order_two=healthy_json()) + items: Final = affinity_turn_one(gateway, group) + register_scenario(group.order_one_scenario, failing()) + mark: Final = marker() + answer: Final = responses_answer( + gateway, HTTPX, group.target, [user_item(f"hello affinity {mark}"), *items, user_item("continue")] + ) + assert_served_by_order_two(answer, group) + record: Final = hop_record(gateway, group) + assert len(record.order_one) == 2 and ORDER_ONE_BLOB in json.dumps(record.order_one[1]), record.order_one + assert len(record.order_two) == 1, record + assert ORDER_ONE_BLOB not in json.dumps(record.order_two[0]), record.order_two[0] + assert record.order_two[0].get("input") == [ + user_item(f"hello affinity {mark}"), + {"type": "reasoning", "summary": [{"type": "summary_text", "text": "nineteen times twenty-one"}]}, + items[1], + user_item("continue"), + ], record.order_two[0] + assert router_setting(gateway, "optional_pre_call_checks") == [], router_setting( + gateway, "optional_pre_call_checks" + ) + + +def test_hop_inside_one_encryption_boundary_keeps_the_reasoning(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + shared: Final = f"shared-{uuid.uuid4().hex[:12]}" + one: Final = scenario_api_base(scenario, f"{shared}-o1", failing()).rsplit("/", 1)[-1] + two: Final = scenario_api_base(scenario, f"{shared}-o2", healthy_json()).rsplit("/", 1)[-1] + shared_base: Final = f"{gateway.upstream_url}/{shared}" + order_one: Final = deployment( + scenario, shared, shared_base, order=1, extra_params={"extra_headers": {"x-scripted-scenario": one}} + ) + order_two: Final = deployment( + scenario, shared, shared_base, order=2, extra_params={"extra_headers": {"x-scripted-scenario": two}} + ) + mark: Final = marker() + answer: Final = responses_answer( + gateway, HTTPX, Target(shared, frozenset({order_one, order_two})), history(mark) + ) + assert answer.status == 200, answer.text + assert answer.headers.get("x-litellm-model-id") == order_two, (order_one, answer.headers) + bodies: Final = bodies_for(observed(gateway), shared) + assert [body.get("input") for body in bodies] == [history(mark), history(mark)], bodies + + +def configured_fallback(scenario: Scenario, source: str, target: str) -> None: + gateway: Final = scenario.gateway + original: Final = router_setting(gateway, "fallbacks") + restored: Final = original if isinstance(original, list) else [] + scenario.cleanups.callback(lambda: gateway.post("/config/update", {"router_settings": {"fallbacks": restored}})) + gateway.post("/config/update", {"router_settings": {"fallbacks": [*restored, {source: [target]}]}}) + + +def assert_cross_group_hop_stripped(gateway: Gateway, source_scenario: str, target_scenario: str, mark: str) -> None: + records: Final = observed(gateway) + assert [body.get("input") for body in bodies_for(records, source_scenario)] == [history(mark)], records + assert [body.get("input") for body in bodies_for(records, target_scenario)] == [stripped_history(mark)], records + + +def test_configured_fallbacks_entry_hop_strips_the_source_groups_reasoning(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + source: Final = f"src-{uuid.uuid4().hex[:12]}" + target: Final = f"dst-{uuid.uuid4().hex[:12]}" + source_id: Final = deployment(scenario, source, scenario_api_base(scenario, source, failing()), order=None) + target_id: Final = deployment(scenario, target, scenario_api_base(scenario, target, healthy_json()), order=None) + configured_fallback(scenario, source, target) + mark: Final = marker() + answer: Final = responses_answer( + gateway, HTTPX, Target(source, frozenset({source_id, target_id})), history(mark) + ) + assert_served_by(answer, target_id) + assert_cross_group_hop_stripped(gateway, source, target, mark) + assert source not in json.dumps(router_setting(gateway, "fallbacks")), router_setting(gateway, "fallbacks") + + +def test_client_fallbacks_entry_hop_strips_the_source_groups_reasoning(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + source: Final = f"src-{uuid.uuid4().hex[:12]}" + target: Final = f"dst-{uuid.uuid4().hex[:12]}" + source_id: Final = deployment(scenario, source, scenario_api_base(scenario, source, failing()), order=None) + target_id: Final = deployment(scenario, target, scenario_api_base(scenario, target, healthy_json()), order=None) + mark: Final = marker() + answer: Final = responses_answer( + gateway, + HTTPX, + Target(source, frozenset({source_id, target_id})), + history(mark), + extra={"fallbacks": [target]}, + ) + assert_served_by(answer, target_id) + assert_cross_group_hop_stripped(gateway, source, target, mark) + + +def assert_served_by_order_one_intact(gateway: Gateway, group: OrderedGroup, answer: Answer, mark: str) -> None: + assert answer.status == 200, answer.text + assert answer.headers.get("x-litellm-model-id") == group.order_one, answer.headers + record: Final = hop_record(gateway, group) + assert [body.get("input") for body in record.order_one] == [history(mark)], record + assert record.order_two == (), record + + +@pytest.mark.parametrize( + "target_order", + [ + pytest.param(2, id="int"), + pytest.param("2", id="str"), + pytest.param([2], id="list"), + pytest.param(True, id="bool"), + pytest.param(99, id="unknown-order"), + pytest.param(None, id="null"), + ], +) +def test_client_sent_target_order_never_moves_a_request_off_order_one( + gateway: Gateway, target_order: JsonValue +) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=healthy_json(ORDER_ONE_BLOB), order_two=healthy_json()) + mark: Final = marker() + answer: Final = responses_answer( + gateway, HTTPX, group.target, history(mark), extra={"_target_order": target_order} + ) + assert_served_by_order_one_intact(gateway, group, answer, mark) + + +@pytest.mark.parametrize( + "fallback_depth", + [ + pytest.param(1, id="int"), + pytest.param(True, id="bool"), + pytest.param("1", id="str"), + pytest.param([1], id="list"), + ], +) +def test_client_sent_fallback_depth_never_strips_a_plain_request(gateway: Gateway, fallback_depth: JsonValue) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=healthy_json(ORDER_ONE_BLOB), order_two=healthy_json()) + mark: Final = marker() + answer: Final = responses_answer( + gateway, HTTPX, group.target, history(mark), extra={"fallback_depth": fallback_depth} + ) + assert_served_by_order_one_intact(gateway, group, answer, mark) + + +def test_client_sent_fallback_depth_cannot_exhaust_the_fallback_budget(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_json()) + mark: Final = marker() + answer: Final = responses_answer(gateway, HTTPX, group.target, history(mark), extra={"fallback_depth": 5}) + assert_served_by_order_two(answer, group) + assert_hop_stripped(hop_record(gateway, group), history(mark), stripped_history(mark)) + + +def test_client_sent_attempted_targets_do_not_reach_the_hop(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_json()) + mark: Final = marker() + answer: Final = responses_answer( + gateway, HTTPX, group.target, history(mark), extra={"attempted_targets": ["forged"]} + ) + assert_served_by_order_two(answer, group) + assert_hop_stripped(hop_record(gateway, group), history(mark), stripped_history(mark)) + + +def test_client_sent_max_fallbacks_zero_still_stops_the_hop(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_json()) + mark: Final = marker() + answer: Final = responses_answer(gateway, HTTPX, group.target, history(mark), extra={"max_fallbacks": 0}) + assert answer.status == 500, answer.text + assert "order one is down" in answer.text, answer.text + record: Final = hop_record(gateway, group) + assert [body.get("input") for body in record.order_one] == [history(mark)], record + assert record.order_two == (), record + + +FIVE_KB: Final = "x" * 5120 + + +@pytest.mark.parametrize( + ("replayed", "expected"), + [ + pytest.param((reasoning_item(7),), (STRIPPED_REASONING,), id="int"), + pytest.param((reasoning_item(["a", "b"]),), (STRIPPED_REASONING,), id="list"), + pytest.param((reasoning_item(FIVE_KB),), (STRIPPED_REASONING,), id="5kb"), + pytest.param( + (reasoning_item(ORDER_ONE_BLOB), reasoning_item(ORDER_ONE_BLOB)), + (STRIPPED_REASONING, STRIPPED_REASONING), + id="duplicated", + ), + pytest.param((reasoning_item(""),), (reasoning_item(""),), id="empty"), + ], +) +def test_hop_strips_malformed_encrypted_content_without_failing( + gateway: Gateway, replayed: tuple[dict[str, JsonValue], ...], expected: tuple[dict[str, JsonValue], ...] +) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_json()) + mark: Final = marker() + answer: Final = responses_answer(gateway, HTTPX, group.target, history(mark, *replayed)) + assert_served_by_order_two(answer, group) + assert_hop_stripped(hop_record(gateway, group), history(mark, *replayed), stripped_history(mark, *expected)) + + +def test_hop_drops_a_reasoning_item_with_nothing_readable_whole(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_json()) + mark: Final = marker() + unreadable: Final = reasoning_item(ORDER_ONE_BLOB, summary=False) + answer: Final = responses_answer(gateway, HTTPX, group.target, history(mark, unreadable)) + assert_served_by_order_two(answer, group) + assert_hop_stripped( + hop_record(gateway, group), + history(mark, unreadable), + [user_item(f"What is 17*23? {mark}"), ASSISTANT_ITEM, user_item("And 19*21?")], + ) + + +def test_every_order_failing_reports_the_failure_to_the_caller(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group( + scenario, order_one=failing("order one is down"), order_two=failing("order two is down") + ) + mark: Final = marker() + answer: Final = responses_answer(gateway, HTTPX, group.target, history(mark)) + assert answer.status == 500, answer.text + assert "is down" in answer.text, answer.text + record: Final = hop_record(gateway, group) + assert [body.get("input") for body in record.order_one] == [history(mark)], record + assert len(record.order_two) == 1, record + + +def test_hop_strips_with_the_affinity_check_list_explicitly_empty(gateway: Gateway) -> None: + assert router_setting(gateway, "optional_pre_call_checks") == [], router_setting( + gateway, "optional_pre_call_checks" + ) + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_json()) + mark: Final = marker() + answer: Final = responses_answer(gateway, HTTPX, group.target, history(mark)) + assert_served_by_order_two(answer, group) + assert_hop_stripped(hop_record(gateway, group), history(mark), stripped_history(mark)) + + +def test_next_turn_without_a_hop_forwards_the_reasoning_untouched(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + group: Final = ordered_group(scenario, order_one=healthy_json(ORDER_ONE_BLOB), order_two=healthy_json()) + mark: Final = marker() + replayed: Final = reasoning_item(ORDER_TWO_BLOB) + answer: Final = responses_answer(gateway, HTTPX, group.target, history(mark, replayed)) + assert answer.status == 200, answer.text + assert answer.headers.get("x-litellm-model-id") == group.order_one, answer.headers + record: Final = hop_record(gateway, group) + assert [body.get("input") for body in record.order_one] == [history(mark, replayed)], record + assert record.order_two == (), record + + +def test_affinity_pin_keeps_the_targets_own_marked_reasoning_and_the_unmarked_history(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + enable_affinity_check(scenario) + group: Final = ordered_group(scenario, order_one=failing(), order_two=healthy_json()) + first: Final = responses_answer( + gateway, + HTTPX, + group.target, + f"hello affinity {marker()}", + extra={"include": ["reasoning.encrypted_content"]}, + ) + assert_served_by_order_two(first, group) + output: Final = object_value(json.loads(first.text)).get("output") + assert isinstance(output, list), first.text + marked: Final = next( + (object_value(item) for item in output if object_value(item).get("type") == "reasoning"), None + ) + assert marked is not None and str(marked["id"]).startswith("encitem_"), first.text + mark: Final = marker() + unmarked: Final = reasoning_item(ORDER_ONE_BLOB) + answer: Final = responses_answer( + gateway, + HTTPX, + group.target, + [user_item(f"hello affinity {mark}"), marked, unmarked, ASSISTANT_ITEM, user_item("continue")], + ) + assert_served_by_order_two(answer, group) + record: Final = hop_record(gateway, group) + assert len(record.order_one) == 1 and mark not in json.dumps(record.order_one), record + assert len(record.order_two) == 2, record + pinned: Final = record.order_two[-1].get("input") + assert isinstance(pinned, list) and mark in json.dumps(pinned[0]), record + own: Final = object_value(pinned[1]) + assert own.get("type") == "reasoning" and own.get("summary") == marked["summary"], pinned + assert str(own.get("id")).startswith(f"rs_{group.order_two_scenario}-"), pinned + assert own.get("encrypted_content") == ORDER_TWO_BLOB, pinned + assert pinned[2] == unmarked, pinned + assert pinned[3:] == [ASSISTANT_ITEM, user_item("continue")], pinned + + +RESPONSES_JSON: Final = "responses" +RESPONSES_STREAM: Final = "responses-stream" +CHAT_JSON: Final = "chat" +CHAT_STREAM: Final = "chat-stream" +MESSAGES_JSON: Final = "messages" +MESSAGES_STREAM: Final = "messages-stream" +BURST_KINDS: Final = (RESPONSES_JSON, CHAT_JSON, MESSAGES_JSON, RESPONSES_STREAM, CHAT_STREAM, MESSAGES_STREAM) + + +def burst_groups(scenario: Scenario) -> Mapping[str, OrderedGroup]: + return MappingProxyType( + { + RESPONSES_JSON: ordered_group(scenario, order_one=failing(), order_two=healthy_json()), + RESPONSES_STREAM: ordered_group(scenario, order_one=failing(), order_two=healthy_responses_stream()), + CHAT_JSON: ordered_group(scenario, order_one=failing(), order_two=healthy_json()), + CHAT_STREAM: ordered_group(scenario, order_one=failing(), order_two=healthy_chat_stream()), + MESSAGES_JSON: ordered_group(scenario, order_one=failing(), order_two=healthy_json(), model=BRIDGED_MODEL), + MESSAGES_STREAM: ordered_group( + scenario, order_one=failing(), order_two=healthy_responses_stream(), model=BRIDGED_MODEL + ), + } + ) + + +@dataclass(frozen=True, slots=True) +class Fired: + kind: str + mark: str + status: int | None + call_id: str | None + text: str + + +def fire(gateway: Gateway, groups: Mapping[str, OrderedGroup], kind: str, mark: str) -> Fired: + group: Final = groups[kind].target + stream: Final = kind.endswith("-stream") + try: + if kind.startswith("responses"): + answer: Final = responses_answer(gateway, HTTPX, group, history(mark), stream=stream) + elif kind.startswith("chat"): + answer = chat_answer(gateway, HTTPX, group, chat_messages(mark), stream=stream) + else: + answer = messages_answer(gateway, HTTPX, group, chat_messages(mark), stream=stream) + except (httpx.HTTPError, AssertionError) as error: + return Fired(kind, mark, None, None, repr(error)) + return Fired(kind, mark, answer.status, answer.headers.get("x-litellm-call-id"), answer.text) + + +def assert_burst_hopped_stripped(gateway: Gateway, groups: Mapping[str, OrderedGroup], fired: Sequence[Fired]) -> None: + failures: Final = [shot for shot in fired if shot.status != 200] + assert not failures, failures + records: Final = observed(gateway) + for shot in fired: + group: Final = groups[shot.kind] + order_one: Final = [ + body for body in bodies_for(records, group.order_one_scenario) if shot.mark in json.dumps(body) + ] + order_two: Final = [ + body for body in bodies_for(records, group.order_two_scenario) if shot.mark in json.dumps(body) + ] + assert len(order_one) == 1 and ORDER_ONE_BLOB in json.dumps(order_one[0]), (shot, order_one) + assert len(order_two) == 1 and ORDER_ONE_BLOB not in json.dumps(order_two[0]), (shot, order_two) + call_ids: Final = [shot.call_id for shot in fired if shot.call_id is not None] + assert len(call_ids) == len(fired), fired + successes: Final = success_rows(call_ids) + assert sorted(string_value(row["litellm_call_id"]) for row in successes) == sorted(call_ids), successes + order_two_ids: Final = {group.order_two for group in groups.values()} + assert all(row["model_id"] in order_two_ids for row in successes), successes + + +@pytest.mark.timeout(300) +def test_concurrent_mixed_burst_hops_every_request_stripped_and_logs_each_once(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + groups: Final = burst_groups(scenario) + shots: Final = tuple((BURST_KINDS[index % len(BURST_KINDS)], marker()) for index in range(30)) + with lanes(gateway, [group.target for group in groups.values()], 30) as pinned: + with ThreadPoolExecutor(max_workers=30) as pool: + fired: Final = tuple( + pool.map( + lambda shot: fire(shot[0], groups, shot[1][0], shot[1][1]), zip(pinned, shots, strict=True) + ) + ) + assert_burst_hopped_stripped(gateway, groups, fired) + + +def peer_reply(request: Request) -> Reply: + identity: Final = f"resp_peer-{uuid.uuid4().hex[:8]}" + body: Final = json.dumps({**_responses_body(ORDER_ONE_BLOB), "id": identity}).replace("$UNIQUE_ID", identity) + return Reply(body=body.encode()) + + +def served_by_order_one(answer: Answer, order_one: str) -> bool: + return answer.status == 200 and answer.headers.get("x-litellm-model-id") == order_one + + +@pytest.mark.timeout(240) +def test_order_one_outage_mid_burst_hops_stripped_and_recovers_after_restart(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + name: Final = f"hop-{uuid.uuid4().hex[:12]}" + two_scenario: Final = f"{name}-o2" + order_two: Final = deployment( + scenario, name, scenario_api_base(scenario, two_scenario, healthy_json()), order=2 + ) + with wire_server(peer_reply) as peer: + order_one: Final = deployment(scenario, name, peer.url, order=1) + target: Final = Target(name, frozenset({order_one, order_two})) + port: Final = int(peer.url.rsplit(":", 1)[-1]) + warm: Final = marker() + assert served_by_order_one(responses_answer(gateway, HTTPX, target, history(warm)), order_one), warm + assert [json.loads(request.body)["input"] for request in peer.drain()] == [history(warm)] + shots: Final = tuple(marker() for _ in range(12)) + with lanes(gateway, [target], 12) as pinned: + with ThreadPoolExecutor(max_workers=12) as pool: + answers: Final = tuple( + pool.map( + lambda shot: responses_answer(shot[0], HTTPX, target, history(shot[1])), + zip(pinned, shots, strict=True), + ) + ) + for answer in answers: + assert answer.status == 200 and answer.headers.get("x-litellm-model-id") == order_two, answer.text + bodies: Final = bodies_for(observed(gateway), two_scenario) + assert sorted(json.dumps(body.get("input")) for body in bodies) == sorted( + json.dumps(stripped_history(mark)) for mark in shots + ), bodies + call_ids: Final = [answer.headers["x-litellm-call-id"] for answer in answers] + assert sorted(string_value(row["litellm_call_id"]) for row in success_rows(call_ids)) == sorted(call_ids) + with wire_server(peer_reply, port=port) as restarted: + recovered: Final = marker() + assert served_by_order_one(responses_answer(gateway, HTTPX, target, history(recovered)), order_one), ( + recovered + ) + assert [json.loads(request.body)["input"] for request in restarted.drain()] == [history(recovered)] + + +WORKER_PID: Final = re.compile(r"Started server process \[(\d+)\]") + + +def worker_pids(log: Path) -> tuple[int, ...]: + return tuple(int(match) for match in WORKER_PID.findall(log.read_text())) + + +def assert_completed_shots_hopped_stripped( + gateway: Gateway, groups: Mapping[str, OrderedGroup], fired: Sequence[Fired] +) -> None: + completed: Final = [shot for shot in fired if shot.status == 200] + assert completed, fired + records: Final = observed(gateway) + for shot in completed: + group: Final = groups[shot.kind] + order_one: Final = [ + body for body in bodies_for(records, group.order_one_scenario) if shot.mark in json.dumps(body) + ] + order_two: Final = [ + body for body in bodies_for(records, group.order_two_scenario) if shot.mark in json.dumps(body) + ] + assert len(order_one) == 1 and ORDER_ONE_BLOB in json.dumps(order_one[0]), (shot, order_one) + assert len(order_two) == 1 and ORDER_ONE_BLOB not in json.dumps(order_two[0]), (shot, order_two) + for shot in fired: + assert shot.status in (200, None) or shot.status >= 500, shot + + +@pytest.mark.timeout(2 * graceful_stop_seconds() + 2 * CONVERGENCE_SECONDS + 120) +def test_worker_killed_mid_burst_leaves_the_survivor_hopping_stripped(tmp_path: Path) -> None: + with gateway_from_environment() as upstream_gateway: + with owned_proxy_process(upstream_gateway, tmp_path, {}, workers=2) as owned: + owned.gateway.post("/config/update", {"router_settings": {"num_retries": 0}}) + with owned.gateway.scenario() as scenario: + groups: Final = burst_groups(scenario) + workers: Final = eventually(lambda: worker_pids(owned.log), lambda pids: len(pids) == 2, seconds=60) + shots: Final = tuple((BURST_KINDS[index % len(BURST_KINDS)], marker()) for index in range(30)) + with lanes(owned.gateway, [group.target for group in groups.values()], 30) as pinned: + with ThreadPoolExecutor(max_workers=30) as pool: + futures: Final = [ + pool.submit(fire, lane, groups, kind, mark) + for lane, (kind, mark) in zip(pinned, shots, strict=True) + ] + psutil.Process(workers[0]).send_signal(signal.SIGKILL) + fired: Final = tuple(future.result() for future in futures) + assert_completed_shots_hopped_stripped(owned.gateway, groups, fired) + probes: Final = tuple(fire(owned.gateway, groups, kind, marker()) for kind in BURST_KINDS) + assert all(probe.status == 200 for probe in probes), probes + assert_burst_hopped_stripped(owned.gateway, groups, probes) diff --git a/tests/integration/routing/test_fallback_retry_header_contracts.py b/tests/integration/routing/test_fallback_retry_header_contracts.py new file mode 100644 index 00000000000..ebed95a683c --- /dev/null +++ b/tests/integration/routing/test_fallback_retry_header_contracts.py @@ -0,0 +1,291 @@ +import threading +import time +import uuid +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from typing import Final, Literal + +import httpx +import pytest +from pydantic import JsonValue, TypeAdapter + +from tests.integration._support.database import read_rows +from tests.integration._support.client import ( + JSON_OBJECT, + Gateway, + Scenario, + eventually, + object_value, + string_value, +) +from tests.integration._support.openai_wire import chat_reply +from tests.integration._support.wire import Reply, Request, wire_server + +UNREACHABLE_API_BASE: Final = "http://127.0.0.1:9/v1" +END_USER_REQUESTS: Final = 10 +CONCURRENT_REQUESTS: Final = 25 +DISTRIBUTION_REQUESTS: Final = 20 +SLOW_UPSTREAM_SECONDS: Final = 3 +HELD_REPLY_SECONDS: Final = 30 +END_USER_ROW_SECONDS: Final = 70 +CUSTOM_FALLBACK_TEXT: Final = "custom fallback prompt" +Caller = Literal["virtual-key", "master-key"] +JSON_OBJECTS: Final = TypeAdapter(list[dict[str, JsonValue]]) + + +def _messages(text: str = "integration control") -> list[JsonValue]: + return [{"role": "user", "content": text}] + + +def _chat( + gateway: Gateway, + body: dict[str, JsonValue], + *, + key: str | None = None, + headers: dict[str, str] | None = None, +) -> httpx.Response: + return gateway.request("POST", "/v1/chat/completions", {"messages": _messages(), **body}, key=key, headers=headers) + + +def _body(response: httpx.Response) -> dict[str, JsonValue]: + return JSON_OBJECT.validate_json(response.content) + + +def _content(response: httpx.Response) -> str: + choices: Final = _body(response)["choices"] + assert isinstance(choices, list) and choices, response.text + return string_value(object_value(object_value(choices[0])["message"])["content"]) + + +def _scripted_model(gateway: Gateway, scenario: Scenario, statuses: list[int], num_retries: int) -> str: + upstream_model: Final = f"integration-{uuid.uuid4().hex}" + script_url: Final = f"{gateway.upstream_url}/__scripts/{upstream_model}" + configured: Final = httpx.post(script_url, json={"statuses": statuses}) + assert configured.status_code == 200, configured.text + scenario.cleanups.callback(httpx.delete, script_url) + return scenario.model(model=f"openai/{upstream_model}", num_retries=num_retries) + + +@dataclass(frozen=True, slots=True) +class UniqueModel: + name: str + upstream: str + deployment_id: str + + +def _unique_model(scenario: Scenario) -> UniqueModel: + upstream_model: Final = f"integration-{uuid.uuid4().hex}" + deployment_id: Final = f"integration-{uuid.uuid4().hex}" + model_info: Final[dict[str, JsonValue]] = {"id": deployment_id} + name: Final = scenario.model(model=f"openai/{upstream_model}", model_info=model_info) + return UniqueModel(name, upstream_model, deployment_id) + + +def _slow_reply(_: Request) -> Reply: + time.sleep(SLOW_UPSTREAM_SECONDS) + return chat_reply("chatcmpl-slow", "gpt-4o-mini", "late", stream=False) + + +def _fallback_reply(_: Request) -> Reply: + return chat_reply("chatcmpl-fallback", "gpt-4o-mini", "served by fallback", stream=False) + + +def _held_reply(release: threading.Event) -> Callable[[Request], Reply]: + def respond(_: Request) -> Reply: + assert release.wait(timeout=HELD_REPLY_SECONDS) + return chat_reply("chatcmpl-held", "gpt-4o-mini", "held", stream=False) + + return respond + + +def _delete_auto_created_end_user(gateway: Gateway, user_id: str) -> None: + eventually( + lambda: read_rows('SELECT user_id FROM "LiteLLM_EndUserTable" WHERE user_id = %s', (user_id,)), + lambda rows: len(rows) == 1, + seconds=END_USER_ROW_SECONDS, + ) + gateway.post("/end_user/delete", {"user_ids": [user_id]}) + assert read_rows('SELECT user_id FROM "LiteLLM_EndUserTable" WHERE user_id = %s', (user_id,)) == [] + + +def _status(gateway: Gateway, model: str) -> int: + return _chat(gateway, {"model": model}).status_code + + +def _served_model_id(gateway: Gateway, model: str) -> str: + response: Final = _chat(gateway, {"model": model}) + assert response.status_code == 200, response.text + return response.headers["x-litellm-model-id"] + + +def _concurrent_round(gateway: Gateway, bad: str, good: str) -> tuple[list[int], list[int]]: + with ThreadPoolExecutor(max_workers=CONCURRENT_REQUESTS * 2) as pool: + bad_calls: Final = [pool.submit(_status, gateway, bad) for _ in range(CONCURRENT_REQUESTS)] + good_calls: Final = [pool.submit(_status, gateway, good) for _ in range(CONCURRENT_REQUESTS)] + return [call.result() for call in bad_calls], [call.result() for call in good_calls] + + +def _deployment(gateway: Gateway, scenario: Scenario, model_name: str, model: str = "openai/gpt-4o-mini") -> str: + created: Final = gateway.post( + "/model/new", + { + "model_name": model_name, + "litellm_params": { + "model": model, + "api_key": "integration-provider-key", + "api_base": f"{gateway.upstream_url}/v1", + }, + }, + ) + deployment_id: Final = string_value(object_value(created["model_info"])["id"]) + scenario.cleanups.callback(scenario.delete_model, deployment_id) + return deployment_id + + +@pytest.mark.parametrize("caller", ["virtual-key", "master-key"]) +def test_end_user_budget_tpm_limit_rate_limits_their_requests(gateway: Gateway, caller: Caller) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + budget: Final = scenario.budget(tpm_limit=2) + end_user: Final = f"integration-{uuid.uuid4().hex}" + gateway.post("/end_user/new", {"user_id": end_user, "budget_id": budget}) + scenario.cleanups.callback(gateway.post, "/end_user/delete", {"user_ids": [end_user]}) + key: Final = scenario.key(models=[model]) if caller == "virtual-key" else gateway.key + control_user: Final = f"integration-{uuid.uuid4().hex}" + scenario.cleanups.callback(_delete_auto_created_end_user, gateway, control_user) + control: Final = _chat(gateway, {"model": model, "user": control_user}, key=key) + assert control.status_code == 200, control.text + statuses: Final = [ + _chat(gateway, {"model": model, "user": end_user}, key=key).status_code for _ in range(END_USER_REQUESTS) + ] + assert statuses.count(200) < 5, statuses + assert set(statuses) <= {200, 429}, statuses + + +def test_client_fallbacks_reach_an_allowed_model_and_name_a_denied_one(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + primary: Final = scenario.model(api_base=UNREACHABLE_API_BASE) + fallback: Final = _unique_model(scenario) + body: Final[dict[str, JsonValue]] = {"model": primary, "fallbacks": [fallback.name]} + served: Final = _chat(gateway, body, key=scenario.key(models=[primary, fallback.name])) + assert served.status_code == 200, served.text + assert _body(served)["model"] == fallback.upstream + assert served.headers["x-litellm-model-id"] == fallback.deployment_id + assert _content(served) + denied: Final = _chat(gateway, body, key=scenario.key(models=[primary])) + assert denied.status_code == 403, denied.text + assert fallback.name in denied.text + + +def test_client_fallback_with_custom_messages_sends_them_to_the_fallback(gateway: Gateway) -> None: + custom: Final = _messages(CUSTOM_FALLBACK_TEXT) + with gateway.scenario() as scenario, wire_server(_fallback_reply) as wire: + primary: Final = scenario.model(api_base=UNREACHABLE_API_BASE) + fallback: Final = scenario.model(api_base=wire.url) + body: Final[dict[str, JsonValue]] = { + "model": primary, + "fallbacks": [{"model": fallback, "messages": custom}], + } + served: Final = _chat(gateway, body, key=scenario.key(models=[primary, fallback])) + assert served.status_code == 200, served.text + assert _content(served) == "served by fallback" + forwarded: Final = [JSON_OBJECT.validate_json(request.body)["messages"] for request in wire.drain()] + assert forwarded == [custom] + denied: Final = _chat(gateway, body, key=scenario.key(models=[primary])) + assert denied.status_code == 403, denied.text + assert fallback in denied.text + assert wire.drain() == () + + +def test_rate_limited_deployment_is_retried_and_reports_retry_counts(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = _scripted_model(gateway, scenario, [429, 200], 50) + response: Final = _chat(gateway, {"model": model}) + assert response.status_code == 200, response.text + assert response.headers["x-litellm-attempted-retries"] == "1" + assert response.headers["x-litellm-max-retries"] == "50" + + +def test_request_fallbacks_reroute_after_a_connection_failure(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + primary: Final = scenario.model(api_base=UNREACHABLE_API_BASE) + fallback: Final = _unique_model(scenario) + response: Final = _chat(gateway, {"model": primary, "fallbacks": [fallback.name]}) + assert response.status_code == 200, response.text + assert _body(response)["model"] == fallback.upstream + assert response.headers["x-litellm-model-id"] == fallback.deployment_id + assert response.headers["x-litellm-attempted-fallbacks"] == "1" + + +def test_model_level_timeout_is_reported_on_a_timed_out_request(gateway: Gateway) -> None: + with gateway.scenario() as scenario, wire_server(_slow_reply) as wire: + response: Final = _chat(gateway, {"model": scenario.model(api_base=wire.url, timeout=1)}) + assert response.status_code == 408, response.text + assert response.headers["x-litellm-timeout"] == "1.0" + + +def test_request_timeout_header_overrides_the_model_timeout(gateway: Gateway) -> None: + with gateway.scenario() as scenario, wire_server(_slow_reply) as wire: + response: Final = _chat( + gateway, + {"model": scenario.model(api_base=wire.url, timeout=1)}, + headers={"x-litellm-timeout": "0.001"}, + ) + assert response.status_code == 408, response.text + assert response.headers["x-litellm-timeout"] == "0.001" + + +def test_failing_model_traffic_does_not_starve_concurrent_good_requests(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + bad: Final = scenario.model(api_base=UNREACHABLE_API_BASE) + good: Final = scenario.model() + for _ in range(2): + bad_calls, good_calls = _concurrent_round(gateway, bad, good) + assert good_calls == [200] * CONCURRENT_REQUESTS + assert 200 not in bad_calls + + +def test_rpm_limited_deployment_rejects_a_second_call_while_the_first_is_in_flight(gateway: Gateway) -> None: + release: Final = threading.Event() + with gateway.scenario() as scenario, wire_server(_held_reply(release)) as wire, ThreadPoolExecutor(1) as pool: + model: Final = scenario.model(api_base=wire.url, rpm=1) + first: Final = pool.submit(_chat, gateway, {"model": model}) + try: + eventually(wire.received.qsize, lambda count: count == 1) + second: Final = _chat(gateway, {"model": model}) + assert second.status_code == 429, second.text + assert wire.received.qsize() == 1 + finally: + release.set() + assert first.result().status_code == 200 + assert len(wire.drain()) == 1 + + +def test_model_group_with_two_deployments_serves_from_both(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + first_id: Final = f"integration-{uuid.uuid4().hex}" + model: Final = scenario.model(model_info={"id": first_id}) + second_id: Final = _deployment(gateway, scenario, model) + served_by: Final = {_served_model_id(gateway, model) for _ in range(DISTRIBUTION_REQUESTS)} + assert served_by == {first_id, second_id} + + +def test_unlisted_provider_model_resolves_through_a_wildcard_deployment(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + prefix: Final = f"integration{uuid.uuid4().hex}" + _deployment(gateway, scenario, f"{prefix}/*", "openai/*") + response: Final = _chat(gateway, {"model": f"{prefix}/gpt-4o-mini"}) + assert response.status_code == 200, response.text + assert _content(response) + + +def test_comma_separated_models_fan_out_to_one_response_per_model(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + first: Final = _unique_model(scenario) + second: Final = _unique_model(scenario) + response: Final = _chat(gateway, {"model": f"{first.name},{second.name}"}) + assert response.status_code == 200, response.text + replies: Final = JSON_OBJECTS.validate_json(response.content) + assert len(replies) == 2, replies + assert {string_value(reply["model"]) for reply in replies} == {first.upstream, second.upstream} diff --git a/tests/integration/routing/test_include_fallback_errors_wire.py b/tests/integration/routing/test_include_fallback_errors_wire.py new file mode 100644 index 00000000000..9075f5e41d3 --- /dev/null +++ b/tests/integration/routing/test_include_fallback_errors_wire.py @@ -0,0 +1,666 @@ +from __future__ import annotations + +import json +import os +import re +import uuid +from collections.abc import Callable, Iterator, Mapping +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from functools import partial +from pathlib import Path +from types import MappingProxyType +from typing import Final, Literal +from urllib.parse import unquote + +import httpx +import pytest +import yaml +from integration._support.anthropic_sse import error_body, message_json, message_stream, parse_sse, stream_reply +from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.database import read_rows +from integration._support.openai_wire import chat_reply, responses_reply +from integration._support.process import graceful_stop_seconds, owned_proxy +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +from litellm.constants import PROXY_CONFIG_RELOAD_INTERVAL_SECONDS + +Endpoint = Literal["chat", "messages", "responses"] + +_OPENAI_MODEL: Final = "gpt-5.4" +_ANTHROPIC_MODEL: Final = "claude-haiku-4-5" +_API_KEY: Final = "synthetic-fallback-errors-key" +_ROUTER_ONLY_KEYS: Final = ("include_fallback_errors", "silent_model") +_ANSWER: Final = "answered by" +_UNAUTHORIZED_MESSAGE: Final = "scripted 401: the primary key was revoked" +_NONCE: Final = re.compile(r"nonce=([0-9a-f]{32})") +_JSON: Final = TypeAdapter(dict[str, JsonValue]) +_ERRORS: Final = TypeAdapter(list[dict[str, JsonValue]]) +_ITEMS: Final = TypeAdapter(list[JsonValue]) +_UPSTREAM_URL_PLACEHOLDER: Final = "upstream-url" +_NOT_YET_ON_EVERY_WORKER: Final = ( + "Invalid model name passed in model=", + "There are no healthy deployments for this model", +) +_PROXY_WORKERS: Final = int(os.environ.get("INTEGRATION_PROXY_WORKERS", "1")) +_WORKER_SYNC_SECONDS: Final = 0.0 if _PROXY_WORKERS == 1 else PROXY_CONFIG_RELOAD_INTERVAL_SECONDS + 5.0 +_FRESH_CONNECTION: Final = MappingProxyType({"Connection": "close"}) +_SPEND_ROW_SECONDS: Final = 70 +_BURST: Final = 24 +_ENDPOINTS: Final[tuple[Endpoint, ...]] = ("chat", "messages", "responses") +_PATHS: Final = MappingProxyType( + {"chat": "/v1/chat/completions", "messages": "/v1/messages", "responses": "/v1/responses"} +) +_PROVIDERS: Final = MappingProxyType( + { + "chat": f"openai/{_OPENAI_MODEL}", + "messages": f"anthropic/{_ANTHROPIC_MODEL}", + "responses": f"openai/{_OPENAI_MODEL}", + } +) +_STREAMING: Final = (pytest.param(False, id="non-stream"), pytest.param(True, id="stream")) +_NON_BOOLEAN_FLAGS: Final = ( + pytest.param(1, True, id="int"), + pytest.param("", False, id="empty-string"), + pytest.param([], False, id="list"), + pytest.param("x" * 5120, True, id="five-kilobytes"), +) +_OPENAI_UNAUTHORIZED: Final = Reply( + status=401, + body=json.dumps( + { + "error": { + "message": _UNAUTHORIZED_MESSAGE, + "type": "invalid_request_error", + "param": None, + "code": "invalid_api_key", + } + } + ).encode(), +) +_ANTHROPIC_UNAUTHORIZED: Final = Reply(status=401, body=error_body(401, _UNAUTHORIZED_MESSAGE)) + +_EXPOSED_OPENAI_PRIMARY: Final = "exposed-openai-primary" +_EXPOSED_OPENAI_BACKUP: Final = "exposed-openai-backup" +_EXPOSED_OPENAI_SERVING: Final = "exposed-openai-serving" +_EXPOSED_OPENAI_FLAKY: Final = "exposed-openai-flaky" +_EXPOSED_ANTHROPIC_PRIMARY: Final = "exposed-anthropic-primary" +_EXPOSED_ANTHROPIC_BACKUP: Final = "exposed-anthropic-backup" + +_CHAT_BRIDGE_INHERITED_LEAKS: Final = frozenset( + f"{deployment}:include_fallback_errors" for deployment in (_EXPOSED_ANTHROPIC_PRIMARY, _EXPOSED_ANTHROPIC_BACKUP) +) + +pytestmark = pytest.mark.timeout(2 * graceful_stop_seconds() + 120) + + +def _prompt(nonce: str, *, outage: bool = False) -> str: + return f"which deployment answers this? nonce={nonce} outage={int(outage)}" + + +def _nonce_of(text: str) -> str: + found: Final = _NONCE.search(text) + assert found is not None, text + return found.group(1) + + +def _outage(deployment: str, raw: str) -> bool: + return deployment.endswith("primary") or (deployment.endswith("flaky") and "outage=1" in raw) + + +def _identity(endpoint: Endpoint, deployment: str, nonce: str) -> str: + match endpoint: + case "chat": + return f"chatcmpl-{deployment}-{nonce}" + case "messages": + return f"msg_{deployment}_{nonce}" + case "responses": + return f"resp_{deployment}_{nonce}" + + +def _peer(request: Request) -> Reply: + if request.method == "GET": + return Reply(body=json.dumps({"object": "list", "data": []}).encode()) + deployment, _, route = unquote(request.target).lstrip("/").partition("/") + raw: Final = request.body.decode() + body: Final = _JSON.validate_json(request.body) + nonce: Final = _nonce_of(raw) + stream: Final = body.get("stream") is True + text: Final = f"{_ANSWER} {deployment}" + if route == "v1/messages": + if _outage(deployment, raw): + return _ANTHROPIC_UNAUTHORIZED + identity: Final = _identity("messages", deployment, nonce) + if stream: + return stream_reply(message_stream(identity, _ANTHROPIC_MODEL, text)) + return Reply(body=message_json(identity, _ANTHROPIC_MODEL, text)) + if route.endswith("responses"): + if _outage(deployment, raw): + return _OPENAI_UNAUTHORIZED + return responses_reply(_identity("responses", deployment, nonce), _OPENAI_MODEL, text, stream=stream) + assert route == "chat/completions", request.target + if _outage(deployment, raw): + return _OPENAI_UNAUTHORIZED + return chat_reply(_identity("chat", deployment, nonce), _OPENAI_MODEL, text, stream=stream) + + +@dataclass(frozen=True, slots=True) +class _Posted: + deployment: str + route: str + raw: str + + def leaked(self) -> tuple[str, ...]: + return tuple(key for key in _ROUTER_ONLY_KEYS if key in self.raw) + + +def _posted_of(request: Request) -> _Posted: + deployment, _, route = unquote(request.target).lstrip("/").partition("/") + return _Posted(deployment, route, request.body.decode()) + + +def _all_posted(wire: Wire) -> tuple[_Posted, ...]: + return tuple(_posted_of(request) for request in wire.drain() if request.method == "POST") + + +def _posted(wire: Wire, nonce: str) -> tuple[_Posted, ...]: + return tuple(item for item in _all_posted(wire) if nonce in item.raw) + + +def _header(response: httpx.Response, name: str) -> str | None: + return response.headers[name] if name in response.headers else None + + +def _leaks(posted: tuple[_Posted, ...]) -> tuple[str, ...]: + return tuple(f"{item.deployment}:{','.join(item.leaked())}" for item in posted if item.leaked()) + + +def _hit(posted: tuple[_Posted, ...]) -> tuple[str, ...]: + return tuple(item.deployment for item in posted) + + +def _body(endpoint: Endpoint, model: str, nonce: str, *, stream: bool, outage: bool = False) -> dict[str, JsonValue]: + prompt: Final = _prompt(nonce, outage=outage) + match endpoint: + case "chat": + return {"model": model, "stream": stream, "messages": [{"role": "user", "content": prompt}]} + case "messages": + return { + "model": model, + "stream": stream, + "max_tokens": 32, + "messages": [{"role": "user", "content": prompt}], + } + case "responses": + return {"model": model, "stream": stream, "input": prompt} + + +def _output_item_identity(item: JsonValue) -> str: + return str(_JSON.validate_python(item)["id"]).removeprefix("msg_") + + +def _served_id(endpoint: Endpoint, response: httpx.Response, *, stream: bool) -> str: + if not stream: + body: Final = _JSON.validate_json(response.content) + if endpoint == "responses": + return _output_item_identity(_ITEMS.validate_python(body["output"])[0]) + return str(body["id"]) + events: Final = parse_sse(response.text) + match endpoint: + case "chat": + return str(events[0].data["id"]) + case "messages": + start: Final = next(event for event in events if event.event == "message_start") + return str(_JSON.validate_python(start.data["message"])["id"]) + case "responses": + done: Final = next(event for event in events if event.data.get("type") == "response.output_item.done") + return _output_item_identity(done.data["item"]) + + +def _settled(text: str) -> bool: + return not any(phrase in text for phrase in _NOT_YET_ON_EVERY_WORKER) + + +@dataclass(frozen=True, slots=True) +class _Sent: + nonce: str + response: httpx.Response + + +def _send( + gateway: Gateway, path: str, body: Callable[[str], Mapping[str, JsonValue]], *, key: str | None = None +) -> _Sent: + def attempt() -> _Sent: + nonce: Final = uuid.uuid4().hex + return _Sent(nonce, gateway.request("POST", path, body(nonce), key=key, headers=_FRESH_CONNECTION)) + + return eventually(attempt, lambda sent: _settled(sent.response.text), seconds=_WORKER_SYNC_SECONDS + 10) + + +@dataclass(frozen=True, slots=True) +class _Observed: + status: int + served_id: str + spend_id: str + text: str + attempted: str | None + errors_header: str | None + hit: tuple[str, ...] + leaks: tuple[str, ...] + + +def _observe(endpoint: Endpoint, wire: Wire, sent: _Sent, *, stream: bool) -> _Observed: + posted: Final = _posted(wire, sent.nonce) + response: Final = sent.response + return _Observed( + status=response.status_code, + served_id=_served_id(endpoint, response, stream=stream) if response.status_code == 200 else "", + spend_id=_spend_id(response) if response.status_code == 200 and not stream else "", + text=response.text, + attempted=_header(response, "x-litellm-attempted-fallbacks"), + errors_header=_header(response, "x-litellm-fallback-errors"), + hit=_hit(posted), + leaks=_leaks(posted), + ) + + +def _spend_id(response: httpx.Response) -> str: + return str(_JSON.validate_json(response.content)["id"]) + + +def _spend_row_lands(spend_id: str) -> None: + eventually( + lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (spend_id,)), + lambda rows: len(rows) == 1, + seconds=_SPEND_ROW_SECONDS, + ) + + +@dataclass(frozen=True, slots=True) +class _Registered: + gateway: Gateway + wire: Wire + primary: Mapping[Endpoint, str] + backup: Mapping[Endpoint, str] + + +@pytest.fixture(scope="module") +def registered() -> Iterator[_Registered]: + with gateway_from_environment() as gateway, wire_server(_peer) as wire, gateway.scenario() as scenario: + primary: Final[dict[Endpoint, str]] = { + endpoint: scenario.model(model=_PROVIDERS[endpoint], api_key=_API_KEY, api_base=f"{wire.url}/primary") + for endpoint in _ENDPOINTS + } + backup: Final[dict[Endpoint, str]] = { + endpoint: scenario.model(model=_PROVIDERS[endpoint], api_key=_API_KEY, api_base=f"{wire.url}/backup") + for endpoint in _ENDPOINTS + } + yield _Registered(gateway, wire, MappingProxyType(primary), MappingProxyType(backup)) + + +@pytest.mark.parametrize("endpoint", _ENDPOINTS) +@pytest.mark.parametrize("stream", _STREAMING) +def test_default_gateway_keeps_the_flag_off_the_wire_on_a_fallback( + registered: _Registered, endpoint: Endpoint, stream: bool +) -> None: + sent: Final = _send( + registered.gateway, + _PATHS[endpoint], + lambda nonce: { + **_body(endpoint, registered.primary[endpoint], nonce, stream=stream), + "fallbacks": [registered.backup[endpoint]], + "include_fallback_errors": True, + }, + ) + observed: Final = _observe(endpoint, registered.wire, sent, stream=stream) + assert observed.status == 200, observed + assert f"{_ANSWER} backup" in observed.text, observed + assert observed.served_id == _identity(endpoint, "backup", sent.nonce), observed + assert observed.leaks == (), observed + assert observed.hit == ("primary", "backup"), observed + assert observed.errors_header is None, observed + if not stream: + _spend_row_lands(observed.spend_id) + if endpoint == "chat" and not stream: + assert observed.attempted == "1", observed + + +@pytest.mark.parametrize(("value", "errors_reported"), _NON_BOOLEAN_FLAGS) +def test_default_gateway_keeps_a_non_boolean_flag_off_the_wire( + registered: _Registered, value: JsonValue, errors_reported: bool +) -> None: + sent: Final = _send( + registered.gateway, + _PATHS["chat"], + lambda nonce: { + **_body("chat", registered.primary["chat"], nonce, stream=False), + "fallbacks": [registered.backup["chat"]], + "include_fallback_errors": value, + }, + ) + observed: Final = _observe("chat", registered.wire, sent, stream=False) + assert observed.status == 200, observed + assert f"{_ANSWER} backup" in observed.text, observed + assert observed.leaks == (), observed + assert observed.hit == ("primary", "backup"), observed + assert observed.errors_header is None, observed + + +def test_default_gateway_keeps_a_duplicated_raw_flag_off_the_wire(registered: _Registered) -> None: + def raw_body(nonce: str) -> str: + body: Final = { + **_body("chat", registered.primary["chat"], nonce, stream=False), + "fallbacks": [registered.backup["chat"]], + } + return json.dumps(body)[:-1] + ', "include_fallback_errors": true, "include_fallback_errors": true}' + + def attempt() -> _Sent: + nonce: Final = uuid.uuid4().hex + response: Final = registered.gateway.client.post( + _PATHS["chat"], + content=raw_body(nonce), + headers={ + "Authorization": f"Bearer {registered.gateway.key}", + "Content-Type": "application/json", + **_FRESH_CONNECTION, + }, + ) + return _Sent(nonce, response) + + sent: Final = eventually(attempt, lambda s: _settled(s.response.text), seconds=_WORKER_SYNC_SECONDS + 10) + observed: Final = _observe("chat", registered.wire, sent, stream=False) + assert observed.status == 200, observed + assert observed.served_id == _identity("chat", "backup", sent.nonce), observed + assert observed.leaks == (), observed + assert observed.hit == ("primary", "backup"), observed + assert observed.errors_header is None, observed + + +def test_default_gateway_rejects_an_unauthenticated_flagged_request_before_the_wire(registered: _Registered) -> None: + nonce: Final = uuid.uuid4().hex + response: Final = registered.gateway.request( + "POST", + _PATHS["chat"], + { + **_body("chat", registered.primary["chat"], nonce, stream=False), + "fallbacks": [registered.backup["chat"]], + "include_fallback_errors": True, + }, + key="sk-not-a-key", + headers=_FRESH_CONNECTION, + ) + assert response.status_code == 401, response.text + assert _posted(registered.wire, nonce) == () + + +def _exposed_deployment(name: str, provider: str) -> dict[str, JsonValue]: + return { + "model_name": name, + "litellm_params": { + "model": provider, + "api_base": f"{_UPSTREAM_URL_PLACEHOLDER}/{name}", + "api_key": _API_KEY, + }, + } + + +def _exposed_config(wire: Wire, directory: Path) -> Path: + config: Final = _JSON.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + config["model_list"] = [ + _exposed_deployment(_EXPOSED_OPENAI_PRIMARY, _PROVIDERS["chat"]), + _exposed_deployment(_EXPOSED_OPENAI_BACKUP, _PROVIDERS["chat"]), + _exposed_deployment(_EXPOSED_OPENAI_SERVING, _PROVIDERS["chat"]), + _exposed_deployment(_EXPOSED_OPENAI_FLAKY, _PROVIDERS["chat"]), + _exposed_deployment(_EXPOSED_ANTHROPIC_PRIMARY, _PROVIDERS["messages"]), + _exposed_deployment(_EXPOSED_ANTHROPIC_BACKUP, _PROVIDERS["messages"]), + ] + config["general_settings"] = { + **_JSON.validate_python(config["general_settings"]), + "expose_fallback_errors_to_caller": True, + } + config["router_settings"] = { + "num_retries": 0, + "disable_cooldowns": True, + "fallbacks": [ + {_EXPOSED_OPENAI_PRIMARY: [_EXPOSED_OPENAI_BACKUP]}, + {_EXPOSED_OPENAI_FLAKY: [_EXPOSED_OPENAI_BACKUP]}, + {_EXPOSED_ANTHROPIC_PRIMARY: [_EXPOSED_ANTHROPIC_BACKUP]}, + ], + } + path: Final = directory / "include-fallback-errors-exposed.yaml" + path.write_text(yaml.safe_dump(config).replace(_UPSTREAM_URL_PLACEHOLDER, wire.url)) + return path + + +@dataclass(frozen=True, slots=True) +class _Exposed: + proxy: Gateway + wire: Wire + + +@pytest.fixture(scope="module") +def exposed(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Exposed]: + directory: Final = tmp_path_factory.mktemp("include-fallback-errors-exposed") + with gateway_from_environment() as gateway, wire_server(_peer) as wire: + with owned_proxy(gateway, directory, {}, config=_exposed_config(wire, directory), workers=2) as proxy: + yield _Exposed(proxy, wire) + + +def _errors_mention_the_outage(errors_header: str | None) -> bool: + if errors_header is None: + return False + errors: Final = _ERRORS.validate_json(errors_header) + return len(errors) == 1 and _UNAUTHORIZED_MESSAGE in str(errors[0]["message"]) + + +@pytest.mark.parametrize("stream", _STREAMING) +def test_exposed_chat_fallback_reports_the_errors_without_putting_the_flag_on_the_wire( + exposed: _Exposed, stream: bool +) -> None: + sent: Final = _send( + exposed.proxy, + _PATHS["chat"], + lambda nonce: { + **_body("chat", _EXPOSED_OPENAI_PRIMARY, nonce, stream=stream), + "include_fallback_errors": True, + }, + ) + observed: Final = _observe("chat", exposed.wire, sent, stream=stream) + assert observed.status == 200, observed + assert f"{_ANSWER} {_EXPOSED_OPENAI_BACKUP}" in observed.text, observed + assert observed.served_id == _identity("chat", _EXPOSED_OPENAI_BACKUP, sent.nonce), observed + assert observed.leaks == (), observed + assert observed.hit == (_EXPOSED_OPENAI_PRIMARY, _EXPOSED_OPENAI_BACKUP), observed + if not stream: + assert observed.attempted == "1", observed + assert _errors_mention_the_outage(observed.errors_header), observed + + +def test_exposed_chat_without_a_fallback_keeps_the_flag_off_the_wire_and_the_errors_header_off( + exposed: _Exposed, +) -> None: + sent: Final = _send( + exposed.proxy, + _PATHS["chat"], + lambda nonce: { + **_body("chat", _EXPOSED_OPENAI_SERVING, nonce, stream=False), + "include_fallback_errors": True, + }, + ) + observed: Final = _observe("chat", exposed.wire, sent, stream=False) + assert observed.status == 200, observed + assert observed.served_id == _identity("chat", _EXPOSED_OPENAI_SERVING, sent.nonce), observed + assert observed.leaks == (), observed + assert observed.hit == (_EXPOSED_OPENAI_SERVING,), observed + assert observed.attempted == "0", observed + assert observed.errors_header is None, observed + + +@pytest.mark.parametrize(("value", "errors_reported"), _NON_BOOLEAN_FLAGS) +def test_exposed_chat_keeps_a_non_boolean_flag_off_the_wire_and_the_errors_header_follows_its_truthiness( + exposed: _Exposed, value: JsonValue, errors_reported: bool +) -> None: + sent: Final = _send( + exposed.proxy, + _PATHS["chat"], + lambda nonce: { + **_body("chat", _EXPOSED_OPENAI_PRIMARY, nonce, stream=False), + "include_fallback_errors": value, + }, + ) + observed: Final = _observe("chat", exposed.wire, sent, stream=False) + assert observed.status == 200, observed + assert observed.served_id == _identity("chat", _EXPOSED_OPENAI_BACKUP, sent.nonce), observed + assert observed.leaks == (), observed + assert observed.hit == (_EXPOSED_OPENAI_PRIMARY, _EXPOSED_OPENAI_BACKUP), observed + assert observed.attempted == "1", observed + assert _errors_mention_the_outage(observed.errors_header) is errors_reported, observed + + +def test_exposed_proxy_rejects_an_unauthenticated_flagged_request_before_the_wire(exposed: _Exposed) -> None: + nonce: Final = uuid.uuid4().hex + response: Final = exposed.proxy.request( + "POST", + _PATHS["chat"], + {**_body("chat", _EXPOSED_OPENAI_PRIMARY, nonce, stream=False), "include_fallback_errors": True}, + key="sk-not-a-key", + headers=_FRESH_CONNECTION, + ) + assert response.status_code == 401, response.text + assert _posted(exposed.wire, nonce) == () + + +@pytest.mark.parametrize("stream", _STREAMING) +def test_exposed_messages_fallback_keeps_the_flag_off_the_wire(exposed: _Exposed, stream: bool) -> None: + sent: Final = _send( + exposed.proxy, + _PATHS["messages"], + lambda nonce: { + **_body("messages", _EXPOSED_ANTHROPIC_PRIMARY, nonce, stream=stream), + "include_fallback_errors": True, + }, + ) + observed: Final = _observe("messages", exposed.wire, sent, stream=stream) + assert observed.status == 200, observed + assert observed.served_id == _identity("messages", _EXPOSED_ANTHROPIC_BACKUP, sent.nonce), observed + assert observed.hit == (_EXPOSED_ANTHROPIC_PRIMARY, _EXPOSED_ANTHROPIC_BACKUP), observed + assert observed.leaks == (), observed + + +@pytest.mark.parametrize("stream", _STREAMING) +def test_exposed_responses_fallback_keeps_the_flag_off_the_wire(exposed: _Exposed, stream: bool) -> None: + sent: Final = _send( + exposed.proxy, + _PATHS["responses"], + lambda nonce: { + **_body("responses", _EXPOSED_OPENAI_PRIMARY, nonce, stream=stream), + "include_fallback_errors": True, + }, + ) + observed: Final = _observe("responses", exposed.wire, sent, stream=stream) + assert observed.status == 200, observed + assert observed.served_id == _identity("responses", _EXPOSED_OPENAI_BACKUP, sent.nonce), observed + assert observed.hit == (_EXPOSED_OPENAI_PRIMARY, _EXPOSED_OPENAI_BACKUP), observed + assert observed.leaks == (), observed + + +@pytest.mark.parametrize("stream", _STREAMING) +def test_exposed_messages_on_the_openai_deployment_bridges_the_fallback_with_a_clean_wire( + exposed: _Exposed, stream: bool +) -> None: + sent: Final = _send( + exposed.proxy, + _PATHS["messages"], + lambda nonce: { + **_body("messages", _EXPOSED_OPENAI_PRIMARY, nonce, stream=stream), + "include_fallback_errors": True, + }, + ) + observed: Final = _observe("messages", exposed.wire, sent, stream=stream) + assert observed.status == 200, observed + assert f"{_ANSWER} {_EXPOSED_OPENAI_BACKUP}" in observed.text, observed + assert observed.hit == (_EXPOSED_OPENAI_PRIMARY, _EXPOSED_OPENAI_BACKUP), observed + assert observed.leaks == (), observed + + +@pytest.mark.parametrize("stream", _STREAMING) +def test_exposed_responses_on_the_anthropic_deployment_falls_back_through_the_chat_bridge_with_no_new_key_on_the_wire( + exposed: _Exposed, stream: bool +) -> None: + sent: Final = _send( + exposed.proxy, + _PATHS["responses"], + lambda nonce: { + **_body("responses", _EXPOSED_ANTHROPIC_PRIMARY, nonce, stream=stream), + "include_fallback_errors": True, + }, + ) + observed: Final = _observe("responses", exposed.wire, sent, stream=stream) + assert observed.status == 200, observed + assert f"{_ANSWER} {_EXPOSED_ANTHROPIC_BACKUP}" in observed.text, observed + assert observed.hit == (_EXPOSED_ANTHROPIC_PRIMARY, _EXPOSED_ANTHROPIC_BACKUP), observed + assert set(observed.leaks) <= _CHAT_BRIDGE_INHERITED_LEAKS, observed + + +@dataclass(frozen=True, slots=True) +class _BurstRequest: + endpoint: Endpoint + stream: bool + outage: bool + nonce: str + + +def _burst_plan() -> tuple[_BurstRequest, ...]: + shapes: Final[tuple[tuple[Endpoint, bool], ...]] = (("chat", False), ("chat", True), ("responses", False)) + return tuple( + _BurstRequest(endpoint, stream, index % 2 == 1, uuid.uuid4().hex) + for index, (endpoint, stream) in enumerate(shapes * (_BURST // len(shapes))) + ) + + +def _fire(proxy: Gateway, request: _BurstRequest) -> httpx.Response: + return proxy.request( + "POST", + _PATHS[request.endpoint], + { + **_body( + request.endpoint, _EXPOSED_OPENAI_FLAKY, request.nonce, stream=request.stream, outage=request.outage + ), + "include_fallback_errors": True, + }, + headers=_FRESH_CONNECTION, + ) + + +def _expected_hit(request: _BurstRequest) -> tuple[str, ...]: + if request.outage: + return (_EXPOSED_OPENAI_FLAKY, _EXPOSED_OPENAI_BACKUP) + return (_EXPOSED_OPENAI_FLAKY,) + + +def _expected_served_id(request: _BurstRequest) -> str: + deployment: Final = _EXPOSED_OPENAI_BACKUP if request.outage else _EXPOSED_OPENAI_FLAKY + return _identity(request.endpoint, deployment, request.nonce) + + +def _check_burst_row(request: _BurstRequest, response: httpx.Response, posted: tuple[_Posted, ...]) -> None: + own: Final = tuple(item for item in posted if request.nonce in item.raw) + assert response.status_code == 200, (request, response.text) + assert _served_id(request.endpoint, response, stream=request.stream) == _expected_served_id(request), request + assert _hit(own) == _expected_hit(request), (request, own) + assert _leaks(own) == (), (request, own) + if request.endpoint == "chat" and not request.stream: + assert _header(response, "x-litellm-attempted-fallbacks") == ("1" if request.outage else "0"), request + assert _errors_mention_the_outage(_header(response, "x-litellm-fallback-errors")) is request.outage, request + + +def test_exposed_burst_with_scripted_outages_lands_every_prompt_once_per_hop_with_a_clean_wire( + exposed: _Exposed, +) -> None: + plan: Final = _burst_plan() + with ThreadPoolExecutor(max_workers=_BURST) as pool: + responses: Final = tuple(pool.map(partial(_fire, exposed.proxy), plan)) + posted: Final = _all_posted(exposed.wire) + for request, response in zip(plan, responses, strict=True): + _check_burst_row(request, response, posted) diff --git a/tests/integration/routing/test_priority_scheduler_queue_cleanup.py b/tests/integration/routing/test_priority_scheduler_queue_cleanup.py index a4606554d34..3c38acadaaf 100644 --- a/tests/integration/routing/test_priority_scheduler_queue_cleanup.py +++ b/tests/integration/routing/test_priority_scheduler_queue_cleanup.py @@ -46,7 +46,6 @@ USAGE: Final[dict[str, JsonValue]] = {"prompt_tokens": 5, "completion_tokens": 3 STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") OWNED_CELL_TIMEOUT: Final = 2 * graceful_stop_seconds() + 120 PAIR_CELL_TIMEOUT: Final = 3 * graceful_stop_seconds() + 120 -WORKER_HEALTHCHECK_ARGUMENTS: Final = ("--timeout_worker_healthcheck", str(int(graceful_stop_seconds()))) PINNED_CONNECTION_TIMEOUT_SECONDS: Final = 30 PINNED_LIMITS: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1, keepalive_expiry=30) ENDPOINTS: Final[tuple[Endpoint, ...]] = ( @@ -541,11 +540,7 @@ def test_in_memory_queue_forgets_served_requests_before_a_cooldown(gateway: Gate with ExitStack() as stack: wire: Final = stack.enter_context(wire_server(answering_model_discovery(upstream.respond))) config: Final = owned_config(tmp_path, wire, (INMEM_GROUP,), cooldown_settings(None)) - owned: Final = stack.enter_context( - owned_proxy_process( - gateway, tmp_path, {}, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS - ) - ) + owned: Final = stack.enter_context(owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2)) worker: Final = pinned(owned.gateway, stack) served: Final = new_marker() assert_served(post(worker, "/v1/chat/completions", chat_body(INMEM_GROUP, served, priority=1)), served) @@ -580,15 +575,9 @@ def pair(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Pair]: wire: Final = stack.enter_context(wire_server(answering_model_discovery(upstream.respond))) config: Final = owned_config(directory, wire, PAIR_GROUPS, cooldown_settings(cache), cancel_on_disconnect=True) overrides: Final = {"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)} - first: Final = stack.enter_context( - owned_proxy_process( - gateway, directory, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS - ) - ) + first: Final = stack.enter_context(owned_proxy_process(gateway, directory, overrides, config=config, workers=2)) second: Final = stack.enter_context( - owned_proxy_process( - gateway, directory, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS - ) + owned_proxy_process(gateway, directory, overrides, config=config, workers=2) ) yield Pair(first.gateway, second.gateway, cache, wire, upstream) @@ -797,9 +786,7 @@ def test_prioritized_requests_survive_a_redis_outage(gateway: Gateway, tmp_path: "REDIS_PORT": str(cache.port), "REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": "1", } - with owned_proxy_process( - gateway, tmp_path, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS - ) as owned: + with owned_proxy_process(gateway, tmp_path, overrides, config=config, workers=2) as owned: before: Final = assert_all_served(burst(owned.gateway, OUTAGE_GROUP, 12)) cache.stop() during: Final = assert_all_served(burst(owned.gateway, OUTAGE_GROUP, 12)) @@ -826,9 +813,7 @@ def test_sibling_worker_keeps_serving_prioritized_requests_after_a_worker_is_kil with owned_redis(tmp_path) as cache, wire_server(answering_model_discovery(upstream.respond)) as wire: config: Final = owned_config(tmp_path, wire, (KILL_GROUP,), redis_settings(cache)) overrides: Final = {"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)} - with owned_proxy_process( - gateway, tmp_path, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS - ) as owned: + with owned_proxy_process(gateway, tmp_path, overrides, config=config, workers=2) as owned: workers: Final = eventually( lambda: tuple(int(found.group(1)) for found in STARTED_WORKER.finditer(owned.log.read_text())), lambda pids: len(pids) == 2, diff --git a/tests/integration/sdk/test_router_include_fallback_errors_wire.py b/tests/integration/sdk/test_router_include_fallback_errors_wire.py new file mode 100644 index 00000000000..fb17e05b02a --- /dev/null +++ b/tests/integration/sdk/test_router_include_fallback_errors_wire.py @@ -0,0 +1,364 @@ +from __future__ import annotations + +import asyncio +import json +import re +import threading +import uuid +from collections.abc import Mapping +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from functools import partial +from typing import Final + +import pytest +from integration._support.client import eventually +from integration._support.openai_wire import chat_reply +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +import litellm +from litellm import CustomStreamWrapper, Router +from litellm.integrations.custom_logger import CustomLogger +from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage +from litellm.types.utils import Choices, ModelResponse, ModelResponseStream + +_MODEL: Final = "gpt-5.4" +_API_KEY: Final = "synthetic-fallback-errors-key" +_ROUTER_ONLY_KEYS: Final = ("include_fallback_errors", "silent_model") +_ANSWER: Final = "answered by" +_UNAUTHORIZED_MESSAGE: Final = "scripted 401: the primary key was revoked" +_NONCE: Final = re.compile(r"nonce=([0-9a-f]{32})") +_JSON: Final = TypeAdapter(dict[str, JsonValue]) +_OBJECT: Final = TypeAdapter(dict[str, object]) +_ERRORS: Final = TypeAdapter(list[dict[str, JsonValue]]) +_BURST: Final = 12 +_CALLBACK_WINDOW: Final = 15.0 +_STREAMING: Final = (pytest.param(False, id="non-stream"), pytest.param(True, id="stream")) +_NON_BOOLEAN_FLAGS: Final = ( + pytest.param(1, True, id="int"), + pytest.param("", False, id="empty-string"), + pytest.param([], False, id="list"), + pytest.param("x" * 5120, True, id="five-kilobytes"), + pytest.param(False, False, id="false"), +) +_UNAUTHORIZED: Final = Reply( + status=401, + body=json.dumps( + { + "error": { + "message": _UNAUTHORIZED_MESSAGE, + "type": "invalid_request_error", + "param": None, + "code": "invalid_api_key", + } + } + ).encode(), +) + + +def _prompt(nonce: str) -> str: + return f"which deployment answers this? nonce={nonce}" + + +def _nonce_of(body: Mapping[str, JsonValue]) -> str: + found: Final = _NONCE.search(json.dumps(body)) + assert found is not None, body + return found.group(1) + + +def _peer(request: Request) -> Reply: + if request.method == "GET": + return Reply(body=json.dumps({"object": "list", "data": []}).encode()) + deployment, _, route = request.target.lstrip("/").partition("/") + assert route == "chat/completions", request.target + if deployment == "primary": + return _UNAUTHORIZED + body: Final = _JSON.validate_json(request.body) + identity: Final = f"chatcmpl-{deployment}-{_nonce_of(body)}" + return chat_reply(identity, _MODEL, f"{_ANSWER} {deployment}", stream=body.get("stream") is True) + + +def _deployment(name: str, wire: Wire, **extra: JsonValue) -> dict[str, JsonValue]: + return { + "model_name": name, + "litellm_params": { + "model": f"openai/{_MODEL}", + "api_base": f"{wire.url}/{name}", + "api_key": _API_KEY, + **extra, + }, + } + + +def _router(wire: Wire, *, cache_responses: bool = False) -> Router: + return Router( + model_list=[ + _deployment("primary", wire), + _deployment("backup", wire), + _deployment("serving", wire), + _deployment("mirrored", wire, silent_model="shadow"), + _deployment("shadow", wire), + ], + fallbacks=[{"primary": ["backup"]}], + num_retries=0, + disable_cooldowns=True, + cache_responses=cache_responses, + ) + + +@dataclass(frozen=True, slots=True) +class _Posted: + deployment: str + raw: str + + def leaked(self) -> tuple[str, ...]: + return tuple(key for key in _ROUTER_ONLY_KEYS if key in self.raw) + + def nonce(self) -> str: + found: Final = _NONCE.search(self.raw) + assert found is not None, self.raw + return found.group(1) + + +def _posted(wire: Wire) -> tuple[_Posted, ...]: + return tuple( + _Posted(request.target.lstrip("/").partition("/")[0], request.body.decode()) + for request in wire.drain() + if request.method == "POST" + ) + + +def _leaks(posted: tuple[_Posted, ...]) -> tuple[str, ...]: + return tuple(f"{item.deployment}:{','.join(item.leaked())}" for item in posted if item.leaked()) + + +def _hit(posted: tuple[_Posted, ...]) -> tuple[str, ...]: + return tuple(item.deployment for item in posted) + + +def _nonces_at(posted: tuple[_Posted, ...], deployment: str) -> tuple[str, ...]: + return tuple(sorted(item.nonce() for item in posted if item.deployment == deployment)) + + +def _hidden(response: object) -> Mapping[str, object]: + return _OBJECT.validate_python(getattr(response, "_hidden_params", None) or {}) + + +def _headers(response: object) -> Mapping[str, object]: + return _OBJECT.validate_python(_hidden(response).get("additional_headers") or {}) + + +def _cache_hit(response: object) -> bool: + return _hidden(response).get("cache_hit") is True + + +def _error_messages(headers: Mapping[str, object]) -> tuple[str, ...]: + raw: Final = headers.get("x-litellm-fallback-errors") + if raw is None: + return () + return tuple(str(error["message"]) for error in _ERRORS.validate_json(str(raw))) + + +def _delta(chunk: object) -> str: + assert isinstance(chunk, ModelResponseStream), chunk + return "".join(str(choice.delta.content or "") for choice in chunk.choices) + + +def _messages(nonce: str) -> list[dict[str, str]]: + return [{"role": "user", "content": _prompt(nonce)}] + + +def _typed_messages(nonce: str) -> list[AllMessageValues]: + return [ChatCompletionUserMessage(role="user", content=_prompt(nonce))] + + +def _content(response: object) -> str: + assert isinstance(response, ModelResponse), response + choice: Final = response.choices[0] + assert isinstance(choice, Choices), choice + return str(choice.message.content) + + +@dataclass(frozen=True, slots=True) +class _Served: + text: str + headers: Mapping[str, object] + cache_hit: bool + + +@dataclass(frozen=True, slots=True) +class _Outcome: + text: str + attempted: object + errors: tuple[str, ...] + hit: tuple[str, ...] + leaks: tuple[str, ...] + + +def _outcome(wire: Wire, served: _Served) -> _Outcome: + posted: Final = _posted(wire) + return _Outcome( + text=served.text, + attempted=served.headers.get("x-litellm-attempted-fallbacks"), + errors=_error_messages(served.headers), + hit=_hit(posted), + leaks=_leaks(posted), + ) + + +def _complete(router: Router, model: str, *, stream: bool, nonce: str | None = None, **request: object) -> _Served: + response: Final = router.completion( + model=model, messages=_messages(nonce or uuid.uuid4().hex), stream=stream, **request + ) + if isinstance(response, CustomStreamWrapper): + return _Served("".join(_delta(chunk) for chunk in response), _headers(response), _cache_hit(response)) + return _Served(_content(response), _headers(response), _cache_hit(response)) + + +async def _acomplete(router: Router, model: str, *, stream: bool, **request: object) -> _Served: + response: Final = await router.acompletion( + model=model, messages=_typed_messages(uuid.uuid4().hex), stream=stream, **request + ) + if isinstance(response, CustomStreamWrapper): + parts: Final = [_delta(chunk) async for chunk in response] + return _Served("".join(parts), _headers(response), _cache_hit(response)) + return _Served(_content(response), _headers(response), _cache_hit(response)) + + +@pytest.mark.parametrize("stream", _STREAMING) +def test_sync_fallback_keeps_the_flag_off_the_wire_and_reports_the_errors(stream: bool) -> None: + with wire_server(_peer) as wire: + router: Final = _router(wire) + outcome: Final = _outcome(wire, _complete(router, "primary", stream=stream, include_fallback_errors=True)) + assert outcome.leaks == (), outcome + assert outcome.hit == ("primary", "backup"), outcome + assert outcome.text == f"{_ANSWER} backup", outcome + assert outcome.attempted == 1, outcome + if not stream: + assert len(outcome.errors) == 1 and _UNAUTHORIZED_MESSAGE in outcome.errors[0], outcome + + +@pytest.mark.parametrize("stream", _STREAMING) +def test_sync_matches_the_async_twin(stream: bool) -> None: + with wire_server(_peer) as wire: + router: Final = _router(wire) + twin: Final = _outcome( + wire, asyncio.run(_acomplete(router, "primary", stream=stream, include_fallback_errors=True)) + ) + observed: Final = _outcome(wire, _complete(router, "primary", stream=stream, include_fallback_errors=True)) + assert observed == twin, (observed, twin) + assert twin.leaks == (), twin + assert twin.hit == ("primary", "backup"), twin + + +def test_without_a_fallback_the_flag_still_stays_off_the_wire() -> None: + with wire_server(_peer) as wire: + router: Final = _router(wire) + outcome: Final = _outcome(wire, _complete(router, "serving", stream=False, include_fallback_errors=True)) + assert outcome.leaks == (), outcome + assert outcome.hit == ("serving",), outcome + assert outcome.text == f"{_ANSWER} serving", outcome + assert outcome.attempted == 0, outcome + assert outcome.errors == (), outcome + + +@pytest.mark.parametrize(("value", "errors_reported"), _NON_BOOLEAN_FLAGS) +def test_a_non_boolean_flag_stays_off_the_wire_and_the_errors_header_follows_its_truthiness( + value: object, errors_reported: bool +) -> None: + with wire_server(_peer) as wire: + router: Final = _router(wire) + outcome: Final = _outcome(wire, _complete(router, "primary", stream=False, include_fallback_errors=value)) + assert outcome.leaks == (), outcome + assert outcome.hit == ("primary", "backup"), outcome + assert outcome.text == f"{_ANSWER} backup", outcome + assert outcome.attempted == 1, outcome + assert (len(outcome.errors) == 1) is errors_reported, outcome + + +class _Recorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.mentions_the_flag: bool | None = None + self.fired = threading.Event() + + def _record(self, kwargs: Mapping[str, object]) -> None: + self.mentions_the_flag = "include_fallback_errors" in json.dumps(kwargs, default=str) + self.fired.set() + + def log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + self._record(kwargs) + + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + self._record(kwargs) + + +async def _acomplete_and_await_the_logger(router: Router, recorder: _Recorder) -> _Served: + served: Final = await _acomplete(router, "serving", stream=False, include_fallback_errors=True) + assert await asyncio.get_running_loop().run_in_executor(None, recorder.fired.wait, _CALLBACK_WINDOW) + return served + + +def test_sync_logger_kwargs_carry_the_flag_exactly_as_the_async_ones_do(monkeypatch: pytest.MonkeyPatch) -> None: + sync_recorder: Final = _Recorder() + async_recorder: Final = _Recorder() + with wire_server(_peer) as wire: + router: Final = _router(wire) + monkeypatch.setattr(litellm, "callbacks", [sync_recorder]) + _complete(router, "serving", stream=False, include_fallback_errors=True) + assert sync_recorder.fired.wait(_CALLBACK_WINDOW) + monkeypatch.setattr(litellm, "callbacks", [async_recorder]) + asyncio.run(_acomplete_and_await_the_logger(router, async_recorder)) + posted: Final = _posted(wire) + assert _leaks(posted) == (), posted + assert (sync_recorder.mentions_the_flag, async_recorder.mentions_the_flag) == (False, False) + + +def test_silent_model_shadow_traffic_carries_neither_router_only_key() -> None: + with wire_server(_peer) as wire: + router: Final = _router(wire) + served: Final = _complete(router, "mirrored", stream=False, include_fallback_errors=True) + eventually(wire.received.qsize, lambda count: count >= 2, seconds=20) + posted: Final = _posted(wire) + assert served.text == f"{_ANSWER} mirrored", served + assert tuple(sorted(_hit(posted))) == ("mirrored", "shadow"), posted + assert _leaks(posted) == (), posted + + +def test_a_cache_hit_repeats_the_answer_without_a_wire_request(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "cache", None) + with wire_server(_peer) as wire: + router: Final = _router(wire, cache_responses=True) + nonce: Final = uuid.uuid4().hex + first: Final = _complete(router, "serving", stream=False, nonce=nonce, include_fallback_errors=True) + posted: Final = _posted(wire) + second: Final = _complete(router, "serving", stream=False, nonce=nonce, include_fallback_errors=True) + again: Final = _posted(wire) + assert _hit(posted) == ("serving",), posted + assert _leaks(posted) == (), posted + assert again == (), again + assert second.text == first.text == f"{_ANSWER} serving" + assert (first.cache_hit, second.cache_hit) == (False, True), (first, second) + + +def _flagged(router: Router, nonce: str) -> _Served: + return _complete(router, "primary", stream=False, nonce=nonce, include_fallback_errors=True) + + +def test_a_sync_burst_lands_every_prompt_once_on_each_side_of_the_fallback() -> None: + nonces: Final = tuple(uuid.uuid4().hex for _ in range(_BURST)) + with wire_server(_peer) as wire: + router: Final = _router(wire) + with ThreadPoolExecutor(max_workers=_BURST) as pool: + served: Final = tuple(pool.map(partial(_flagged, router), nonces)) + posted: Final = _posted(wire) + assert tuple(item.text for item in served) == (f"{_ANSWER} backup",) * _BURST, served + assert tuple(item.headers.get("x-litellm-attempted-fallbacks") for item in served) == (1,) * _BURST, served + assert _leaks(posted) == (), posted + assert _nonces_at(posted, "primary") == tuple(sorted(nonces)), posted + assert _nonces_at(posted, "backup") == tuple(sorted(nonces)), posted diff --git a/tests/integration/sdk/test_router_sync_stream_fallback_wire.py b/tests/integration/sdk/test_router_sync_stream_fallback_wire.py index f5095c5a778..fc5fb028123 100644 --- a/tests/integration/sdk/test_router_sync_stream_fallback_wire.py +++ b/tests/integration/sdk/test_router_sync_stream_fallback_wire.py @@ -199,28 +199,23 @@ def test_router_retries_configured(client: str) -> None: assert _deployments_hit(wire) == ("primary", "backup") -@dataclass(frozen=True, slots=True) -class _Outcome: - text: str | None - error: str | None - hit: tuple[str, ...] - - -def _outcome(client: str, wire: Wire, router: Router, **request: object) -> _Outcome: - try: - streamed: Final = _stream(client, router, **request) - except litellm.APIConnectionError as error: - return _Outcome(text=None, error=type(error).__name__, hit=_deployments_hit(wire)) - return _Outcome(text=streamed.text, error=None, hit=_deployments_hit(wire)) - - -def test_per_request_fallback_list_behaves_like_the_async_twin() -> None: +@pytest.mark.parametrize("client", _CLIENTS) +def test_per_request_fallback_list(client: str) -> None: with wire_server(_peer(_PRIMARY_DIES)) as wire: router: Final = _router(wire, ("primary", "backup")) - twin: Final = _outcome("async", wire, router, fallbacks=_PRIMARY_TO_BACKUP) - observed: Final = _outcome("sync", wire, router, fallbacks=_PRIMARY_TO_BACKUP) - assert observed == twin, (observed, twin) - assert observed.hit[:1] == ("primary",), observed + streamed: Final = _stream(client, router, fallbacks=_PRIMARY_TO_BACKUP) + assert streamed.text == "answered by the backup", streamed + assert streamed.attempted_fallbacks == 1, streamed + assert _deployments_hit(wire) == ("primary", "backup") + + +@pytest.mark.parametrize("client", _CLIENTS) +def test_per_request_fallbacks_none_turns_the_router_list_off(client: str) -> None: + with wire_server(_peer(_PRIMARY_DIES)) as wire: + router: Final = _router(wire, ("primary", "backup"), fallbacks=_PRIMARY_TO_BACKUP) + with pytest.raises(litellm.APIConnectionError, match="overloaded"): + _stream(client, router, fallbacks=None) + assert _deployments_hit(wire) == ("primary",) @pytest.mark.parametrize("client", _CLIENTS) diff --git a/tests/integration/spend/test_background_response_poll_retirement.py b/tests/integration/spend/test_background_response_poll_retirement.py new file mode 100644 index 00000000000..f1b1e96901c --- /dev/null +++ b/tests/integration/spend/test_background_response_poll_retirement.py @@ -0,0 +1,370 @@ +import itertools +import os +import uuid +from collections.abc import Iterator, Sequence +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows, scratch_database, write_rows +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario +from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse +from pydantic import JsonValue + +from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE +from litellm.responses.utils import ResponsesAPIRequestUtils +from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN + +INPUT_COST_PER_TOKEN: Final = 0.001 +OUTPUT_COST_PER_TOKEN: Final = 0.002 +INPUT_TOKENS: Final = 19 +OUTPUT_TOKENS: Final = 7 +SCHEDULER_PERIOD_CEILING_SECONDS: Final = 31 +ONE_POLL_SECONDS: Final = 2 * SCHEDULER_PERIOD_CEILING_SECONDS + 8 +TWO_POLLS_SECONDS: Final = 4 * SCHEDULER_PERIOD_CEILING_SECONDS + 8 +OWNED_PROXY_CELL_SECONDS: Final = 2 * graceful_stop_seconds() + 240 +PROVIDER_ID: Final = "resp_$REQUEST_ID" + + +def _response_body(status: str) -> dict[str, JsonValue]: + completed: Final = status == "completed" + return { + "id": PROVIDER_ID, + "object": "response", + "created_at": 1, + "status": status, + "background": True, + "store": False, + "error": None, + "incomplete_details": None, + "model": "gpt-4o-mini", + "output": ( + [ + { + "id": "msg_$REQUEST_ID", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "pong", "annotations": []}], + } + ] + if completed + else [] + ), + "parallel_tool_calls": True, + "temperature": 1.0, + "tool_choice": "auto", + "tools": [], + "top_p": 1.0, + "truncation": "disabled", + "usage": ( + { + "input_tokens": INPUT_TOKENS, + "output_tokens": OUTPUT_TOKENS, + "total_tokens": INPUT_TOKENS + OUTPUT_TOKENS, + "input_tokens_details": {"cached_tokens": 0}, + "output_tokens_details": {"reasoning_tokens": 0}, + } + if completed + else None + ), + "metadata": {}, + } + + +def _json(body: dict[str, JsonValue], status: int = 200) -> JsonResponse: + return JsonResponse(content_type="application/json", body=body, status=status) + + +def _provider_404(message: str) -> JsonResponse: + return _json({"error": {"message": message, "type": "invalid_request_error", "param": None, "code": None}}, 404) + + +def _routes(submission_status: str, retrieve: JsonResponse) -> RoutedResponse: + return RoutedResponse( + content_type="application/x-routed", + routes={"POST /responses": _json(_response_body(submission_status)), f"GET /responses/{PROVIDER_ID}": retrieve}, + ) + + +def _gone_routes(submission_status: str = "queued") -> RoutedResponse: + return _routes(submission_status, _provider_404(f"Response with id '{PROVIDER_ID}' not found.")) + + +def _scenario_id(marker: str) -> str: + return f"bg-{marker}-{uuid.uuid4().hex[:12]}" + + +def _scripted_deployment(scenario: Scenario, marker: str, routes: RoutedResponse) -> tuple[str, ScenarioHandle]: + return _deployment_on(scenario, _scenario_id(marker), routes) + + +def _register_deployment( + scenario: Scenario, scenario_id: str, routes: RoutedResponse +) -> tuple[str, str, ScenarioHandle]: + handle: Final = register_scenario(scenario_id, routes) + scenario.cleanups.callback(delete_scenario, handle) + created: Final = scenario.gateway.post( + "/model/new", + { + "model_name": f"bg-{sha256(scenario_id.encode()).hexdigest()[:12]}", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "sk-scripted-provider", + "api_base": handle.api_base(), + "input_cost_per_token": INPUT_COST_PER_TOKEN, + "output_cost_per_token": OUTPUT_COST_PER_TOKEN, + }, + }, + ) + return string_value(created["model_name"]), string_value(object_value(created["model_info"])["id"]), handle + + +def _deployment_on(scenario: Scenario, scenario_id: str, routes: RoutedResponse) -> tuple[str, ScenarioHandle]: + model_name, model_id, handle = _register_deployment(scenario, scenario_id, routes) + scenario.cleanups.callback(scenario.delete_model, model_id) + return model_name, handle + + +def _forget_row(unified_id: str, database_url: str | None = None) -> None: + write_rows( + 'DELETE FROM "LiteLLM_ManagedObjectTable" WHERE unified_object_id = %s', + (unified_id,), + database_url=database_url, + ) + + +def _submit_background_response(scenario: Scenario, key: str, model: str, database_url: str | None = None) -> str: + response: Final = scenario.gateway.request( + "POST", "/v1/responses", {"model": model, "input": "poll me", "background": True, "store": False}, key=key + ) + assert response.status_code == 200, response.text + unified_id: Final = string_value(object_value(response.json())["id"]) + scenario.cleanups.callback(_forget_row, unified_id, database_url) + return unified_id + + +def _row_status(unified_id: str, database_url: str | None = None) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT status FROM "LiteLLM_ManagedObjectTable" WHERE unified_object_id = %s', + (unified_id,), + database_url=database_url, + ) + + +def _await_status(unified_id: str, status: str, seconds: float, database_url: str | None = None) -> None: + assert eventually( + lambda: _row_status(unified_id, database_url), lambda rows: rows == [{"status": status}], seconds=seconds + ) == [{"status": status}] + + +def _observed_requests(gateway: Gateway) -> tuple[dict[str, JsonValue], ...]: + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream: + return tuple(map(object_value, upstream.get("/__observations").json()["requests"])) + + +def _calls_to(requests: Sequence[dict[str, JsonValue]], handle: ScenarioHandle) -> tuple[dict[str, JsonValue], ...]: + return tuple(request for request in requests if string_value(request["path"]).startswith(f"/{handle.scenario_id}/")) + + +def _upstream_log(gateway: Gateway, handle: ScenarioHandle) -> Iterator[tuple[dict[str, JsonValue], ...]]: + fresh: Final = (_calls_to(_observed_requests(gateway), handle) for _ in itertools.count()) + return itertools.accumulate(fresh, lambda seen, calls: (*seen, *calls)) + + +def _polls(log: Sequence[dict[str, JsonValue]], handle: ScenarioHandle) -> tuple[dict[str, JsonValue], ...]: + poll_path: Final = f"/{handle.scenario_id}/responses/resp_{handle.scenario_id}" + return tuple(call for call in log if call["method"] == "GET" and call["path"] == poll_path) + + +def _await_polls(gateway: Gateway, handle: ScenarioHandle, at_least: int, seconds: float) -> None: + log: Final = _upstream_log(gateway, handle) + polls: Final = _polls( + eventually(lambda: next(log), lambda seen: len(_polls(seen, handle)) >= at_least, seconds), handle + ) + assert len(polls) >= at_least, polls + + +def _submissions(log: Sequence[dict[str, JsonValue]], handle: ScenarioHandle) -> tuple[dict[str, JsonValue], ...]: + return tuple( + call for call in log if call["method"] == "POST" and call["path"] == f"/{handle.scenario_id}/responses" + ) + + +@pytest.mark.timeout(240) +@pytest.mark.parametrize("submission_status", ["queued", "in_progress"]) +def test_a_response_gone_at_the_provider_is_retired_from_polling(gateway: Gateway, submission_status: str) -> None: + with gateway.scenario() as scenario: + key: Final = scenario.key() + model, handle = _scripted_deployment(scenario, "gone", _gone_routes(submission_status)) + log: Final = _upstream_log(gateway, handle) + unified_id: Final = _submit_background_response(scenario, key, model) + _await_status(unified_id, "stale_expired", ONE_POLL_SECONDS) + seen: Final = next(log) + assert [call["body"] for call in _submissions(seen, handle)] == [ + {"model": "gpt-4o-mini", "input": "poll me", "background": True, "store": False} + ] + assert len(_polls(seen, handle)) >= 1, seen + caller_view: Final = gateway.request("GET", f"/v1/responses/{unified_id}", key=key) + assert caller_view.status_code == 404, caller_view.text + assert f"Response with id 'resp_{handle.scenario_id}' not found." in caller_view.text + first_clock, _ = _scripted_deployment(scenario, "clock1", _gone_routes()) + _await_status(_submit_background_response(scenario, key, first_clock), "stale_expired", ONE_POLL_SECONDS) + polls_after_first_clock: Final = len(_polls(next(log), handle)) + second_clock, _ = _scripted_deployment(scenario, "clock2", _gone_routes()) + _await_status(_submit_background_response(scenario, key, second_clock), "stale_expired", ONE_POLL_SECONDS) + assert len(_polls(next(log), handle)) == polls_after_first_clock + assert _row_status(unified_id) == [{"status": "stale_expired"}] + + +@pytest.mark.timeout(180) +def test_a_404_that_does_not_name_the_response_keeps_the_row_queued_for_retry(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model, handle = _scripted_deployment(scenario, "vague404", _routes("queued", _provider_404("Not found."))) + unified_id: Final = _submit_background_response(scenario, scenario.key(), model) + _await_polls(gateway, handle, 2, TWO_POLLS_SECONDS) + assert _row_status(unified_id) == [{"status": "queued"}] + + +@pytest.mark.timeout(180) +def test_a_provider_error_other_than_404_keeps_the_row_queued_for_retry(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + outage: Final = _json({"error": {"message": "The server had an error.", "type": "server_error"}}, 500) + model, handle = _scripted_deployment(scenario, "outage", _routes("queued", outage)) + unified_id: Final = _submit_background_response(scenario, scenario.key(), model) + _await_polls(gateway, handle, 2, TWO_POLLS_SECONDS) + assert _row_status(unified_id) == [{"status": "queued"}] + + +def _polled_provider_id(row: dict[str, JsonValue]) -> str: + return ResponsesAPIRequestUtils.decode_responses_api_response_id(string_value(row["request_id"]))["response_id"] + + +def _poll_spend_rows(handle: ScenarioHandle) -> list[dict[str, JsonValue]]: + billed: Final = read_rows( + 'SELECT request_id, call_type, status, prompt_tokens, completion_tokens, spend FROM "LiteLLM_SpendLogs" ' + "WHERE metadata->>'internal_call_origin' = %s", + (BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN,), + ) + return [ + {column: value for column, value in row.items() if column != "request_id"} + for row in billed + if _polled_provider_id(row) == f"resp_{handle.scenario_id}" + ] + + +@pytest.mark.timeout(180) +def test_a_completed_response_is_marked_completed_and_its_poll_is_billed(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model, handle = _scripted_deployment(scenario, "done", _routes("queued", _json(_response_body("completed")))) + unified_id: Final = _submit_background_response(scenario, scenario.key(), model) + _await_status(unified_id, "completed", ONE_POLL_SECONDS) + spend_rows: Final = eventually(lambda: _poll_spend_rows(handle), lambda rows: len(rows) >= 1, seconds=70) + assert spend_rows[0] == { + "call_type": "aget_responses", + "status": "success", + "prompt_tokens": INPUT_TOKENS, + "completion_tokens": OUTPUT_TOKENS, + "spend": pytest.approx(INPUT_TOKENS * INPUT_COST_PER_TOKEN + OUTPUT_TOKENS * OUTPUT_COST_PER_TOKEN), + } + + +def _insert_unreadable_row(scenario: Scenario) -> str: + unified_id: Final = f"resp_unreadable-{uuid.uuid4().hex[:12]}" + write_rows( + 'INSERT INTO "LiteLLM_ManagedObjectTable" ' + '("id", "unified_object_id", "model_object_id", "file_object", "file_purpose", "status", "created_at", "updated_at") ' + "VALUES (%s, %s, %s, '[]'::jsonb, 'response', 'queued', NOW() - INTERVAL '1 hour', NOW())", + (str(uuid.uuid4()), unified_id, unified_id), + ) + scenario.cleanups.callback(_forget_row, unified_id) + return unified_id + + +@pytest.mark.timeout(180) +def test_a_row_the_poll_cannot_prepare_skips_only_that_row(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + unreadable_id: Final = _insert_unreadable_row(scenario) + model, _ = _scripted_deployment(scenario, "gone", _gone_routes()) + unified_id: Final = _submit_background_response(scenario, scenario.key(), model) + _await_status(unified_id, "stale_expired", ONE_POLL_SECONDS) + assert _row_status(unreadable_id) == [{"status": "queued"}] + + +@pytest.mark.timeout(300) +def test_responses_gone_at_the_provider_do_not_starve_a_newer_response_out_of_cost_polling(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + key: Final = scenario.key() + gone_ids: Final = tuple( + _submit_background_response( + scenario, key, _scripted_deployment(scenario, f"gone{index}", _gone_routes())[0] + ) + for index in range(MAX_OBJECTS_PER_POLL_CYCLE) + ) + completable_model, completable = _scripted_deployment( + scenario, "done", _routes("queued", _json(_response_body("completed"))) + ) + completable_id: Final = _submit_background_response(scenario, key, completable_model) + _await_status(completable_id, "completed", TWO_POLLS_SECONDS) + assert [_row_status(gone_id) for gone_id in gone_ids] == [ + [{"status": "stale_expired"}] + ] * MAX_OBJECTS_PER_POLL_CYCLE + spend_rows: Final = eventually(lambda: _poll_spend_rows(completable), lambda rows: len(rows) >= 1, seconds=70) + assert spend_rows[0]["spend"] == pytest.approx( + INPUT_TOKENS * INPUT_COST_PER_TOKEN + OUTPUT_TOKENS * OUTPUT_COST_PER_TOKEN + ) + + +@pytest.mark.timeout(OWNED_PROXY_CELL_SECONDS) +def test_a_404_on_a_response_whose_deployment_left_the_router_keeps_the_row_queued_for_retry( + gateway: Gateway, tmp_path: Path +) -> None: + deployment_scenario_id: Final = _scenario_id("left") + provider_id: Final = f"resp_{deployment_scenario_id}" + env_handle: Final = register_scenario( + _scenario_id("envbase"), + RoutedResponse( + content_type="application/x-routed", + routes={f"GET /responses/{provider_id}": _provider_404(f"Response with id '{provider_id}' not found.")}, + ), + ) + try: + with scratch_database() as database_url: + overrides: Final = { + "DATABASE_URL": database_url, + "OPENAI_API_BASE": env_handle.api_base(), + "OPENAI_API_KEY": "sk-scripted-provider", + } + with owned_proxy_process( + gateway, tmp_path, overrides, remove_environment=("DATABASE_URL_READ_REPLICA",) + ) as owned: + with owned.gateway.scenario() as scenario: + outage: Final = _json( + {"error": {"message": "The server had an error.", "type": "server_error"}}, 500 + ) + model, model_id, _ = _register_deployment( + scenario, deployment_scenario_id, _routes("queued", outage) + ) + unified_id: Final = _submit_background_response(scenario, scenario.key(), model, database_url) + owned.gateway.post("/model/delete", {"id": model_id}) + log: Final = _upstream_log(gateway, env_handle) + fallback_polls: Final = eventually( + lambda: next(log), + lambda seen: len(_fallback_polls(seen, env_handle, provider_id)) >= 2, + seconds=TWO_POLLS_SECONDS, + ) + assert len(_fallback_polls(fallback_polls, env_handle, provider_id)) >= 2, fallback_polls + assert _row_status(unified_id, database_url) == [{"status": "queued"}] + finally: + delete_scenario(env_handle) + + +def _fallback_polls( + log: Sequence[dict[str, JsonValue]], env_handle: ScenarioHandle, provider_id: str +) -> tuple[dict[str, JsonValue], ...]: + poll_path: Final = f"/{env_handle.scenario_id}/responses/{provider_id}" + return tuple(call for call in log if call["method"] == "GET" and call["path"] == poll_path) diff --git a/tests/integration/spend/test_batch_completion_accounting.py b/tests/integration/spend/test_batch_completion_accounting.py index 4ecab10f942..bbc0300b3f0 100644 --- a/tests/integration/spend/test_batch_completion_accounting.py +++ b/tests/integration/spend/test_batch_completion_accounting.py @@ -63,7 +63,7 @@ def _failed_line(index: int) -> str: ) -def _batch_routes(model: str) -> RoutedResponse: +def batch_routes(model: str) -> RoutedResponse: output_lines: Final = ( _succeeded_line(1, model, **FIRST_LINE), _succeeded_line(2, model, **SECOND_LINE), @@ -139,7 +139,7 @@ def _batches_total_tokens(endpoints: dict[str, object] | None) -> int | None: return int(total_tokens) if isinstance(total_tokens, (int, float, str)) else None -def _input_file(model: str) -> bytes: +def batch_input_file(model: str) -> bytes: return ( "\n".join( json.dumps( @@ -166,13 +166,13 @@ def test_completed_batch_spend_row_records_reasoning_tokens_and_error_file_failu with gateway.scenario() as scenario: key: Final = scenario.key() scenario_id: Final = f"batch-accounting-{uuid.uuid4().hex[:12]}" - handle: Final = register_scenario(scenario_id, _batch_routes("gpt-4o-mini")) + handle: Final = register_scenario(scenario_id, batch_routes("gpt-4o-mini")) scenario.cleanups.callback(delete_scenario, handle) model: Final = scenario.model(api_base=handle.api_base()) file_response: Final = gateway.request_multipart( "/v1/files", {"purpose": "batch", "model": model}, - {"file": ("in.jsonl", _input_file(model), "application/jsonl")}, + {"file": ("in.jsonl", batch_input_file(model), "application/jsonl")}, key=key, ) assert file_response.status_code == 200, file_response.text @@ -235,7 +235,7 @@ BATCH_SPEND: Final = (BATCH_PROMPT_TOKENS * INPUT_COST_PER_TOKEN + BATCH_COMPLET def test_completed_batch_spend_lands_under_batches_in_team_endpoint_activity(gateway: Gateway) -> None: with gateway.scenario() as scenario: scenario_id: Final = f"batch-endpoint-{uuid.uuid4().hex[:12]}" - handle: Final = register_scenario(scenario_id, _batch_routes("gpt-4o-mini")) + handle: Final = register_scenario(scenario_id, batch_routes("gpt-4o-mini")) scenario.cleanups.callback(delete_scenario, handle) model: Final = scenario.model( api_base=handle.api_base(), @@ -247,7 +247,7 @@ def test_completed_batch_spend_lands_under_batches_in_team_endpoint_activity(gat file_response: Final = gateway.request_multipart( "/v1/files", {"purpose": "batch", "model": model}, - {"file": ("in.jsonl", _input_file(model), "application/jsonl")}, + {"file": ("in.jsonl", batch_input_file(model), "application/jsonl")}, key=key, ) assert file_response.status_code == 200, file_response.text diff --git a/tests/integration/spend/test_budget_limit_envelopes.py b/tests/integration/spend/test_budget_limit_envelopes.py new file mode 100644 index 00000000000..6ac2a11e657 --- /dev/null +++ b/tests/integration/spend/test_budget_limit_envelopes.py @@ -0,0 +1,172 @@ +from __future__ import annotations + +import json +import re +import uuid +from collections.abc import Sequence +from hashlib import sha256 +from typing import Final, Literal + +import httpx +import pytest +from pydantic import JsonValue + +from tests.integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value +from tests.integration._support.database import read_rows +from tests.integration._support.provider import PROVIDER_URL, SharedProvider +from tests.integration._support.wire import Reply + +_CALL_COST: Final = 0.02 +_TINY_BUDGET: Final = 0.0000000005 +_LIMIT_FIELDS: Final = ("max_budget", "rpm_limit", "tpm_limit") + + +def _completion() -> Reply: + return Reply( + body=json.dumps( + { + "id": f"chatcmpl-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "scripted"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + ).encode() + ) + + +def _priced_model(scenario: Scenario) -> str: + return scenario.model(api_base=f"{PROVIDER_URL}/v1", input_cost_per_token=0.001, output_cost_per_token=0.002) + + +def _ask(gateway: Gateway, model: str, key: str) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"budget {uuid.uuid4().hex}"}]}, + key=key, + ) + + +def _served(gateway: Gateway, provider: SharedProvider, model: str, key: str) -> None: + provider.expect(_completion()) + response: Final = _ask(gateway, model, key) + assert response.status_code == 200, response.text + assert len(provider.received()) == 1 + + +def _spend_reaches(table: Literal["key", "team"], identity: str, amount: float) -> None: + query: Final = ( + 'SELECT spend FROM "LiteLLM_VerificationToken" WHERE token = %s' + if table == "key" + else 'SELECT spend FROM "LiteLLM_TeamTable" WHERE team_id = %s' + ) + eventually( + lambda: read_rows(query, (identity,)), + lambda rows: len(rows) == 1 and float(str(rows[0]["spend"])) >= amount, + seconds=70, + ) + + +_COST_AND_LIMIT: Final = re.compile(r"Current cost: ([^\s,]+), Max budget: ([^\s,]+)") + + +def _budget_refusal( + gateway: Gateway, provider: SharedProvider, model: str, key: str, *, spent: float, limit: float +) -> str: + refused: Final = _ask(gateway, model, key) + assert refused.status_code == 422, f"{refused.status_code} {refused.text}" + error: Final = object_value(JSON_OBJECT.validate_json(refused.content)["error"]) + assert error["type"] == "budget_exceeded", refused.text + assert error["code"] == "422", refused.text + message: Final = string_value(error["message"]) + assert "Budget has been exceeded!" in message, refused.text + figures: Final = _COST_AND_LIMIT.search(message) + assert figures is not None, message + assert float(figures[1]) == pytest.approx(spent), message + assert float(figures[2]) == pytest.approx(limit), message + assert provider.received() == () + return message + + +def _hashed(key: str) -> str: + return sha256(key.encode()).hexdigest() + + +def test_a_key_with_a_tiny_budget_serves_once_and_then_answers_budget_exceeded( + gateway: Gateway, provider: SharedProvider +) -> None: + with gateway.scenario() as scenario: + model: Final = _priced_model(scenario) + key: Final = scenario.key(models=[model], max_budget=_TINY_BUDGET) + _served(gateway, provider, model, key) + _spend_reaches("key", _hashed(key), _CALL_COST) + _budget_refusal(gateway, provider, model, key, spent=_CALL_COST, limit=_TINY_BUDGET) + + +def test_a_key_with_a_zero_budget_is_refused_before_the_provider_is_called( + gateway: Gateway, provider: SharedProvider +) -> None: + with gateway.scenario() as scenario: + model: Final = _priced_model(scenario) + key: Final = scenario.key(models=[model], max_budget=0) + _budget_refusal(gateway, provider, model, key, spent=0.0, limit=0.0) + + +def test_a_key_with_room_for_two_calls_serves_both_before_answering_budget_exceeded( + gateway: Gateway, provider: SharedProvider +) -> None: + with gateway.scenario() as scenario: + model: Final = _priced_model(scenario) + limit: Final = _CALL_COST * 1.5 + key: Final = scenario.key(models=[model], max_budget=limit) + _served(gateway, provider, model, key) + _spend_reaches("key", _hashed(key), _CALL_COST) + _served(gateway, provider, model, key) + _spend_reaches("key", _hashed(key), _CALL_COST * 2) + _budget_refusal(gateway, provider, model, key, spent=_CALL_COST * 2, limit=limit) + + +def test_a_team_key_serves_once_and_then_answers_the_team_budget_envelope( + gateway: Gateway, provider: SharedProvider +) -> None: + with gateway.scenario() as scenario: + model: Final = _priced_model(scenario) + team: Final = scenario.team(models=[model], max_budget=_TINY_BUDGET) + key: Final = scenario.key(team_id=team, models=[model]) + _served(gateway, provider, model, key) + _spend_reaches("team", team, _CALL_COST) + message: Final = _budget_refusal(gateway, provider, model, key, spent=_CALL_COST, limit=_TINY_BUDGET) + assert f"Team={team}" in message, message + + +def _limits(record: dict[str, JsonValue]) -> Sequence[JsonValue]: + return [record[field] for field in _LIMIT_FIELDS] + + +@pytest.mark.parametrize("field", _LIMIT_FIELDS) +def test_a_key_limit_is_set_by_update_and_reset_to_null(gateway: Gateway, field: str) -> None: + with gateway.scenario() as scenario: + key: Final = scenario.key(max_budget=None, rpm_limit=None, tpm_limit=None) + raised: Final = gateway.post("/key/update", {"key": key, field: 10}) + assert raised[field] == 10, raised + assert [value for name, value in zip(_LIMIT_FIELDS, _limits(raised)) if name != field] == [None, None] + cleared: Final = gateway.post("/key/update", {"key": key, field: None}) + assert _limits(cleared) == [None, None, None], cleared + saved: Final = object_value(gateway.get("/key/info", {"key": key})["info"]) + assert _limits(saved) == [None, None, None], saved + + +@pytest.mark.parametrize("field", _LIMIT_FIELDS) +def test_a_team_limit_is_set_by_update_and_reset_to_null(gateway: Gateway, field: str) -> None: + with gateway.scenario() as scenario: + team: Final = scenario.team(max_budget=None, rpm_limit=None, tpm_limit=None) + raised: Final = object_value(gateway.post("/team/update", {"team_id": team, field: 10})["data"]) + assert raised[field] == 10, raised + cleared: Final = object_value(gateway.post("/team/update", {"team_id": team, field: None})["data"]) + assert _limits(cleared) == [None, None, None], cleared + saved: Final = object_value(gateway.get("/team/info", {"team_id": team})["team_info"]) + assert _limits(saved) == [None, None, None], saved diff --git a/tests/integration/spend/test_spend_attribution_contracts.py b/tests/integration/spend/test_spend_attribution_contracts.py new file mode 100644 index 00000000000..f5036160ce1 --- /dev/null +++ b/tests/integration/spend/test_spend_attribution_contracts.py @@ -0,0 +1,102 @@ +from contextlib import ExitStack +from itertools import chain +from typing import Final + +from pydantic import JsonValue + +from tests.integration._support.client import ( + JSON_OBJECT, + Gateway, + Scenario, + delete_key_if_present, + eventually, + object_value, + string_value, +) +from tests.integration._support.database import read_rows + +MEMBER_BUDGET: Final = 0.0000001 +SPEND_LANDING_SECONDS: Final = 70 + + +def test_spend_log_of_an_org_team_key_records_org_and_team(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + organization: Final = scenario.organization(models=[model]) + team: Final = scenario.team(organization_id=organization, models=[model]) + key: Final = scenario.key(team_id=team, models=[model]) + request_id: Final = string_value(gateway.chat(model, key=key)["id"]) + rows: Final = eventually( + lambda: read_rows( + 'SELECT team_id, metadata::text AS metadata FROM "LiteLLM_SpendLogs" WHERE request_id = %s', + (request_id,), + ), + lambda found: len(found) == 1, + seconds=SPEND_LANDING_SECONDS, + ) + metadata: Final = JSON_OBJECT.validate_json(string_value(rows[0]["metadata"])) + assert metadata["user_api_key_org_id"] == organization + assert metadata["user_api_key_team_id"] == team + assert rows[0]["team_id"] == team + + +def _team_memberships(team: JsonValue, user_id: str) -> list[dict[str, JsonValue]]: + entries: Final = object_value(team).get("team_memberships") or [] + assert isinstance(entries, list) + return [object_value(entry) for entry in entries if object_value(entry).get("user_id") == user_id] + + +def _membership(gateway: Gateway, user_id: str, team_id: str) -> dict[str, JsonValue]: + teams: Final = gateway.get("/user/info", {"user_id": user_id})["teams"] + assert isinstance(teams, list) + memberships: Final = list( + chain.from_iterable( + _team_memberships(team, user_id) for team in teams if object_value(team)["team_id"] == team_id + ) + ) + assert len(memberships) == 1, teams + return memberships[0] + + +def test_team_member_budget_blocks_the_member_after_spend_lands(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + team: Final = scenario.team() + created: Final = gateway.post( + "/user/new", + {"user_id": f"integration-{model}", "team_id": team, "models": [model], "max_budget": 10.0}, + ) + user: Final = string_value(created["user_id"]) + key: Final = string_value(created["key"]) + with ExitStack() as member_cleanups: + member_cleanups.callback(scenario.delete_user, user) + member_cleanups.callback(delete_key_if_present, gateway, key) + gateway.post("/team/member_update", {"team_id": team, "user_id": user, "max_budget_in_team": MEMBER_BUDGET}) + membership: Final = _membership(gateway, user, team) + scenario.cleanups.callback(scenario.delete_budget, string_value(membership["budget_id"])) + scenario.cleanups.push(member_cleanups.pop_all()) + assert object_value(membership["litellm_budget_table"])["max_budget"] == MEMBER_BUDGET + assert ( + gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "first"}]}, + key=key, + ).status_code + == 200 + ) + eventually( + lambda: read_rows( + 'SELECT spend FROM "LiteLLM_TeamMembership" WHERE user_id = %s AND team_id = %s', (user, team) + ), + lambda rows: len(rows) == 1 and isinstance(spend := rows[0]["spend"], float) and spend >= MEMBER_BUDGET, + seconds=SPEND_LANDING_SECONDS, + ) + blocked: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "second"}]}, + key=key, + ) + assert blocked.status_code != 200, blocked.text + assert "Budget has been exceeded" in blocked.text diff --git a/tests/integration/spend/test_spend_logs_metadata_fields.py b/tests/integration/spend/test_spend_logs_metadata_fields.py new file mode 100644 index 00000000000..d7f1dbfd998 --- /dev/null +++ b/tests/integration/spend/test_spend_logs_metadata_fields.py @@ -0,0 +1,509 @@ +import json +from collections.abc import Callable, Generator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from hashlib import sha256 +from pathlib import Path +from typing import Final +from uuid import uuid4 + +import httpx +import pytest +from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows, scratch_database +from integration._support.process import owned_proxy +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.spend.test_batch_completion_accounting import batch_input_file, batch_routes +from integration.streaming.test_stream_contracts import text_stream +from pydantic import JsonValue + +CACHED_PROMPT_TOKENS: Final = 4 + + +def _config( + tmp_path: Path, + spend_logs_metadata_fields: Mapping[str, JsonValue] | None, + *, + model_list: tuple[Mapping[str, JsonValue], ...] = (), + litellm_settings: Mapping[str, JsonValue] | None = None, + guardrails: tuple[Mapping[str, JsonValue], ...] = (), +) -> Path: + retention: Final = ( + {} if spend_logs_metadata_fields is None else {"spend_logs_metadata_fields": dict(spend_logs_metadata_fields)} + ) + config: Final = tmp_path / f"spend-logs-metadata-fields-{uuid4()}.json" + config.write_text( + json.dumps( + { + "model_list": [dict(model) for model in model_list], + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + "store_model_in_db": True, + "proxy_batch_write_at": 1, + "proxy_batch_polling_interval": 1, + **retention, + }, + "litellm_settings": dict(litellm_settings or {}), + "guardrails": [dict(guardrail) for guardrail in guardrails], + } + ) + ) + return config + + +def _respond(request: Request) -> Reply: + if request.method == "GET": + return Reply(body=b'{"object":"list","data":[]}') + return Reply( + body=json.dumps( + { + "id": f"chatcmpl-{uuid4()}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 2, + "total_tokens": 12, + "prompt_tokens_details": {"cached_tokens": CACHED_PROMPT_TOKENS}, + }, + } + ).encode() + ) + + +def _stored_row_after_one_chat( + gateway: Gateway, tmp_path: Path, spend_logs_metadata_fields: Mapping[str, JsonValue] +) -> tuple[dict[str, JsonValue], str]: + with ( + wire_server(_respond) as wire, + owned_proxy(gateway, tmp_path, {}, config=_config(tmp_path, spend_logs_metadata_fields)) as isolated, + isolated.scenario() as scenario, + ): + model: Final = scenario.model(model="openai/gpt-4o-mini", api_base=wire.url, api_key="synthetic-openai-key") + api_key: Final = scenario.key(key_alias=f"metadata-fields-{uuid4()}", models=[model]) + response_id: Final = string_value(isolated.chat(model, key=api_key)["id"]) + rows: Final = eventually( + lambda: read_rows( + 'SELECT metadata, proxy_server_request, response FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (response_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + return rows[0], sha256(api_key.encode()).hexdigest() + + +def test_excluded_metadata_fields_are_not_stored_but_still_reach_daily_spend(gateway: Gateway, tmp_path: Path) -> None: + row, hashed_key = _stored_row_after_one_chat( + gateway, tmp_path, {"exclude": ["model_map_information", "usage_object"]} + ) + + metadata: Final = object_value(row["metadata"]) + assert "model_map_information" not in metadata, metadata + assert "usage_object" not in metadata, metadata + assert {"status", "cold_storage_object_key"} <= set(metadata), metadata + assert string_value(metadata["user_api_key_alias"]).startswith("metadata-fields-") + assert row["proxy_server_request"] == {} + assert row["response"] == {} + daily: Final = eventually( + lambda: read_rows( + 'SELECT prompt_tokens, cache_read_input_tokens FROM "LiteLLM_DailyUserSpend" WHERE api_key=%s', + (hashed_key,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert daily[0]["cache_read_input_tokens"] == CACHED_PROMPT_TOKENS, daily + + +def test_included_metadata_fields_are_the_only_ones_stored_besides_always_kept( + gateway: Gateway, tmp_path: Path +) -> None: + row, _ = _stored_row_after_one_chat(gateway, tmp_path, {"include": ["user_api_key_alias"]}) + + assert set(object_value(row["metadata"])) == {"status", "cold_storage_object_key", "user_api_key_alias"} + + +def test_guardrail_usage_is_tracked_when_guardrail_information_is_not_stored(gateway: Gateway, tmp_path: Path) -> None: + guardrail_name: Final = f"metadata-fields-guardrail-{uuid4()}" + with ( + wire_server(lambda _: Reply(body=b'{"action":"NONE"}')) as policy, + wire_server(_respond) as wire, + owned_proxy( + gateway, + tmp_path, + {}, + config=_config( + tmp_path, + {"exclude": ["guardrail_information"]}, + guardrails=( + { + "guardrail_name": guardrail_name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": policy.url, + "api_key": "synthetic-guardrail-key", + }, + }, + ), + ), + ) as isolated, + isolated.scenario() as scenario, + ): + model: Final = scenario.model(model="openai/gpt-4o-mini", api_base=wire.url, api_key="synthetic-openai-key") + response_id: Final = string_value(isolated.chat(model, key=scenario.key(models=[model]))["id"]) + assert len(policy.drain()) == 1 + rows: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (response_id,)), + lambda values: len(values) == 1, + seconds=70, + ) + assert "guardrail_information" not in object_value(rows[0]["metadata"]), rows[0] + indexed: Final = eventually( + lambda: read_rows( + 'SELECT guardrail_id FROM "LiteLLM_SpendLogGuardrailIndex" WHERE request_id=%s', (response_id,) + ), + lambda values: len(values) == 1, + seconds=70, + ) + metrics: Final = eventually( + lambda: read_rows( + 'SELECT requests_evaluated, passed_count FROM "LiteLLM_DailyGuardrailMetrics" WHERE guardrail_id=%s', + (string_value(indexed[0]["guardrail_id"]),), + ), + lambda values: values == [{"requests_evaluated": 1, "passed_count": 1}], + seconds=70, + ) + assert metrics == [{"requests_evaluated": 1, "passed_count": 1}], metrics + + +EXCLUDED: Final = ("model_map_information", "user_api_key_alias") + + +@dataclass(frozen=True, slots=True) +class _Isolated: + proxy: Gateway + scenario: Scenario + key: str + database_url: str | None = None + + @property + def hashed_key(self) -> str: + return sha256(self.key.encode()).hexdigest() + + def rows(self, count: int, where: str = "TRUE", seconds: int = 70) -> tuple[dict[str, JsonValue], ...]: + return tuple( + eventually( + lambda: read_rows( + "SELECT request_id, call_type, status, spend, cache_hit, metadata " + f'FROM "LiteLLM_SpendLogs" WHERE api_key=%s AND {where} ORDER BY "startTime"', + (self.hashed_key,), + database_url=self.database_url, + ), + lambda values: len(values) == count, + seconds=seconds, + ) + ) + + +@contextmanager +def _isolated(gateway: Gateway, config: Path, tmp_path: Path, database_url: str | None = None) -> Generator[_Isolated]: + database: Final = {} if database_url is None else {"DATABASE_URL": database_url} + replica: Final = () if database_url is None else ("DATABASE_URL_READ_REPLICA",) + with ( + owned_proxy(gateway, tmp_path, database, config=config, remove_environment=replica) as proxy, + proxy.scenario() as scenario, + ): + key: Final = scenario.key(key_alias=f"metadata-fields-{uuid4()}") + yield _Isolated(proxy, scenario, key, database_url) + + +def _assert_filtered(row: Mapping[str, JsonValue], *kept: str) -> dict[str, JsonValue]: + metadata: Final = object_value(row["metadata"]) + assert not set(EXCLUDED) & set(metadata), metadata + assert {"status", "cold_storage_object_key", *kept} <= set(metadata), metadata + return metadata + + +def _anthropic_event(event: Mapping[str, JsonValue]) -> bytes: + return f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() + + +def _anthropic_stream(model: JsonValue) -> tuple[bytes, ...]: + message: Final = {**_anthropic_message(model), "content": [], "stop_reason": None} + events: Final[tuple[Mapping[str, JsonValue], ...]] = ( + {"type": "message_start", "message": message}, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "ok"}}, + {"type": "content_block_stop", "index": 0}, + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 2}}, + {"type": "message_stop"}, + ) + return tuple(_anthropic_event(event) for event in events) + + +def _respond_streaming_or_not(request: Request) -> Reply: + if request.method == "GET": + return Reply(body=b'{"object":"list","data":[]}') + if request.target.startswith("/v1/messages"): + return _anthropic(request) + if JSON_OBJECT.validate_json(request.body).get("stream") is True: + return Reply(content_type="text/event-stream", chunks=text_stream(f"chatcmpl-{uuid4()}")) + return _respond(request) + + +def _stream(proxy: Gateway, key: str, path: str, body: Mapping[str, JsonValue]) -> None: + with proxy.client.stream("POST", path, json=dict(body), headers={"Authorization": f"Bearer {key}"}) as response: + assert response.status_code == 200, response.read().decode() + lines: Final = tuple(response.iter_lines()) + assert any(line.startswith("data:") for line in lines), lines + + +def _call(proxy: Gateway, key: str, path: str, body: Mapping[str, JsonValue]) -> None: + response: Final = proxy.request("POST", path, body, key=key) + assert response.status_code == 200, response.text + + +def test_every_inference_surface_and_stream_stores_filtered_metadata(gateway: Gateway, tmp_path: Path) -> None: + with ( + wire_server(_respond_streaming_or_not) as wire, + _isolated(gateway, _config(tmp_path, {"exclude": list(EXCLUDED)}), tmp_path) as isolated, + ): + model: Final = isolated.scenario.model( + model="deepseek/gpt-4o-mini", api_base=wire.url + "/v1", api_key="synthetic-key" + ) + anthropic: Final = isolated.scenario.model( + model="anthropic/claude-sonnet-4-6", api_base=wire.url, api_key="synthetic-key" + ) + prompt: Final[list[JsonValue]] = [{"role": "user", "content": f"surfaces {uuid4()}"}] + calls: Final[tuple[Callable[[Gateway, str, str, Mapping[str, JsonValue]], None], ...]] = ( + _stream, + _call, + _stream, + _call, + _stream, + ) + requests: Final[tuple[tuple[str, Mapping[str, JsonValue]], ...]] = ( + ("/v1/chat/completions", {"model": model, "messages": prompt, "stream": True}), + ("/v1/messages", {"model": anthropic, "max_tokens": 16, "messages": prompt}), + ("/v1/messages", {"model": anthropic, "max_tokens": 16, "messages": prompt, "stream": True}), + ("/v1/responses", {"model": model, "input": f"surfaces {uuid4()}"}), + ("/v1/responses", {"model": model, "input": f"surfaces {uuid4()}", "stream": True}), + ) + for send, (path, body) in zip(calls, requests, strict=True): + send(isolated.proxy, isolated.key, path, body) + + rows: Final = isolated.rows(len(requests)) + for row in rows: + assert row["status"] == "success", row + _assert_filtered(row, "usage_object") + + +def test_failed_request_keeps_its_error_and_status_but_drops_excluded_metadata( + gateway: Gateway, tmp_path: Path +) -> None: + def fail(request: Request) -> Reply: + if request.method == "GET": + return Reply(body=b'{"object":"list","data":[]}') + return Reply(status=500, body=b'{"error":{"message":"scripted upstream failure","type":"server_error"}}') + + with ( + wire_server(fail) as wire, + _isolated(gateway, _config(tmp_path, {"exclude": list(EXCLUDED)}), tmp_path) as isolated, + ): + model: Final = isolated.scenario.model(model="openai/gpt-4o-mini", api_base=wire.url, api_key="synthetic") + response: Final = isolated.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "fail"}]}, + key=isolated.key, + ) + assert response.status_code >= 500, response.text + + row: Final = isolated.rows(1)[0] + assert row["status"] == "failure", row + metadata: Final = _assert_filtered(row, "error_information") + assert metadata["status"] == "failure", metadata + assert "scripted upstream failure" in json.dumps(metadata["error_information"]), metadata + + +def test_response_cache_hit_row_is_filtered_and_charged_nothing(gateway: Gateway, tmp_path: Path) -> None: + config: Final = _config( + tmp_path, {"exclude": list(EXCLUDED)}, litellm_settings={"cache": True, "cache_params": {"type": "local"}} + ) + with wire_server(_respond) as wire, _isolated(gateway, config, tmp_path) as isolated: + model: Final = isolated.scenario.model(model="openai/gpt-4o-mini", api_base=wire.url, api_key="synthetic") + text: Final = f"cache {uuid4()}" + first: Final = string_value(isolated.proxy.chat(model, key=isolated.key, text=text)["id"]) + isolated.proxy.chat(model, key=isolated.key, text=text) + + rows: Final = isolated.rows(2) + assert len([request for request in wire.drain() if request.method == "POST"]) == 1 + paid, hit = sorted(rows, key=lambda row: row["cache_hit"] == "True") + assert paid["request_id"] == first and float(str(paid["spend"])) > 0, rows + assert hit["cache_hit"] == "True" and string_value(hit["request_id"]).startswith(first + "_cache_hit"), rows + assert float(str(hit["spend"])) == 0, rows + for row in rows: + _assert_filtered(row, "usage_object") + + +def test_batch_cost_row_is_filtered_and_charged_once(gateway: Gateway, tmp_path: Path) -> None: + with _isolated(gateway, _config(tmp_path, {"exclude": list(EXCLUDED)}), tmp_path) as isolated: + handle: Final = register_scenario(f"metadata-fields-batch-{uuid4().hex[:12]}", batch_routes("gpt-4o-mini")) + isolated.scenario.cleanups.callback(delete_scenario, handle) + model: Final = isolated.scenario.model(api_base=handle.api_base()) + uploaded: Final = isolated.proxy.request_multipart( + "/v1/files", + {"purpose": "batch", "model": model}, + {"file": ("in.jsonl", batch_input_file(model), "application/jsonl")}, + key=isolated.key, + ) + assert uploaded.status_code == 200, uploaded.text + batch: Final = isolated.proxy.post( + "/v1/batches", + { + "input_file_id": string_value(JSON_OBJECT.validate_json(uploaded.content)["id"]), + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": model, + }, + key=isolated.key, + ) + batch_path: Final = f"/v1/batches/{string_value(batch['id'])}" + retrievals: Final = tuple(isolated.proxy.request("GET", batch_path, key=isolated.key) for _ in range(2)) + assert all(r.status_code == 200 and r.json()["status"] == "completed" for r in retrievals), retrievals + + row: Final = isolated.rows(1, "call_type='aretrieve_batch'")[0] + assert float(str(row["spend"])) > 0, row + metadata: Final = _assert_filtered(row, "usage_object") + assert (metadata["batch_successful_requests"], metadata["batch_failed_requests"]) == (2, 3), metadata + + +def _update_retention(proxy: Gateway, value: JsonValue) -> httpx.Response: + return proxy.request( + "POST", + "/config/field/update", + {"field_name": "spend_logs_metadata_fields", "field_value": value, "config_type": "general_settings"}, + ) + + +def test_runtime_retention_update_rejects_invalid_values_and_applies_valid_ones_without_restart( + gateway: Gateway, tmp_path: Path +) -> None: + with ( + scratch_database() as database_url, + wire_server(_respond) as wire, + _isolated(gateway, _config(tmp_path, None), tmp_path, database_url) as isolated, + ): + model: Final = isolated.scenario.model(model="openai/gpt-4o-mini", api_base=wire.url, api_key="synthetic") + invalid: Final[tuple[JsonValue, ...]] = ( + {"include": ["usage_object"], "exclude": ["model_map_information"]}, + {}, + {"exclude": ["not_a_metadata_field"]}, + {"exclude": ["status"]}, + ) + + isolated.proxy.chat(model, key=isolated.key) + assert set(EXCLUDED) <= set(object_value(isolated.rows(1)[0]["metadata"])) + for value in invalid: + assert _update_retention(isolated.proxy, value).status_code == 400, value + isolated.proxy.chat(model, key=isolated.key) + assert set(EXCLUDED) <= set(object_value(isolated.rows(2)[-1]["metadata"])) + + assert _update_retention(isolated.proxy, {"exclude": list(EXCLUDED)}).status_code == 200 + isolated.proxy.chat(model, key=isolated.key) + _assert_filtered(isolated.rows(3)[-1]) + + for value in invalid: + assert _update_retention(isolated.proxy, value).status_code == 400, value + isolated.proxy.chat(model, key=isolated.key) + _assert_filtered(isolated.rows(4)[-1]) + + +SAVINGS_FIELDS: Final = ("autorouter_savings_estimate", "autorouter_savings") + + +def _anthropic_message(model: JsonValue) -> dict[str, JsonValue]: + return { + "id": f"msg_{uuid4().hex}", + "type": "message", + "role": "assistant", + "model": model, + "content": [{"type": "text", "text": "ok"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": { + "input_tokens": 12, + "output_tokens": 2, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 0}, + }, + } + + +def _anthropic(request: Request) -> Reply: + if request.target.startswith("/v1/messages/count_tokens"): + return Reply(body=b'{"input_tokens":12}') + body: Final = JSON_OBJECT.validate_json(request.body) + if body.get("stream") is True: + return Reply(content_type="text/event-stream", chunks=_anthropic_stream(body["model"])) + return Reply(body=json.dumps(_anthropic_message(body["model"])).encode()) + + +def _router_models(wire: Wire, router: str) -> tuple[Mapping[str, JsonValue], ...]: + tiers: Final = {"cheap": "anthropic/claude-sonnet-4-6", "frontier": "anthropic/claude-opus-4-8"} + return ( + *( + {"model_name": name, "litellm_params": {"model": model, "api_base": wire.url, "api_key": "synthetic"}} + for name, model in tiers.items() + ), + { + "model_name": router, + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": {"SIMPLE": "cheap", "MEDIUM": "cheap", "COMPLEX": "frontier", "REASONING": "frontier"}, + }, + }, + }, + ) + + +@pytest.mark.parametrize( + ("retention", "stored"), + [(None, SAVINGS_FIELDS), ({"exclude": list(SAVINGS_FIELDS)}, ())], + ids=["unset-keeps-savings", "excluded-savings-stay-out"], +) +@pytest.mark.timeout(240) +def test_delayed_autorouter_savings_publication_respects_retention( + gateway: Gateway, tmp_path: Path, retention: Mapping[str, JsonValue] | None, stored: tuple[str, ...] +) -> None: + router: Final = f"router-{uuid4().hex[:8]}" + with wire_server(_anthropic) as wire: + config: Final = _config(tmp_path, retention, model_list=_router_models(wire, router)) + with _isolated(gateway, config, tmp_path) as isolated: + response: Final = isolated.proxy.request( + "POST", + "/v1/messages", + {"model": router, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}, + key=isolated.key, + headers={"x-litellm-session-id": f"session-{uuid4()}"}, + ) + assert response.status_code == 200, response.text + + published: Final = isolated.rows( + 1, + "metadata::jsonb ? 'routing_decision' AND NOT metadata::jsonb ? 'autorouter_baseline_observation' " + 'AND request_id IN (SELECT request_id FROM "LiteLLM_AutoRouterBaselineObservation" ' + "WHERE publication IS NOT NULL)", + seconds=150, + )[0] + metadata: Final = object_value(published["metadata"]) + assert {name for name in SAVINGS_FIELDS if name in metadata} == set(stored), metadata diff --git a/tests/integration/streaming/test_responses_bridge_stream_chaos.py b/tests/integration/streaming/test_responses_bridge_stream_chaos.py index ca891584a8e..6b2ea47a849 100644 --- a/tests/integration/streaming/test_responses_bridge_stream_chaos.py +++ b/tests/integration/streaming/test_responses_bridge_stream_chaos.py @@ -230,7 +230,7 @@ def _assert_answered_in_its_own_shape(served: _Served) -> None: frame_error: Final = object_value(events[-1][1]["error"]) assert frame_error["type"] == "rate_limit_error", frame_error frame_message: Final = string_value(frame_error["message"]) - assert frame_message.count(_SENTINEL_PREFIX) == 1 and _RATE_LIMIT_PREFIX in frame_message, frame_message + assert frame_message.startswith(_RATE_LIMIT_PREFIX) and _SENTINEL_PREFIX not in frame_message, frame_message case "responses_limited": assert served.status == 200, served.text kinds: Final = [frame["type"] for frame in _data_frames(served.text)] diff --git a/tests/integration/streaming/test_responses_bridge_stream_errors.py b/tests/integration/streaming/test_responses_bridge_stream_errors.py index 9700658fa70..162ddc07bad 100644 --- a/tests/integration/streaming/test_responses_bridge_stream_errors.py +++ b/tests/integration/streaming/test_responses_bridge_stream_errors.py @@ -392,8 +392,8 @@ def assert_messages_errorframe(events: Sequence[tuple[str, Mapping[str, JsonValu error: Final = object_value(events[-1][1]["error"]) assert error["type"] == "rate_limit_error", error message: Final = string_value(error["message"]) - assert message.startswith(_SENTINEL_PREFIX + _RATE_LIMIT_PREFIX) and RATE_LIMIT_MESSAGE in message, message - assert message.count(_SENTINEL_PREFIX) == 1, message + assert message.startswith(_RATE_LIMIT_PREFIX) and RATE_LIMIT_MESSAGE in message, message + assert message.count(_RATE_LIMIT_PREFIX) == 1 and _SENTINEL_PREFIX not in message, message def test_messages_over_the_bridged_stream_carry_the_provider_error_once_in_the_errorframe(gateway: Gateway) -> None: diff --git a/tests/integration/translation/decisions/bases/cloudflare.py b/tests/integration/translation/decisions/bases/cloudflare.py index 1668dc215a8..ae86a2ccd87 100644 --- a/tests/integration/translation/decisions/bases/cloudflare.py +++ b/tests/integration/translation/decisions/bases/cloudflare.py @@ -6,7 +6,7 @@ from integration.translation.case import TranslationTestCase """ CLEF_TEST_CASE: Final = TranslationTestCase( scenario="basic", - litellm_endpoint="/v1/decisions", + litellm_endpoint="/v1/systemone", litellm_request={ "model": "cloudflare/@cf/cloudflare/clef", "state": "Ticket (billing): The export job hangs at 99% and never finishes", diff --git a/tests/integration/translation/decisions/bases/openrouter.py b/tests/integration/translation/decisions/bases/openrouter.py index f6217bcf818..b65acff6b06 100644 --- a/tests/integration/translation/decisions/bases/openrouter.py +++ b/tests/integration/translation/decisions/bases/openrouter.py @@ -6,7 +6,7 @@ from integration.translation.case import TranslationTestCase """ TYPESAFE_JEV_1_13_TEST_CASE: Final = TranslationTestCase( scenario="basic", - litellm_endpoint="/v1/decisions", + litellm_endpoint="/v1/systemone", litellm_request={ "model": "openrouter/typesafe/jev-1.13", "state": "Ticket (billing): The export job hangs at 99% and never finishes", diff --git a/tests/integration/translation/decisions/bases/perplexity.py b/tests/integration/translation/decisions/bases/perplexity.py index 71ef7a85220..02b1de51f24 100644 --- a/tests/integration/translation/decisions/bases/perplexity.py +++ b/tests/integration/translation/decisions/bases/perplexity.py @@ -6,7 +6,7 @@ from integration.translation.case import TranslationTestCase """ PPLX_DECIDER_V1_27B_TEST_CASE: Final = TranslationTestCase( scenario="basic", - litellm_endpoint="/v1/decisions", + litellm_endpoint="/v1/systemone", litellm_request={ "model": "perplexity/pplx-decider-v1-27b", "state": "Ticket (billing): The export job hangs at 99% and never finishes", diff --git a/tests/integration/translation/decisions/bases/strands_decider.py b/tests/integration/translation/decisions/bases/strands_decider.py index 83d7ae5e33a..8a1a51f8c3e 100644 --- a/tests/integration/translation/decisions/bases/strands_decider.py +++ b/tests/integration/translation/decisions/bases/strands_decider.py @@ -6,7 +6,7 @@ from integration.translation.case import TranslationTestCase """ STRANDS_DECIDER_2B_HOBSON_V19_TEST_CASE: Final = TranslationTestCase( scenario="basic", - litellm_endpoint="/v1/decisions", + litellm_endpoint="/v1/systemone", litellm_request={ "model": "strands_decider/strands-decider-2B-hobson-v19", "state": "Help! My payouts have been failing for 3 days!", diff --git a/tests/integration/translation/decisions/bases/typesafe.py b/tests/integration/translation/decisions/bases/typesafe.py index 2fb384a94ae..5da5b05cb09 100644 --- a/tests/integration/translation/decisions/bases/typesafe.py +++ b/tests/integration/translation/decisions/bases/typesafe.py @@ -6,7 +6,7 @@ from integration.translation.case import TranslationTestCase """ JEV_1_13_0_TEST_CASE: Final = TranslationTestCase( scenario="basic", - litellm_endpoint="/v1/decisions", + litellm_endpoint="/v1/systemone", litellm_request={ "model": "typesafe/jev-1.13.0", "state": "Ticket (billing): The export job hangs at 99% and never finishes", diff --git a/tests/litellm_utils_tests/base_token_counter_test.py b/tests/litellm_utils_tests/base_token_counter_test.py index ddce27522c2..f4c90118a03 100644 --- a/tests/litellm_utils_tests/base_token_counter_test.py +++ b/tests/litellm_utils_tests/base_token_counter_test.py @@ -17,7 +17,6 @@ import pytest from litellm.llms.base_llm.base_utils import BaseTokenCounter -from litellm.types.utils import TokenCountResponse class BaseTokenCounterTest(ABC): @@ -71,69 +70,3 @@ class BaseTokenCounterTest(ABC): ): pytest.skip(f"Missing or invalid credentials: {e}") raise - - @pytest.mark.asyncio - async def test_count_tokens_basic(self): - """ - Test basic token counting functionality. - - Verifies that: - - Token counter returns a TokenCountResponse - - total_tokens is greater than 0 - - tokenizer_type is set - - No error occurred - """ - token_counter = self.get_token_counter() - model = self.get_test_model() - messages = self.get_test_messages() - deployment = self.get_deployment_config() - - result = await token_counter.count_tokens( - model_to_use=model, - messages=messages, - contents=None, - deployment=deployment, - request_model=model, - ) - - print(f"Token count result: {result}") - - assert result is not None, "Token counter should return a result" - assert isinstance( - result, TokenCountResponse - ), "Result should be TokenCountResponse" - assert ( - result.total_tokens > 0 - ), f"Token count should be > 0, got {result.total_tokens}" - assert result.tokenizer_type is not None, "tokenizer_type should be set" - assert ( - result.error is not True - ), f"Token counting should not error: {result.error_message}" - - def test_should_use_token_counting_api(self): - """ - Test that should_use_token_counting_api returns True for the correct provider. - - Verifies that the token counter correctly identifies when it should be used - based on the custom_llm_provider. - """ - token_counter = self.get_token_counter() - provider = self.get_custom_llm_provider() - - result = token_counter.should_use_token_counting_api( - custom_llm_provider=provider - ) - - assert ( - result is True - ), f"should_use_token_counting_api should return True for {provider}" - - # Also verify it returns False for other providers - other_provider = "some_other_provider_that_doesnt_exist" - result_other = token_counter.should_use_token_counting_api( - custom_llm_provider=other_provider - ) - - assert ( - result_other is False - ), f"should_use_token_counting_api should return False for {other_provider}" diff --git a/tests/litellm_utils_tests/test_bedrock_token_counter.py b/tests/litellm_utils_tests/test_bedrock_token_counter.py deleted file mode 100644 index b4c05cb0cd7..00000000000 --- a/tests/litellm_utils_tests/test_bedrock_token_counter.py +++ /dev/null @@ -1,107 +0,0 @@ -""" -Bedrock Token Counter Tests. - -Tests for the Bedrock token counter implementation using the base test suite. - -Note: Not all Bedrock models support token counting. The CountTokens API -is only available for specific models. If the model doesn't support token -counting, the test will be skipped. -""" - -import os -from typing import Any, Dict, List - -import pytest - - -from litellm.llms.base_llm.base_utils import BaseTokenCounter -from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter -from tests.litellm_utils_tests.base_token_counter_test import BaseTokenCounterTest - - -class TestBedrockTokenCounter(BaseTokenCounterTest): - """Test suite for Bedrock token counter. - - Note: Bedrock CountTokens API support varies by model. Some models - (like older Claude versions) may not support token counting. - Use amazon.nova-* models for reliable token counting support. - """ - - def get_token_counter(self) -> BaseTokenCounter: - return BedrockTokenCounter() - - def get_test_model(self) -> str: - # Use Amazon Nova model which supports token counting - # Alternatively, use environment variable to override - return os.getenv("BEDROCK_TEST_MODEL", "amazon.nova-lite-v1:0") - - def get_test_messages(self) -> List[Dict[str, Any]]: - return [{"role": "user", "content": "Hello, how are you today?"}] - - def get_deployment_config(self) -> Dict[str, Any]: - # Bedrock uses AWS credentials from environment - # Check for AWS credentials - aws_access_key = os.getenv("AWS_ACCESS_KEY_ID") - aws_secret_key = os.getenv("AWS_SECRET_ACCESS_KEY") - aws_region = os.getenv("AWS_REGION_NAME", "us-east-1") - - if not aws_access_key or not aws_secret_key: - pytest.skip( - "AWS credentials not set (AWS_ACCESS_KEY_ID, AWS_SECRET_ACCESS_KEY)" - ) - - return { - "litellm_params": { - "aws_access_key_id": aws_access_key, - "aws_secret_access_key": aws_secret_key, - "aws_region_name": aws_region, - } - } - - def get_custom_llm_provider(self) -> str: - return "bedrock" - - @pytest.mark.asyncio - async def test_count_tokens_basic(self): - """ - Test basic token counting functionality. - - Override to handle models that don't support token counting. - """ - from litellm.types.utils import TokenCountResponse - - token_counter = self.get_token_counter() - model = self.get_test_model() - messages = self.get_test_messages() - deployment = self.get_deployment_config() - - result = await token_counter.count_tokens( - model_to_use=model, - messages=messages, - contents=None, - deployment=deployment, - request_model=model, - ) - - print(f"Token count result: {result}") - - assert result is not None, "Token counter should return a result" - assert isinstance( - result, TokenCountResponse - ), "Result should be TokenCountResponse" - - # Check if the model doesn't support token counting - if result.error and "doesn't support counting tokens" in str( - result.error_message - ): - pytest.skip( - f"Model {model} doesn't support token counting: {result.error_message}" - ) - - assert ( - result.total_tokens > 0 - ), f"Token count should be > 0, got {result.total_tokens}" - assert result.tokenizer_type is not None, "tokenizer_type should be set" - assert ( - result.error is not True - ), f"Token counting should not error: {result.error_message}" diff --git a/tests/litellm_utils_tests/test_health_check.py b/tests/litellm_utils_tests/test_health_check.py deleted file mode 100644 index 9a79db89164..00000000000 --- a/tests/litellm_utils_tests/test_health_check.py +++ /dev/null @@ -1,326 +0,0 @@ -#### What this tests #### -# This tests if ahealth_check() actually works - -import asyncio -import os -from unittest.mock import AsyncMock, patch - -import pytest - -import litellm - - -@pytest.mark.asyncio -async def test_azure_health_check(): - response = await litellm.ahealth_check( - model_params={ - "model": "azure/gpt-4.1-mini", - "messages": [{"role": "user", "content": "Hey, how's it going?"}], - "api_key": os.getenv("AZURE_AI_API_KEY"), - "api_base": os.getenv("AZURE_AI_API_BASE"), - "api_version": os.getenv("AZURE_AI_API_VERSION"), - } - ) - print(f"response: {response}") - - assert "x-ratelimit-remaining-tokens" in response - return response - - -# asyncio.run(test_azure_health_check()) - - - - -@pytest.mark.asyncio -async def test_azure_embedding_health_check(): - response = await litellm.ahealth_check( - model_params={ - "model": "azure/text-embedding-ada-002", - "api_key": os.getenv("AZURE_AI_API_KEY"), - "api_base": os.getenv("AZURE_AI_API_BASE"), - "api_version": os.getenv("AZURE_AI_API_VERSION"), - }, - input=["test for litellm"], - mode="embedding", - ) - print(f"response: {response}") - - assert "x-ratelimit-remaining-tokens" in response - return response - - -@pytest.mark.asyncio -async def test_openai_img_gen_health_check(): - response = await litellm.ahealth_check( - model_params={ - "model": "gpt-image-1", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - mode="image_generation", - prompt="cute baby sea otter", - ) - print(f"response: {response}") - - assert isinstance(response, dict) and "error" not in response - return response - - -# asyncio.run(test_openai_img_gen_health_check()) - - - - - - -# asyncio.run(test_sagemaker_embedding_health_check()) - - -@pytest.mark.asyncio -async def test_groq_health_check(): - """ - This should not fail - - ensure that provider wildcard model passes health check - """ - litellm.set_verbose = True - response = await litellm.ahealth_check( - model_params={ - "api_key": os.environ.get("GROQ_API_KEY"), - "model": "groq/*", - "messages": [{"role": "user", "content": "What's 1 + 1?"}], - }, - mode=None, - prompt="What's 1 + 1?", - input=["test from litellm"], - ) - print(f"response: {response}") - assert response == {} - - return response - - -@pytest.mark.asyncio -async def test_cohere_rerank_health_check(): - response = await litellm.ahealth_check( - model_params={ - "model": "cohere/rerank-english-v3.0", - "api_key": os.getenv("COHERE_API_KEY"), - }, - mode="rerank", - prompt="Hey, how's it going", - ) - - assert "error" not in response - - print(response) - - -@pytest.mark.asyncio -async def test_audio_speech_health_check(): - response = await litellm.ahealth_check( - model_params={ - "model": "openai/tts-1", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - mode="audio_speech", - prompt="Hey", - ) - - assert "error" not in response - - print(response) - - -@pytest.mark.asyncio -async def test_audio_speech_health_check_with_another_voice(): - response = await litellm.ahealth_check( - model_params={ - "model": "openai/tts-1", - "api_key": os.getenv("OPENAI_API_KEY"), - "health_check_voice": "en-US-JennyNeural", - }, - mode="audio_speech", - prompt="Hey", - ) - - assert "error" not in response - - print(response) - - -@pytest.mark.asyncio -async def test_audio_transcription_health_check(): - litellm.set_verbose = True - response = await litellm.ahealth_check( - model_params={ - "model": "openai/whisper-1", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - mode="audio_transcription", - ) - - print(f"response: {response}") - - assert "error" not in response - - print(response) - - - - - - - - - - -@pytest.mark.asyncio -async def test_health_check_bad_model(): - import time - - from litellm.proxy.health_check import _perform_health_check - - model_list = [ - { - "model_name": "openai-gpt-4o", - "litellm_params": { - "api_key": "sk-9876", - "api_base": "https://exampleopenaiendpoint-production.up.railway.app", - "model": "openai/my-fake-openai-endpoint", - "mock_timeout": True, - "timeout": 60, - }, - "model_info": { - "id": "ca27ca2eeea2f9e38bb274ead831948a26621a3738d06f1797253f0e6c4278c0", - "db_model": False, - "health_check_timeout": 1, - }, - }, - ] - details = None - healthy_endpoints, unhealthy_endpoints, _ = await _perform_health_check( - model_list, details - ) - print(f"healthy_endpoints: {healthy_endpoints}") - print(f"unhealthy_endpoints: {unhealthy_endpoints}") - - # Track which model is actually used in the health check - health_check_calls = [] - - async def mock_health_check(litellm_params, **kwargs): - health_check_calls.append(litellm_params["model"]) - await asyncio.sleep(10) - return {"status": "healthy"} - - with patch( - "litellm.ahealth_check", side_effect=mock_health_check - ) as mock_health_check: - start_time = time.time() - healthy_endpoints, unhealthy_endpoints, _ = await _perform_health_check( - model_list - ) - end_time = time.time() - print("health check calls: ", health_check_calls) - assert len(healthy_endpoints) == 0 - assert len(unhealthy_endpoints) == 1 - assert ( - end_time - start_time < 2 - ), "Health check took longer than health_check_timeout" - - -@pytest.mark.asyncio -async def test_health_check_respects_concurrency_limit(): - from litellm.proxy.health_check import _perform_health_check - - model_list = [ - {"litellm_params": {"model": f"openai/gpt-4o-mini-{i}", "api_key": "fake-key"}} - for i in range(6) - ] - - active = 0 - max_active = 0 - - async def mock_health_check(litellm_params, **kwargs): - nonlocal active, max_active - active += 1 - max_active = max(max_active, active) - await asyncio.sleep(0.05) - active -= 1 - return {"status": "healthy"} - - with patch("litellm.ahealth_check", side_effect=mock_health_check): - await _perform_health_check(model_list, max_concurrency=2) - - assert max_active <= 2 - - -@pytest.mark.asyncio -async def test_health_check_creates_only_bounded_initial_tasks(): - from litellm.proxy.health_check import _perform_health_check - - model_list = [ - {"litellm_params": {"model": f"openai/gpt-4o-mini-{i}", "api_key": "fake-key"}} - for i in range(10) - ] - release_event = asyncio.Event() - create_task_call_count = 0 - real_create_task = asyncio.create_task - - async def mock_health_check(litellm_params, **kwargs): - await release_event.wait() - return {"status": "healthy"} - - def tracked_create_task(coro): - nonlocal create_task_call_count - create_task_call_count += 1 - return real_create_task(coro) - - with ( - patch("litellm.ahealth_check", side_effect=mock_health_check), - patch( - "litellm.proxy.health_check.asyncio.create_task", - side_effect=tracked_create_task, - ), - ): - perform_task = real_create_task( - _perform_health_check(model_list, max_concurrency=2) - ) - await asyncio.sleep(0.05) - assert create_task_call_count == 2 - release_event.set() - await perform_task - - -@pytest.mark.asyncio -async def test_timeout_does_not_cancel_other_health_checks(): - from litellm.proxy.health_check import _perform_health_check - - model_list = [ - { - "litellm_params": {"model": "openai/slow-model", "api_key": "fake-key"}, - "model_info": {"health_check_timeout": 0.05}, - }, - { - "litellm_params": {"model": "openai/fast-model", "api_key": "fake-key"}, - "model_info": {"health_check_timeout": 1}, - }, - ] - - async def mock_health_check(litellm_params, **kwargs): - if litellm_params["model"] == "openai/slow-model": - await asyncio.sleep(0.2) - return {"status": "healthy"} - await asyncio.sleep(0.01) - return {"status": "healthy"} - - with patch("litellm.ahealth_check", side_effect=mock_health_check): - healthy_endpoints, unhealthy_endpoints, _ = await _perform_health_check( - model_list, max_concurrency=1 - ) - - healthy_models = {endpoint["model"] for endpoint in healthy_endpoints} - unhealthy_models = {endpoint["model"] for endpoint in unhealthy_endpoints} - - assert "openai/fast-model" in healthy_models - assert "openai/slow-model" in unhealthy_models diff --git a/tests/litellm_utils_tests/test_secret_manager.py b/tests/litellm_utils_tests/test_secret_manager.py index 7cc5c72faa7..6d088b0c0e2 100644 --- a/tests/litellm_utils_tests/test_secret_manager.py +++ b/tests/litellm_utils_tests/test_secret_manager.py @@ -1,4 +1,3 @@ -import base64 import hashlib import json import os @@ -7,7 +6,6 @@ from dotenv import load_dotenv load_dotenv() import tempfile -from unittest.mock import MagicMock, patch from typing import Final import pytest @@ -93,8 +91,6 @@ def test_oidc_google(): ) print(f"secret_val: {redact_oidc_signature(secret_val)}") - - @pytest.mark.skipif( os.environ.get("ACTIONS_ID_TOKEN_REQUEST_TOKEN") is None, reason="Cannot run without being in GitHub Actions", @@ -105,112 +101,3 @@ def test_oidc_github(): ) print(f"secret_val: {redact_oidc_signature(secret_val)}") - - -@pytest.mark.skipif( - os.environ.get("CIRCLE_OIDC_TOKEN") is None, - reason="Cannot run without being in CircleCI Runner", -) -def test_oidc_circleci(): - secret_val = get_secret("oidc/circleci/") - - print(f"secret_val: {redact_oidc_signature(secret_val)}") - - -@pytest.mark.skipif( - os.environ.get("CIRCLE_OIDC_TOKEN_V2") is None, - reason="Cannot run without being in CircleCI Runner", -) -def test_oidc_circleci_v2(): - secret_val = get_secret( - "oidc/circleci_v2/https://bedrock-runtime.us-east-1.amazonaws.com/model/amazon.titan-text-express-v1/invoke" - ) - - print(f"secret_val: {redact_oidc_signature(secret_val)}") - - - - - - - - - - - - -def test_google_secret_manager(): - """ - Test that we can get a secret from Google Secret Manager - """ - os.environ["GOOGLE_SECRET_MANAGER_PROJECT_ID"] = "litellm-ci-cd" - - from litellm.secret_managers.google_secret_manager import GoogleSecretManager - - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.json.return_value = { - "payload": { - "data": base64.b64encode(b"anything").decode("utf-8"), - } - } - - with ( - patch("litellm.proxy.proxy_server.premium_user", True), - patch.object( - GoogleSecretManager, - "sync_construct_request_headers", - return_value={"Authorization": "Bearer mock_token"}, - ), - ): - secret_manager = GoogleSecretManager() - secret_manager.sync_httpx_client = MagicMock() - secret_manager.sync_httpx_client.get.return_value = mock_response - - secret_val = secret_manager.get_secret_from_google_secret_manager( - secret_name="OPENAI_API_KEY" - ) - print("secret_val: {}".format(secret_val)) - - assert ( - secret_val == "anything" - ), "did not get expected secret value. expect 'anything', got '{}'".format( - secret_val - ) - - secret_manager.sync_httpx_client.get.assert_called_once() - call_url = secret_manager.sync_httpx_client.get.call_args[1]["url"] - assert "projects/litellm-ci-cd/secrets/OPENAI_API_KEY" in call_url - - -def test_google_secret_manager_read_in_memory(): - """ - Test that Google Secret manager returns in memory value when it exists - """ - from litellm.secret_managers.google_secret_manager import GoogleSecretManager - - os.environ["GOOGLE_SECRET_MANAGER_PROJECT_ID"] = "litellm-ci-cd" - - with ( - patch("litellm.proxy.proxy_server.premium_user", True), - patch.object( - GoogleSecretManager, - "sync_construct_request_headers", - return_value={"Authorization": "Bearer mock_token"}, - ), - ): - secret_manager = GoogleSecretManager() - secret_manager.cache.cache_dict["UNIQUE_KEY"] = None - secret_manager.cache.cache_dict["UNIQUE_KEY_2"] = "lite-llm" - - secret_val = secret_manager.get_secret_from_google_secret_manager( - secret_name="UNIQUE_KEY" - ) - print("secret_val: {}".format(secret_val)) - assert secret_val is None - - secret_val = secret_manager.get_secret_from_google_secret_manager( - secret_name="UNIQUE_KEY_2" - ) - print("secret_val: {}".format(secret_val)) - assert secret_val == "lite-llm" diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py index de29496a79b..5ec7b96dcfa 100644 --- a/tests/litellm_utils_tests/test_utils.py +++ b/tests/litellm_utils_tests/test_utils.py @@ -1,7 +1,5 @@ import copy import logging -import time -from datetime import datetime from unittest import mock from dotenv import load_dotenv @@ -15,22 +13,14 @@ import pytest import litellm from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, headers -from litellm.litellm_core_utils.duration_parser import duration_in_seconds -from litellm.litellm_core_utils.duration_parser import ( - get_last_day_of_month, - _extract_from_regex, -) from litellm.utils import ( - check_valid_key, get_llm_provider, get_supported_openai_params, get_token_count, - get_valid_models, trim_messages, validate_environment, ) from litellm.llms.openai_like.json_loader import JSONProviderRegistry -from unittest.mock import AsyncMock, MagicMock, patch # Assuming your trim_messages, shorten_message_to_fit_limit, and get_token_count functions are all in a module named 'message_utils' @@ -47,676 +37,29 @@ def reset_mock_cache(): # test_basic_trimming() - - # test_basic_trimming_no_max_tokens_specified() - - # test_multiple_messages_trimming() - - # test_multiple_messages_no_trimming() - - # test_large_trimming() - - - - - - - - - - - - - - - - -@pytest.mark.parametrize("custom_llm_provider", ["anthropic", "xai"]) -def test_get_valid_models_with_custom_llm_provider(custom_llm_provider): - from litellm.utils import ProviderConfigManager - from litellm.types.utils import LlmProviders - - provider_config = ProviderConfigManager.get_provider_model_info( - model=None, - provider=LlmProviders(custom_llm_provider), - ) - assert provider_config is not None - valid_models = get_valid_models( - check_provider_endpoint=True, custom_llm_provider=custom_llm_provider - ) - print(valid_models) - assert len(valid_models) > 0 - assert set(provider_config.get_models()) == set(valid_models) - - # test_get_valid_models() -def test_bad_key(): - key = "bad-key" - response = check_valid_key(model="gpt-5-mini", api_key=key) - print(response, key) - assert response == False - - -def test_good_key(): - key = os.environ["OPENAI_API_KEY"] - response = check_valid_key(model="gpt-5-mini", api_key=key) - assert response == True - - # test validate environment - - - - - - - - - - - - -def test_function_to_dict(): - print("testing function to dict for get current weather") - - def get_current_weather(location: str, unit: str): - """Get the current weather in a given location - - Parameters - ---------- - location : str - The city and state, e.g. San Francisco, CA - unit : {'celsius', 'fahrenheit'} - Temperature unit - - Returns - ------- - str - a sentence indicating the weather - """ - if location == "Boston, MA": - return "The weather is 12F" - - function_json = litellm.utils.function_to_dict(get_current_weather) - print(function_json) - - expected_output = { - "name": "get_current_weather", - "description": "Get the current weather in a given location", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "The city and state, e.g. San Francisco, CA", - }, - "unit": { - "type": "string", - "description": "Temperature unit", - "enum": "['fahrenheit', 'celsius']", - }, - }, - "required": ["location", "unit"], - }, - } - print(expected_output) - - assert function_json["name"] == expected_output["name"] - assert function_json["description"] == expected_output["description"] - assert function_json["parameters"]["type"] == expected_output["parameters"]["type"] - assert ( - function_json["parameters"]["properties"]["location"] - == expected_output["parameters"]["properties"]["location"] - ) - - # the enum can change it can be - which is why we don't assert on unit - # {'type': 'string', 'description': 'Temperature unit', 'enum': "['fahrenheit', 'celsius']"} - # {'type': 'string', 'description': 'Temperature unit', 'enum': "['celsius', 'fahrenheit']"} - - assert ( - function_json["parameters"]["required"] - == expected_output["parameters"]["required"] - ) - - print("passed") - - -# test_function_to_dict() - - - - - - - - - - - - - - -def test_duration_in_seconds(): - """ - Test if duration int is correctly calculated for different str - """ - import time - - now = time.time() - current_time = datetime.fromtimestamp(now) - - if current_time.month == 12: - target_year = current_time.year + 1 - target_month = 1 - else: - target_year = current_time.year - target_month = current_time.month + 1 - - # Determine the day to set for next month - target_day = current_time.day - last_day_of_target_month = get_last_day_of_month(target_year, target_month) - - if target_day > last_day_of_target_month: - target_day = last_day_of_target_month - - next_month = datetime( - year=target_year, - month=target_month, - day=target_day, - hour=current_time.hour, - minute=current_time.minute, - second=current_time.second, - microsecond=current_time.microsecond, - ) - - # Calculate the duration until the first day of the next month - duration_until_next_month = next_month - current_time - expected_duration = int(duration_until_next_month.total_seconds()) - - value = duration_in_seconds(duration="1mo") - - assert value - expected_duration < 2 - - - - - - -@pytest.mark.parametrize("langfuse_trace_id", [None, "my-unique-trace-id"]) -@pytest.mark.parametrize( - "langfuse_existing_trace_id", [None, "my-unique-existing-trace-id"] -) -def test_logging_trace_id(langfuse_trace_id, langfuse_existing_trace_id): - """ - - Unit test for `_get_trace_id` function in Logging obj - """ - from litellm.integrations.langfuse.langfuse_sdk import resolve_trace_id - from litellm.litellm_core_utils.litellm_logging import Logging - - litellm.success_callback = ["langfuse"] - litellm_call_id = "my-unique-call-id" - litellm_logging_obj = Logging( - model="gpt-5-mini", - messages=[{"role": "user", "content": "hi"}], - stream=False, - call_type="acompletion", - litellm_call_id=litellm_call_id, - start_time=datetime.now(), - function_id="1234", - ) - - metadata = {} - - if langfuse_trace_id is not None: - metadata["trace_id"] = langfuse_trace_id - if langfuse_existing_trace_id is not None: - metadata["existing_trace_id"] = langfuse_existing_trace_id - - litellm.completion( - model="gpt-5-mini", - messages=[{"role": "user", "content": "Hey how's it going?"}], - mock_response="Hey!", - litellm_logging_obj=litellm_logging_obj, - metadata=metadata, - ) - - time.sleep(3) - assert litellm_logging_obj.get_trace_id(service_name="langfuse") is not None - - # langfuse addresses a trace by a 32-hex id, so the id litellm reports back is the - # resolved form of whichever source won; that is what the alerting deep link needs - if langfuse_existing_trace_id is not None: - expected_source = langfuse_existing_trace_id - elif langfuse_trace_id is not None: - expected_source = langfuse_trace_id - else: - expected_source = litellm_logging_obj.litellm_trace_id - - assert litellm_logging_obj.get_trace_id(service_name="langfuse") == resolve_trace_id( - expected_source - ) - - - - - - - - -def test_is_base64_encoded(): - import base64 - - import requests - - litellm.set_verbose = True - url = "https://dummyimage.com/100/100/fff&text=Test+image" - response = requests.get(url) - file_data = response.content - - encoded_file = base64.b64encode(file_data).decode("utf-8") - base64_image = f"data:image/png;base64,{encoded_file}" - - from litellm.utils import is_base64_encoded - - assert is_base64_encoded(s=base64_image) is True - - - - - - - - - - - - - - - - - - - - - - -def test_is_prompt_caching_enabled_return_default_image_dimensions(): - """ - Assert that `is_prompt_caching_valid_prompt` counts tokens with use_default_image_token_count=True - when processing messages containing images - - IMPORTANT: Ensures Get token counter does not make a GET request to the image url - """ - mock_token_counter = MagicMock(return_value=False) - with patch( - "litellm.utils.get_messages_reach_token_count", - return_value=mock_token_counter, - ): - litellm.utils.is_prompt_caching_valid_prompt( - messages=[ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is in this image?"}, - { - "type": "image_url", - "image_url": { - "url": "https://www.gstatic.com/webp/gallery/1.webp", - "detail": "high", - }, - }, - ], - } - ], - tools=None, - custom_llm_provider="openai", - model="gpt-4o-mini", - ) - - # Assert token_counter was called with use_default_image_token_count=True - args_to_mock_token_counter = mock_token_counter.call_args[1] - print("args_to_mock", args_to_mock_token_counter) - assert args_to_mock_token_counter["use_default_image_token_count"] is True - - - - - - - - -def test_get_valid_models_fireworks_ai(monkeypatch): - from litellm.utils import get_valid_models - import litellm - - litellm.turn_on_debug() - - monkeypatch.setenv("FIREWORKS_API_KEY", "sk-9876") - monkeypatch.setenv("FIREWORKS_ACCOUNT_ID", "1234") - monkeypatch.setattr(litellm, "provider_list", ["fireworks_ai"]) - - mock_response_data = { - "models": [ - { - "name": "accounts/fireworks/models/llama-3.1-8b-instruct", - "displayName": "", - "description": "", - "createTime": "2023-11-07T05:31:56Z", - "createdBy": "", - "state": "STATE_UNSPECIFIED", - "status": {"code": "OK", "message": ""}, - "kind": "KIND_UNSPECIFIED", - "githubUrl": "", - "huggingFaceUrl": "", - "baseModelDetails": { - "worldSize": 123, - "checkpointFormat": "CHECKPOINT_FORMAT_UNSPECIFIED", - "parameterCount": "", - "moe": True, - "tunable": True, - }, - "peftDetails": { - "baseModel": "", - "r": 123, - "targetModules": [""], - }, - "teftDetails": {}, - "public": True, - "conversationConfig": { - "style": "", - "system": "", - "template": "", - }, - "contextLength": 123, - "supportsImageInput": True, - "supportsTools": True, - "importedFrom": "", - "fineTuningJob": "", - "defaultDraftModel": "", - "defaultDraftTokenCount": 123, - "precisions": ["PRECISION_UNSPECIFIED"], - "deployedModelRefs": [ - { - "name": "", - "deployment": "", - "state": "STATE_UNSPECIFIED", - "default": True, - "public": True, - } - ], - "cluster": "", - "deprecationDate": {"year": 123, "month": 123, "day": 123}, - } - ], - "nextPageToken": "", - "totalSize": 123, - } - - # Create a mock response object - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.json.return_value = mock_response_data - - with patch.object( - litellm.module_level_client, "get", return_value=mock_response - ) as mock_post: - valid_models = get_valid_models(check_provider_endpoint=True) - print("valid_models", valid_models) - mock_post.assert_called_once() - assert ( - "fireworks_ai/accounts/fireworks/models/llama-3.1-8b-instruct" - in valid_models - ) - - -def test_get_valid_models_default(monkeypatch): - """ - Ensure that the default models is used when error retrieving from model api. - - Prevent regression for existing usage. - """ - from litellm.utils import get_valid_models - - monkeypatch.setenv("FIREWORKS_API_KEY", "sk-9876") - valid_models = get_valid_models() - assert len(valid_models) > 0 - - - - - - -def test_add_custom_logger_callback_to_specific_event(monkeypatch): - from litellm.utils import add_custom_logger_callback_to_specific_event - - monkeypatch.setattr(litellm, "success_callback", []) - monkeypatch.setattr(litellm, "failure_callback", []) - - add_custom_logger_callback_to_specific_event("langfuse", "success") - - assert len(litellm.success_callback) == 1 - assert len(litellm.failure_callback) == 0 - - - - - - -@pytest.mark.asyncio -async def test_add_custom_logger_callback_to_specific_event_with_duplicates( - monkeypatch, -): - """ - Test that when a callback exists in both success_callback and _async_success_callback, - it's not added again - """ - from litellm.integrations.langfuse.langfuse_prompt_management import ( - LangfusePromptManagement, - ) - - # Reset all callback lists - monkeypatch.setattr(litellm, "callbacks", []) - monkeypatch.setattr(litellm, "_async_success_callback", []) - monkeypatch.setattr(litellm, "_async_failure_callback", []) - monkeypatch.setattr(litellm, "success_callback", []) - monkeypatch.setattr(litellm, "failure_callback", []) - - # Add logger to both success_callback and _async_success_callback - langfuse_logger = LangfusePromptManagement() - litellm.success_callback.append(langfuse_logger) - litellm._async_success_callback.append(langfuse_logger) - - # Get initial lengths - initial_success_callback_len = len(litellm.success_callback) - initial_async_success_callback_len = len(litellm._async_success_callback) - - # Make a completion call - await litellm.acompletion( - model="gpt-5-mini", - messages=[{"role": "user", "content": "Hello, world!"}], - mock_response="Testing duplicate callbacks", - ) - - # Assert no new callbacks were added - assert len(litellm.success_callback) == initial_success_callback_len - assert len(litellm._async_success_callback) == initial_async_success_callback_len - - -@pytest.mark.asyncio -async def test_add_custom_logger_callback_to_specific_event_with_duplicates_success_callback( - monkeypatch, -): - """ - Test that when a callback exists in both success_callback and _async_success_callback, - it's not added again - """ - from litellm.integrations.langfuse.langfuse_prompt_management import ( - LangfusePromptManagement, - ) - - # Reset all callback lists - monkeypatch.setattr(litellm, "callbacks", []) - monkeypatch.setattr(litellm, "_async_success_callback", []) - monkeypatch.setattr(litellm, "_async_failure_callback", []) - monkeypatch.setattr(litellm, "success_callback", []) - monkeypatch.setattr(litellm, "failure_callback", []) - - # Add logger to both success_callback and _async_success_callback - langfuse_logger = LangfusePromptManagement() - litellm.success_callback.append(langfuse_logger) - - # Get initial lengths - initial_success_callback_len = len(litellm.success_callback) - initial_async_success_callback_len = len(litellm._async_success_callback) - - # Make a completion call - await litellm.acompletion( - model="gpt-5-mini", - messages=[{"role": "user", "content": "Hello, world!"}], - mock_response="Testing duplicate callbacks", - ) - - # Assert no new callbacks were added - assert len(litellm.success_callback) == initial_success_callback_len - assert len(litellm._async_success_callback) == initial_async_success_callback_len - - -@pytest.mark.asyncio -async def test_add_custom_logger_callback_to_specific_event_with_duplicates_callbacks( - monkeypatch, -): - """ - Test that when a callback exists in both success_callback and _async_success_callback, - it's not added again - """ - from litellm.integrations.langfuse.langfuse_prompt_management import ( - LangfusePromptManagement, - ) - - # Reset all callback lists - monkeypatch.setattr(litellm, "callbacks", []) - monkeypatch.setattr(litellm, "_async_success_callback", []) - monkeypatch.setattr(litellm, "_async_failure_callback", []) - monkeypatch.setattr(litellm, "success_callback", []) - monkeypatch.setattr(litellm, "failure_callback", []) - - # Add logger to both success_callback and _async_success_callback - langfuse_logger = LangfusePromptManagement() - litellm.callbacks.append(langfuse_logger) - - # Make a completion call - await litellm.acompletion( - model="gpt-5-mini", - messages=[{"role": "user", "content": "Hello, world!"}], - mock_response="Testing duplicate callbacks", - ) - - # Assert no new callbacks were added - initial_callbacks_len = len(litellm.callbacks) - initial_async_success_callback_len = len(litellm._async_success_callback) - initial_success_callback_len = len(litellm.success_callback) - print( - f"Num callbacks before: litellm.callbacks: {len(litellm.callbacks)}, litellm._async_success_callback: {len(litellm._async_success_callback)}, litellm.success_callback: {len(litellm.success_callback)}" - ) - - for _ in range(10): - await litellm.acompletion( - model="gpt-5-mini", - messages=[{"role": "user", "content": "Hello, world!"}], - mock_response="Testing duplicate callbacks", - ) - - assert len(litellm.callbacks) == initial_callbacks_len - assert len(litellm._async_success_callback) == initial_async_success_callback_len - assert len(litellm.success_callback) == initial_success_callback_len - - print( - f"Num callbacks after 10 mock calls: litellm.callbacks: {len(litellm.callbacks)}, litellm._async_success_callback: {len(litellm._async_success_callback)}, litellm.success_callback: {len(litellm.success_callback)}" - ) - - - - - - - - - - from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.utils import get_applied_guardrails from unittest.mock import Mock - - - - - - - - - - -def test_get_provider_audio_transcription_config(): - from litellm.utils import ProviderConfigManager - from litellm.types.utils import LlmProviders - - for provider in LlmProviders: - config = ProviderConfigManager.get_provider_audio_transcription_config( - model="whisper-1", provider=provider - ) - - - - - - - - -def test_get_valid_models_from_dynamic_api_key(): - """ - Test that get_valid_models returns the correct models for a given provider - """ - from litellm.utils import get_valid_models - from litellm.types.router import CredentialLiteLLMParams - - creds = CredentialLiteLLMParams(api_key="123") - - valid_models = get_valid_models( - custom_llm_provider="anthropic", - litellm_params=creds, - check_provider_endpoint=True, - ) - assert len(valid_models) == 0 - - creds = CredentialLiteLLMParams(api_key=os.getenv("ANTHROPIC_API_KEY")) - valid_models = get_valid_models( - custom_llm_provider="anthropic", - litellm_params=creds, - check_provider_endpoint=True, - ) - assert len(valid_models) > 0 - assert "anthropic/claude-sonnet-4-6" in valid_models - - def test_get_whitelisted_models(): """ Snapshot of all bedrock models as of 12/24/2024. diff --git a/tests/llm_responses_api_testing/base_responses_api.py b/tests/llm_responses_api_testing/base_responses_api.py deleted file mode 100644 index 8f3867a9201..00000000000 --- a/tests/llm_responses_api_testing/base_responses_api.py +++ /dev/null @@ -1,658 +0,0 @@ -import httpx -import json -import pytest -from typing import Any, Dict, List, Optional -from unittest.mock import MagicMock, Mock, patch -from litellm._uuid import uuid -import time -import base64 - -import litellm -from abc import ABC, abstractmethod - -from litellm.integrations.custom_logger import CustomLogger -from litellm.types.utils import StandardLoggingPayload -from litellm.types.llms.openai import ( - ResponseCompletedEvent, - ResponsesAPIResponse, - ResponseAPIUsage, - IncompleteDetails, -) -from openai.types.responses.response_create_params import ( - ResponseInputParam, -) -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler -import openai - - -def validate_responses_api_response(response, final_chunk: bool = False): - """ - Validate that a response from litellm.responses() or litellm.aresponses() - conforms to the expected ResponsesAPIResponse structure. - - Args: - response: The response object to validate - - Raises: - AssertionError: If the response doesn't match the expected structure - """ - # Validate response structure - print("response=", json.dumps(response, indent=4, default=str)) - assert isinstance( - response, ResponsesAPIResponse - ), "Response should be an instance of ResponsesAPIResponse" - - # Required fields - assert "id" in response and isinstance( - response["id"], str - ), "Response should have a string 'id' field" - assert "created_at" in response and isinstance( - response["created_at"], int - ), "Response should have an integer 'created_at' field" - if response.get("status") == "completed": - assert "output" in response and isinstance( - response["output"], list - ), "Response should have a list 'output' field" - - # Optional fields with their expected types - optional_fields = { - "error": (dict, type(None)), # error can be dict or None - "incomplete_details": (IncompleteDetails, type(None)), - "instructions": (str, type(None)), - "metadata": dict, - "model": str, - "object": str, - "parallel_tool_calls": (bool, type(None)), - "temperature": (int, float, type(None)), - "tool_choice": (dict, str, type(None)), - "tools": (list, type(None)), - "top_p": (int, float, type(None)), - "max_output_tokens": (int, type(None)), - "previous_response_id": (str, type(None)), - "reasoning": (dict, type(None)), - "status": str, - "text": dict, - "truncation": (str, type(None)), - "usage": ResponseAPIUsage, - "user": (str, type(None)), - "store": (bool, type(None)), - } - if final_chunk is False: - optional_fields["usage"] = type(None) - - for field, expected_type in optional_fields.items(): - if field in response: - assert isinstance( - response[field], expected_type - ), f"Field '{field}' should be of type {expected_type}, but got {type(response[field])}" - - # Check if output has at least one item - if final_chunk is True and response.get("status") == "completed": - assert ( - len(response["output"]) > 0 - ), "Response 'output' field should have at least one item" - - return True # Return True if validation passes - - -class BaseResponsesAPITest(ABC): - """ - Abstract base test class that enforces a common test across all test classes. - """ - - @abstractmethod - def get_base_completion_call_args(self) -> dict: - """Must return the base completion call args""" - pass - - def get_advanced_model_for_shell_tool(self) -> Optional[str]: - """If specified, overrides the model used by test_responses_api_shell_tool_streaming_sees_shell_output (e.g. openai/gpt-5.2 for shell support).""" - return None - - @pytest.mark.parametrize("sync_mode", [True, False]) - @pytest.mark.asyncio - async def test_basic_openai_responses_api(self, sync_mode): - litellm.turn_on_debug() - litellm.set_verbose = True - base_completion_call_args = self.get_base_completion_call_args() - try: - if sync_mode: - response = litellm.responses( - input="Basic ping", - max_output_tokens=20, - **base_completion_call_args, - ) - else: - response = await litellm.aresponses( - input="Basic ping", - max_output_tokens=20, - **base_completion_call_args, - ) - except litellm.InternalServerError: - pytest.skip("Skipping test due to litellm.InternalServerError") - print("litellm response=", json.dumps(response, indent=4, default=str)) - - # Use the helper function to validate the response - validate_responses_api_response(response, final_chunk=True) - - @pytest.mark.parametrize("sync_mode", [True, False]) - @pytest.mark.asyncio - @pytest.mark.flaky(retries=3, delay=2) - async def test_basic_openai_responses_api_streaming(self, sync_mode): - litellm.turn_on_debug() - # Enable cost calculation for streaming usage - litellm.include_cost_in_streaming_usage = True - base_completion_call_args = self.get_base_completion_call_args() - collected_content_string = "" - response_completed_event = None - if sync_mode: - response = litellm.responses( - input="Basic ping", stream=True, **base_completion_call_args - ) - for event in response: - print("litellm response=", json.dumps(event, indent=4, default=str)) - if event.type == "response.output_text.delta": - collected_content_string += event.delta - elif event.type == "response.completed": - response_completed_event = event - else: - response = await litellm.aresponses( - input="Basic ping", stream=True, **base_completion_call_args - ) - async for event in response: - print("litellm response=", json.dumps(event, indent=4, default=str)) - if event.type == "response.output_text.delta": - collected_content_string += event.delta - elif event.type == "response.completed": - response_completed_event = event - - # assert the response completed event is not None - assert response_completed_event is not None - - # assert the response completed event has a response - assert response_completed_event.response is not None - - # For async agent APIs (like Manus), the response may be in 'running' state - # without content yet - this is valid behavior - response_status = response_completed_event.response.status - if response_status in ["running", "pending"]: - # Running/pending state is acceptable - task started successfully - print( - f"Response is in '{response_status}' state - async agent API behavior" - ) - assert response_completed_event.response.id is not None - else: - # For completed responses, validate content and usage - # assert the delta chunks content had len(collected_content_string) > 0 - # this content is typically rendered on chat ui's - assert len(collected_content_string) > 0 - - # assert the response completed event includes the usage - assert response_completed_event.response.usage is not None - - # basic test assert the usage seems reasonable - print( - "response_completed_event.response.usage=", - response_completed_event.response.usage, - ) - assert ( - response_completed_event.response.usage.input_tokens > 0 - and response_completed_event.response.usage.input_tokens < 100 - ) - assert ( - response_completed_event.response.usage.output_tokens > 0 - and response_completed_event.response.usage.output_tokens < 2000 - ) - assert ( - response_completed_event.response.usage.total_tokens > 0 - and response_completed_event.response.usage.total_tokens < 2000 - ) - - # total tokens should be the sum of input and output tokens - assert ( - response_completed_event.response.usage.total_tokens - == response_completed_event.response.usage.input_tokens - + response_completed_event.response.usage.output_tokens - ) - - # assert the response completed event includes cost when include_cost_in_streaming_usage is True - assert hasattr( - response_completed_event.response.usage, "cost" - ), "Cost should be included in streaming responses API usage object" - assert ( - response_completed_event.response.usage.cost > 0 - ), "Cost should be greater than 0" - print( - f"Cost found in streaming response: {response_completed_event.response.usage.cost}" - ) - - # Reset the setting - litellm.include_cost_in_streaming_usage = False - - @pytest.mark.parametrize("sync_mode", [False, True]) - @pytest.mark.asyncio - async def test_basic_openai_responses_delete_endpoint(self, sync_mode): - litellm.turn_on_debug() - litellm.set_verbose = True - base_completion_call_args = self.get_base_completion_call_args() - if sync_mode: - response = litellm.responses( - input="Basic ping", max_output_tokens=20, **base_completion_call_args - ) - - # delete the response - if isinstance(response, ResponsesAPIResponse): - litellm.delete_responses( - response_id=response.id, **base_completion_call_args - ) - else: - raise ValueError("response is not a ResponsesAPIResponse") - else: - response = await litellm.aresponses( - input="Basic ping", max_output_tokens=20, **base_completion_call_args - ) - - # async delete the response - if isinstance(response, ResponsesAPIResponse): - await litellm.adelete_responses( - response_id=response.id, **base_completion_call_args - ) - else: - raise ValueError("response is not a ResponsesAPIResponse") - - @pytest.mark.parametrize("sync_mode", [True, False]) - @pytest.mark.flaky(retries=3, delay=2) - @pytest.mark.asyncio - async def test_basic_openai_responses_streaming_delete_endpoint(self, sync_mode): - # litellm.turn_on_debug() - # litellm.set_verbose = True - base_completion_call_args = self.get_base_completion_call_args() - response_id = None - if sync_mode: - response_id = None - response = litellm.responses( - input="Basic ping", - max_output_tokens=20, - stream=True, - **base_completion_call_args, - ) - for event in response: - print("litellm response=", json.dumps(event, indent=4, default=str)) - if "response" in event: - response_obj = event.get("response") - if response_obj is not None: - response_id = response_obj.get("id") - print("got response_id=", response_id) - - # delete the response - assert response_id is not None - litellm.delete_responses( - response_id=response_id, **base_completion_call_args - ) - else: - response = await litellm.aresponses( - input="Basic ping", - max_output_tokens=20, - stream=True, - **base_completion_call_args, - ) - async for event in response: - print("litellm response=", json.dumps(event, indent=4, default=str)) - if "response" in event: - response_obj = event.get("response") - if response_obj is not None: - response_id = response_obj.get("id") - print("got response_id=", response_id) - - # delete the response - assert response_id is not None - await litellm.adelete_responses( - response_id=response_id, **base_completion_call_args - ) - - @pytest.mark.parametrize("sync_mode", [False, True]) - @pytest.mark.flaky(retries=3, delay=2) - @pytest.mark.asyncio - async def test_basic_openai_responses_get_endpoint(self, sync_mode): - litellm.turn_on_debug() - litellm.set_verbose = True - base_completion_call_args = self.get_base_completion_call_args() - if sync_mode: - response = litellm.responses( - input="Basic ping", max_output_tokens=20, **base_completion_call_args - ) - - # get the response - if isinstance(response, ResponsesAPIResponse): - result = litellm.get_responses( - response_id=response.id, **base_completion_call_args - ) - assert result is not None - assert result.id == response.id - assert result.output_text == response.output_text - else: - raise ValueError("response is not a ResponsesAPIResponse") - else: - response = await litellm.aresponses( - input="Basic ping", max_output_tokens=20, **base_completion_call_args - ) - # async get the response - if isinstance(response, ResponsesAPIResponse): - result = await litellm.aget_responses( - response_id=response.id, **base_completion_call_args - ) - assert result is not None - assert result.id == response.id - assert result.output_text == response.output_text - else: - raise ValueError("response is not a ResponsesAPIResponse") - - @pytest.mark.asyncio - async def test_multiturn_responses_api(self): - litellm.turn_on_debug() - litellm.set_verbose = True - try: - base_completion_call_args = self.get_base_completion_call_args() - response_1 = await litellm.aresponses( - input="Basic ping", max_output_tokens=20, **base_completion_call_args - ) - - # follow up with a second request - response_1_id = response_1.id - response_2 = await litellm.aresponses( - input="Basic ping", - max_output_tokens=20, - previous_response_id=response_1_id, - **base_completion_call_args, - ) - - # assert the response is not None - assert response_1 is not None - assert response_2 is not None - except litellm.InternalServerError: - pytest.skip("Skipping test due to litellm.InternalServerError") - - @pytest.mark.asyncio - async def test_responses_api_with_tool_calls(self): - """Test that calls the Responses API with tool calls including function call and output""" - litellm.turn_on_debug() - litellm.set_verbose = True - base_completion_call_args = self.get_base_completion_call_args() - - # Define the input with message, function call, and function call output - input_data: ResponseInputParam = [ - { - "type": "message", - "role": "user", - "content": "How is the weather in São Paulo today ?", - }, - { - "type": "function_call", - "arguments": '{"location": "São Paulo, Brazil"}', - "call_id": "fc_1fe70e2a-a596-45ef-b72c-9b8567c460e5", - "name": "get_weather", - "id": "fc_1fe70e2a-a596-45ef-b72c-9b8567c460e5", - "status": "completed", - }, - { - "type": "function_call_output", - "call_id": "fc_1fe70e2a-a596-45ef-b72c-9b8567c460e5", - "output": "Rainy", - }, - ] - - # Define the tools - tools = [ - { - "type": "function", - "name": "get_weather", - "description": "Get current temperature for a given location.", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "City and country e.g. Bogotá, Colombia", - } - }, - "required": ["location"], - "additionalProperties": False, - }, - } - ] - - try: - # Make the responses API call - response = await litellm.aresponses( - input=input_data, store=False, tools=tools, **base_completion_call_args - ) - except litellm.InternalServerError: - pytest.skip("Skipping test due to litellm.InternalServerError") - - print("litellm response=", json.dumps(response, indent=4, default=str)) - - # Validate the response structure - validate_responses_api_response(response, final_chunk=True) - - # Additional assertions specific to tool calls - assert response is not None - assert "output" in response - # For async agent APIs (like Manus), the response may be in 'running' state - # without output yet - this is valid behavior - if response.get("status") in ["running", "pending"]: - print( - f"Response is in '{response.get('status')}' state - async agent API behavior" - ) - assert response.get("id") is not None - else: - assert len(response["output"]) > 0 - - def test_openai_responses_api_dict_input_filtering(self): - """ - Test that regular dict inputs with status fields are properly filtered - to replicate exclude_unset=True behavior for non-Pydantic objects. - """ - from litellm.llms.openai.responses.transformation import ( - OpenAIResponsesAPIConfig, - ) - - # Test input with regular dict objects (like from JSON) - test_input = [ - {"role": "user", "content": "test"}, - { - "id": "rs_123", - "summary": [{"text": "test", "type": "summary_text"}], - "type": "reasoning", - "content": None, # Should be filtered out - "encrypted_content": None, # Should be filtered out - "status": None, # Should be filtered out - }, - { - "arguments": "{}", - "call_id": "call_123", - "name": "get_today", - "type": "function_call", - "id": "fc_123", - "status": "completed", # Should be preserved (not a default field) - }, - ] - - config = OpenAIResponsesAPIConfig() - validated_input = config._validate_input_param(test_input) - - # Verify the results - assert len(validated_input) == 3 - - # Check reasoning item (index 1) - reasoning_item = validated_input[1] - assert reasoning_item["type"] == "reasoning" - assert ( - "status" not in reasoning_item - ), "status field should be filtered out from reasoning item" - assert ( - "content" not in reasoning_item - ), "content field should be filtered out from reasoning item" - assert ( - "encrypted_content" not in reasoning_item - ), "encrypted_content field should be filtered out from reasoning item" - # Note: ID auto-generation was disabled, so reasoning items may not have IDs - # Only check for ID if it was present in the original input - if "id" in reasoning_item: - assert reasoning_item["id"] == "rs_123", "ID should be preserved if present" - assert "summary" in reasoning_item, "summary field should be preserved" - - # Check function call item (index 2) - function_call_item = validated_input[2] - assert function_call_item["type"] == "function_call" - assert ( - "status" in function_call_item - ), "status field should be preserved in function call item" - assert ( - function_call_item["status"] == "completed" - ), "status value should be preserved" - - print("✅ OpenAI Responses API dict input filtering test passed") - - @pytest.mark.parametrize("sync_mode", [False, True]) - @pytest.mark.flaky(retries=3, delay=2) - @pytest.mark.asyncio - async def test_basic_openai_responses_cancel_endpoint(self, sync_mode): - try: - litellm.turn_on_debug() - litellm.set_verbose = True - base_completion_call_args = self.get_base_completion_call_args() - if sync_mode: - response = litellm.responses( - input="Basic ping", - max_output_tokens=20, - background=True, - **base_completion_call_args, - ) - - # cancel the response - if isinstance(response, ResponsesAPIResponse): - cancel_result = litellm.cancel_responses( - response_id=response.id, **base_completion_call_args - ) - assert cancel_result is not None - assert hasattr(cancel_result, "id") - # The actual response structure depends on the provider implementation - assert isinstance(cancel_result, ResponsesAPIResponse) - else: - raise ValueError("response is not a ResponsesAPIResponse") - else: - response = await litellm.aresponses( - input="Basic ping", - max_output_tokens=20, - background=True, - **base_completion_call_args, - ) - - # async cancel the response - if isinstance(response, ResponsesAPIResponse): - cancel_result = await litellm.acancel_responses( - response_id=response.id, **base_completion_call_args - ) - assert cancel_result is not None - assert hasattr(cancel_result, "id") - # The actual response structure depends on the provider implementation - assert isinstance(cancel_result, ResponsesAPIResponse) - else: - raise ValueError("response is not a ResponsesAPIResponse") - except Exception as e: - if "Cannot cancel a completed response" in str(e): - pass - else: - raise e - - @pytest.mark.parametrize("sync_mode", [False, True]) - @pytest.mark.asyncio - async def test_cancel_responses_invalid_response_id(self, sync_mode): - """Test cancel_responses with invalid response ID should raise appropriate error""" - base_completion_call_args = self.get_base_completion_call_args() - - if sync_mode: - with pytest.raises(openai.APIError): - litellm.cancel_responses( - response_id="invalid_response_id_12345", **base_completion_call_args - ) - else: - with pytest.raises(openai.APIError): - await litellm.acancel_responses( - response_id="invalid_response_id_12345", **base_completion_call_args - ) - - @pytest.mark.asyncio - async def test_responses_api_context_management_server_side_compaction(self): - """ - E2E test for server-side compaction (context_management) on OpenAI Responses API. - Passes context_management with compact_threshold; validates that the request is - accepted and returns a valid response. Compaction may not run for short inputs. - """ - base_completion_call_args = self.get_base_completion_call_args() - model = base_completion_call_args.get("model") or "" - # Azure does not support compaction context_management (only clear_tool_results) - if "azure/" in str(model): - pytest.skip("context_management compaction is not supported on Azure") - if "openai/" not in str(model): - pytest.skip( - "context_management server-side compaction e2e is only run for OpenAI" - ) - context_management = [{"type": "compaction", "compact_threshold": 200000}] - try: - response = await litellm.aresponses( - input="Short ping to verify context_management is accepted.", - max_output_tokens=20, - context_management=context_management, - **base_completion_call_args, - ) - except litellm.InternalServerError: - pytest.skip("Skipping test due to litellm.InternalServerError") - validate_responses_api_response(response, final_chunk=True) - assert response.get("id") is not None - assert response.get("status") is not None - - @pytest.mark.asyncio - async def test_responses_api_shell_tool(self): - """ - E2E test for Shell tool on OpenAI Responses API. - Passes tools=[{"type": "shell", "environment": {"type": "container_auto"}}]; - validates that the request is accepted and returns a valid response. - Only runs for OpenAI; offline coverage for the Azure route lives in - tests/unit/responses/test_responses_api_request_body.py. - """ - base_completion_call_args = self.get_base_completion_call_args() - model = ( - self.get_advanced_model_for_shell_tool() - or base_completion_call_args.get("model") - or "" - ) - if "openai/" not in str(model): - pytest.skip( - "Shell tool e2e is OpenAI-only; no Azure deployment supports the shell tool yet, re-enable once one exists" - ) - tools = [{"type": "shell", "environment": {"type": "container_auto"}}] - input_msg = "List files in /mnt/data and show python --version." - try: - response = await litellm.aresponses( - **{**base_completion_call_args, "model": model}, - input=input_msg, - max_output_tokens=256, - tools=tools, - tool_choice="auto", - timeout=90, - ) - except litellm.Timeout: - pytest.skip("Provider did not answer the shell tool request within 90s") - except litellm.InternalServerError: - pytest.skip("Skipping test due to litellm.InternalServerError") - except litellm.BadRequestError as e: - if "shell" in str(e).lower() and "not supported" in str(e).lower(): - pytest.skip( - "Shell tool is not supported for this model (e.g. gpt-5.5); use a model that supports shell" - ) - raise - validate_responses_api_response(response, final_chunk=True) - assert response.get("id") is not None - assert response.get("status") is not None - diff --git a/tests/llm_responses_api_testing/conftest.py b/tests/llm_responses_api_testing/conftest.py deleted file mode 100644 index 5501d99cb22..00000000000 --- a/tests/llm_responses_api_testing/conftest.py +++ /dev/null @@ -1,112 +0,0 @@ -# conftest.py - -import asyncio -import importlib - -import pytest - - -import litellm # noqa: E402 - -from tests._vcr_conftest_common import ( # noqa: E402,F401 - VerboseReporterState, - _pin_multipart_boundary, - apply_vcr_auto_marker_to_items, - emit_cassette_cache_session_banner, - emit_vcr_classification_summary, - emit_vcr_diagnostic_log, - install_live_call_probe, - record_vcr_outcome, - register_persister_if_enabled, - reset_vcr_diag_dir, - vcr_config_dict, -) - -_verbose_state = VerboseReporterState() - - -@pytest.fixture(scope="module") -def vcr_config(): - return vcr_config_dict() - - -def pytest_recording_configure(config, vcr): - register_persister_if_enabled(vcr) - - -@pytest.hookimpl(hookwrapper=True) -def pytest_runtest_makereport(item, call): - outcome = yield - rep = outcome.get_result() - setattr(item, f"rep_{rep.when}", rep) - - -@pytest.fixture(autouse=True) -def _vcr_outcome_gate(request, vcr): - install_live_call_probe(request, vcr) - yield - record_vcr_outcome(request, vcr) - - -def pytest_configure(config): - _verbose_state.remember_pluginmanager(config) - reset_vcr_diag_dir() - - -def pytest_runtest_logreport(report): - _verbose_state.maybe_emit_verdict(report) - - -@pytest.fixture(scope="session") -def event_loop(): - try: - loop = asyncio.get_running_loop() - except RuntimeError: - loop = asyncio.new_event_loop() - yield loop - loop.close() - - -@pytest.fixture(scope="function", autouse=True) -def setup_and_teardown(): - """ - This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained. - """ - - - importlib.reload(litellm) - - try: - if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"): - importlib.reload(litellm.proxy.proxy_server) - except Exception as e: - print(f"Error reloading litellm.proxy.proxy_server: {e}") - - loop = asyncio.get_event_loop_policy().new_event_loop() - asyncio.set_event_loop(loop) - print(litellm) - yield - - # Teardown code (executes after the yield point) - loop.close() # Close the loop created earlier - asyncio.set_event_loop(None) # Remove the reference to the loop - - -def pytest_collection_modifyitems(config, items): - apply_vcr_auto_marker_to_items(items) - - custom_logger_tests = [ - item for item in items if "custom_logger" in item.parent.name - ] - other_tests = [item for item in items if "custom_logger" not in item.parent.name] - - custom_logger_tests.sort(key=lambda x: x.name) - other_tests.sort(key=lambda x: x.name) - - items[:] = custom_logger_tests + other_tests - - -def pytest_terminal_summary(terminalreporter, exitstatus, config): - emit_cassette_cache_session_banner(terminalreporter) - emit_vcr_classification_summary(terminalreporter) - emit_vcr_diagnostic_log(terminalreporter) diff --git a/tests/llm_responses_api_testing/test_anthropic_responses_api.py b/tests/llm_responses_api_testing/test_anthropic_responses_api.py deleted file mode 100644 index 89751d7fdd0..00000000000 --- a/tests/llm_responses_api_testing/test_anthropic_responses_api.py +++ /dev/null @@ -1,105 +0,0 @@ -import pytest - -import litellm -from base_responses_api import BaseResponsesAPITest -from openai.types.responses.function_tool import FunctionTool - - -class TestAnthropicResponsesAPITest(BaseResponsesAPITest): - def get_base_completion_call_args(self): - # litellm.turn_on_debug() - return { - "model": "anthropic/claude-sonnet-4-5", - } - - async def test_basic_openai_responses_delete_endpoint(self, sync_mode=False): - pytest.skip("DELETE responses is not supported for anthropic") - - async def test_basic_openai_responses_streaming_delete_endpoint( - self, sync_mode=False - ): - pytest.skip("DELETE responses is not supported for anthropic") - - async def test_basic_openai_responses_get_endpoint(self, sync_mode=False): - pytest.skip("GET responses is not supported for anthropic") - - async def test_basic_openai_responses_cancel_endpoint(self, sync_mode=False): - pytest.skip("CANCEL responses is not supported for anthropic") - - async def test_cancel_responses_invalid_response_id(self, sync_mode=False): - pytest.skip("CANCEL responses is not supported for anthropic") - - -def test_multiturn_tool_calls(): - # Test streaming response with tools for Anthropic - litellm.turn_on_debug() - shell_tool = dict( - FunctionTool( - type="function", - name="shell", - description="Runs a shell command, and returns its output.", - parameters={ - "type": "object", - "properties": { - "command": {"type": "array", "items": {"type": "string"}}, - "workdir": { - "type": "string", - "description": "The working directory for the command.", - }, - }, - "required": ["command"], - }, - strict=True, - ) - ) - - # Step 1: Initial request with the tool - response = litellm.responses( - input=[ - { - "role": "user", - "content": [ - {"type": "input_text", "text": "make a hello world html file"} - ], - "type": "message", - } - ], - model="anthropic/claude-haiku-4-5-20251001", - instructions="You are a helpful coding assistant.", - tools=[shell_tool], - ) - - print("response=", response) - - # Step 2: Send the results of the tool call back to the model - # Get the response ID and tool call ID from the response - - response_id = response.id - tool_call_id = None - for item in response.output: - if hasattr(item, "type") and item.type == "function_call": - tool_call_id = getattr(item, "call_id", None) - if tool_call_id: - break - - # Validate that we got a tool call with a valid call_id - if not tool_call_id: - raise AssertionError( - f"Expected a function_call with a valid call_id in response.output, but got: {response.output}" - ) - - # Use await with asyncio.run for the async function - follow_up_response = litellm.responses( - model="anthropic/claude-haiku-4-5-20251001", - previous_response_id=response_id, - input=[ - { - "type": "function_call_output", - "call_id": tool_call_id, - "output": '{"output":"\\n\\n Hello Page\\n\\n\\n

Hi

\\n

Welcome to this simple webpage!

\\n\\n > index.html\\n","metadata":{"exit_code":0,"duration_seconds":0}}', - } - ], - tools=[shell_tool], - ) - - print("follow_up_response=", follow_up_response) diff --git a/tests/llm_responses_api_testing/test_azure_responses_api.py b/tests/llm_responses_api_testing/test_azure_responses_api.py deleted file mode 100644 index e2b1de33182..00000000000 --- a/tests/llm_responses_api_testing/test_azure_responses_api.py +++ /dev/null @@ -1,35 +0,0 @@ -import os -import pytest - -import litellm -from base_responses_api import BaseResponsesAPITest - - -class TestAzureResponsesAPITest(BaseResponsesAPITest): - test_multiturn_responses_api = None - test_responses_api_with_tool_calls = None - - def get_base_completion_call_args(self): - return { - "model": "azure/gpt-4.1-mini", - "truncation": "auto", - "api_base": os.getenv("AZURE_AI_API_BASE"), - "api_key": os.getenv("AZURE_AI_API_KEY"), - "api_version": "2025-03-01-preview", - } - - -@pytest.mark.asyncio -async def test_azure_responses_api_preview_api_version(): - """ - Ensure new azure preview api version is working - """ - litellm.turn_on_debug() - response = await litellm.aresponses( - model="azure/gpt-5-mini", - truncation="auto", - api_version="preview", - api_base=os.getenv("AZURE_AI_API_BASE"), - api_key=os.getenv("AZURE_AI_API_KEY"), - input="Hello, can you tell me a short joke?", - ) diff --git a/tests/llm_responses_api_testing/test_google_ai_studio_responses_api.py b/tests/llm_responses_api_testing/test_google_ai_studio_responses_api.py deleted file mode 100644 index 068de771e2a..00000000000 --- a/tests/llm_responses_api_testing/test_google_ai_studio_responses_api.py +++ /dev/null @@ -1,233 +0,0 @@ -import os -import pytest - -import litellm -import json -from base_responses_api import BaseResponsesAPITest - - -@pytest.mark.asyncio -async def test_basic_google_ai_studio_responses_api_with_tools(): - litellm.turn_on_debug() - litellm.set_verbose = True - request_model = "gemini/gemini-2.5-flash" - response = await litellm.aresponses( - model=request_model, - input="what is the latest version of supabase python package and when was it released?", - tools=[{"type": "web_search_preview", "search_context_size": "low"}], - ) - print("litellm response=", json.dumps(response, indent=4, default=str)) - - -@pytest.mark.asyncio -async def test_gemini_3_responses_api_with_thought_signatures(): - """ - Test that Gemini 3 Responses API preserves thought signatures in function calls. - This test verifies that provider_specific_fields with thought_signature are correctly - preserved when using the Responses API with Gemini 3. - """ - if not os.getenv("GEMINI_API_KEY"): - pytest.skip("GEMINI_API_KEY not set") - - litellm.set_verbose = False - request_model = "gemini/gemini-3.1-pro-preview" - - tools = [ - { - "type": "function", - "name": "get_weather", - "description": "Get the current weather for a location", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "City and country e.g. Mumbai, India", - }, - "units": { - "type": "string", - "enum": ["celsius", "fahrenheit"], - "description": "Units the temperature will be returned in.", - }, - }, - "required": ["location", "units"], - "additionalProperties": False, - }, - "strict": True, - } - ] - - # Step 1: Initial request with tools - response = await litellm.aresponses( - model=request_model, - input="What is the weather in Mumbai?", - tools=tools, - reasoning_effort="low", - ) - - # Validate response structure - from litellm.types.llms.openai import ResponsesAPIResponse - - assert isinstance( - response, ResponsesAPIResponse - ), "Response should be a ResponsesAPIResponse" - assert ( - hasattr(response, "output") or "output" in response - ), "Response should have 'output' field" - assert isinstance(response.output, list), "Output should be a list" - - # Find function call in output - function_call_item = None - for item in response.output: - # Convert to dict if it's a Pydantic model for easier access - if hasattr(item, "model_dump"): - item_dict = item.model_dump() - elif hasattr(item, "__dict__"): - item_dict = dict(item) if not isinstance(item, dict) else item - else: - item_dict = item if isinstance(item, dict) else {} - - if isinstance(item_dict, dict) and item_dict.get("type") == "function_call": - function_call_item = item_dict - break - - # Verify function call exists - assert ( - function_call_item is not None - ), "Response should contain a function_call item" - assert ( - function_call_item.get("name") == "get_weather" - ), "Function call should be for get_weather" - - # Verify thought signature is present in provider_specific_fields - provider_specific_fields = function_call_item.get("provider_specific_fields") - assert ( - provider_specific_fields is not None - ), "Function call should have provider_specific_fields" - assert ( - "thought_signature" in provider_specific_fields - ), "provider_specific_fields should contain thought_signature" - assert isinstance( - provider_specific_fields["thought_signature"], str - ), "thought_signature should be a string" - assert ( - len(provider_specific_fields["thought_signature"]) > 0 - ), "thought_signature should not be empty" - - print( - f"✅ Thought signature preserved: {provider_specific_fields['thought_signature'][:50]}..." - ) - - -@pytest.mark.asyncio -async def test_gemini_3_responses_api_streaming_with_thought_signatures(): - """ - Test that Gemini 3 Responses API preserves thought signatures in streaming mode. - This test verifies that provider_specific_fields with thought_signature are correctly - preserved when using streaming Responses API with Gemini 3. - """ - if not os.getenv("GEMINI_API_KEY"): - pytest.skip("GEMINI_API_KEY not set") - - litellm.set_verbose = False - request_model = "gemini/gemini-3.1-pro-preview" - - tools = [ - { - "type": "function", - "name": "get_weather", - "description": "Get the current weather for a location", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "City and country e.g. Mumbai, India", - }, - "units": { - "type": "string", - "enum": ["celsius", "fahrenheit"], - "description": "Units the temperature will be returned in.", - }, - }, - "required": ["location", "units"], - "additionalProperties": False, - }, - "strict": True, - } - ] - - # Step 1: Streaming request with tools - response_stream = await litellm.aresponses( - model=request_model, - input="What is the weather in Mumbai?", - tools=tools, - stream=True, - reasoning_effort="low", - ) - - # Collect all chunks - chunks = [] - completed_response = None - - async for chunk in response_stream: - chunks.append(chunk) - # Check if this is the completed response event - if hasattr(chunk, "type") and chunk.type == "response.completed": - completed_response = chunk.response - elif isinstance(chunk, dict) and chunk.get("type") == "response.completed": - completed_response = chunk.get("response") - - # Verify we got chunks - assert len(chunks) > 0, "Should receive at least one chunk" - - # If we have a completed response, check for thought signatures - if completed_response: - output = completed_response.get("output", []) - function_call_item = None - for item in output: - if isinstance(item, dict) and item.get("type") == "function_call": - function_call_item = item - break - - if function_call_item: - provider_specific_fields = function_call_item.get( - "provider_specific_fields" - ) - if provider_specific_fields: - thought_signature = provider_specific_fields.get("thought_signature") - if thought_signature: - assert isinstance( - thought_signature, str - ), "thought_signature should be a string" - assert ( - len(thought_signature) > 0 - ), "thought_signature should not be empty" - print( - f"✅ Streaming thought signature preserved: {thought_signature[:50]}..." - ) - - print(f"✅ Collected {len(chunks)} streaming chunks") - - -class TestGoogleAIStudioResponsesAPITest(BaseResponsesAPITest): - def get_base_completion_call_args(self): - # litellm.turn_on_debug() - return {"model": "gemini/gemini-2.5-flash-lite"} - - async def test_basic_openai_responses_delete_endpoint(self, sync_mode=False): - pytest.skip("DELETE responses is not supported for Google AI Studio") - - async def test_basic_openai_responses_streaming_delete_endpoint( - self, sync_mode=False - ): - pytest.skip("DELETE responses is not supported for Google AI Studio") - - async def test_basic_openai_responses_get_endpoint(self, sync_mode=False): - pytest.skip("GET responses is not supported for Google AI Studio") - - async def test_basic_openai_responses_cancel_endpoint(self, sync_mode=False): - pytest.skip("CANCEL responses is not supported for Google AI Studio") - - async def test_cancel_responses_invalid_response_id(self, sync_mode=False): - pytest.skip("CANCEL responses is not supported for Google AI Studio") diff --git a/tests/llm_responses_api_testing/test_openai_responses_api.py b/tests/llm_responses_api_testing/test_openai_responses_api.py deleted file mode 100644 index 2705b91ee3a..00000000000 --- a/tests/llm_responses_api_testing/test_openai_responses_api.py +++ /dev/null @@ -1,805 +0,0 @@ -import asyncio -import json -import os -import time -from typing import Optional, cast - -import pytest -from base_responses_api import BaseResponsesAPITest, validate_responses_api_response - -import litellm -from litellm.integrations.custom_logger import CustomLogger -from litellm.types.llms.openai import ( - ResponseAPIUsage, - ResponseCompletedEvent, - ResponsesAPIResponse, -) -from litellm.types.utils import StandardLoggingPayload - - -class TestOpenAIResponsesAPITest(BaseResponsesAPITest): - test_responses_api_with_tool_calls = None - - def get_base_completion_call_args(self): - return { - "model": "openai/gpt-5.5", - } - - def get_advanced_model_for_shell_tool(self): - return "openai/gpt-5.2" - - -class TestCustomLogger(CustomLogger): - def __init__( - self, - ): - self.standard_logging_object: Optional[StandardLoggingPayload] = None - - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - print("in async_log_success_event") - print("kwargs=", json.dumps(kwargs, indent=4, default=str)) - self.standard_logging_object = kwargs["standard_logging_object"] - pass - - -def validate_standard_logging_payload( - slp: StandardLoggingPayload, response: ResponsesAPIResponse, request_model: str -): - """ - Validate that a StandardLoggingPayload object matches the expected response - - Args: - slp (StandardLoggingPayload): The standard logging payload object to validate - response (dict): The litellm response to compare against - request_model (str): The model name that was requested - """ - # Validate payload exists - assert slp is not None, "Standard logging payload should not be None" - - # Validate token counts - print( - "VALIDATING STANDARD LOGGING PAYLOAD. response=", - json.dumps(response, indent=4, default=str), - ) - print("FIELDS IN SLP=", json.dumps(slp, indent=4, default=str)) - print("SLP PROMPT TOKENS=", slp["prompt_tokens"]) - print("RESPONSE PROMPT TOKENS=", response["usage"]["input_tokens"]) - assert ( - slp["prompt_tokens"] == response["usage"]["input_tokens"] - ), "Prompt tokens mismatch" - assert ( - slp["completion_tokens"] == response["usage"]["output_tokens"] - ), "Completion tokens mismatch" - assert ( - slp["total_tokens"] - == response["usage"]["input_tokens"] + response["usage"]["output_tokens"] - ), "Total tokens mismatch" - - # Validate spend and response metadata - assert slp["response_cost"] > 0, "Response cost should be greater than 0" - assert slp["id"] == response["id"], "Response ID mismatch" - assert slp["model"] == request_model, "Model name mismatch" - - # Validate messages - assert slp["messages"] == [{"content": "hi", "role": "user"}], "Messages mismatch" - - # Validate complete response structure - validate_responses_match(slp["response"], response) - - -@pytest.mark.asyncio -def test_basic_openai_responses_api_streaming_with_logging(): - litellm.turn_on_debug() - litellm.set_verbose = True - test_custom_logger = TestCustomLogger() - litellm.callbacks = [test_custom_logger] - request_model = "gpt-5.5" - response = litellm.responses( - model=request_model, - input="hi", - stream=True, - ) - final_response: Optional[ResponseCompletedEvent] = None - for event in response: - if event.type == "response.completed": - final_response = event - print("litellm response=", json.dumps(event, indent=4, default=str)) - - print("sleeping for 2 seconds...") - time.sleep(2) - print( - "standard logging payload=", - json.dumps(test_custom_logger.standard_logging_object, indent=4, default=str), - ) - - assert final_response is not None - assert test_custom_logger.standard_logging_object is not None - - validate_standard_logging_payload( - slp=test_custom_logger.standard_logging_object, - response=final_response.response, - request_model=request_model, - ) - - -def validate_responses_match(slp_response, litellm_response): - """Validate that the standard logging payload OpenAI response matches the litellm response""" - # Validate core fields - assert slp_response["id"] == litellm_response["id"], "ID mismatch" - assert slp_response["model"] == litellm_response["model"], "Model mismatch" - assert ( - slp_response["created_at"] == litellm_response["created_at"] - ), "Created at mismatch" - - # Validate usage - assert ( - slp_response["usage"]["prompt_tokens"] - == litellm_response["usage"]["input_tokens"] - ), "Input tokens mismatch" - assert ( - slp_response["usage"]["completion_tokens"] - == litellm_response["usage"]["output_tokens"] - ), "Output tokens mismatch" - assert ( - slp_response["usage"]["total_tokens"] - == litellm_response["usage"]["total_tokens"] - ), "Total tokens mismatch" - - # Validate output/messages - assert len(slp_response["output"]) == len( - litellm_response["output"] - ), "Output length mismatch" - for slp_msg, litellm_msg in zip(slp_response["output"], litellm_response["output"]): - assert slp_msg["role"] == litellm_msg.role, "Message role mismatch" - # Access the content's text field for the litellm response - litellm_content = litellm_msg.content[0].text if litellm_msg.content else "" - assert ( - slp_msg["content"][0]["text"] == litellm_content - ), f"Message content mismatch. Expected {litellm_content}, Got {slp_msg['content']}" - assert slp_msg["status"] == litellm_msg.status, "Message status mismatch" - - -@pytest.mark.asyncio -async def test_basic_openai_responses_api_non_streaming_with_logging(): - litellm.turn_on_debug() - litellm.set_verbose = True - test_custom_logger = TestCustomLogger() - litellm.callbacks = [test_custom_logger] - request_model = "gpt-5.5" - response = await litellm.aresponses( - model=request_model, - input="hi", - ) - - print("litellm response=", json.dumps(response, indent=4, default=str)) - print("response hidden params=", response._hidden_params) - - print("sleeping for 2 seconds...") - await asyncio.sleep(5) - print( - "standard logging payload=", - json.dumps(test_custom_logger.standard_logging_object, indent=4, default=str), - ) - print("response usage=", response.usage) - - assert response is not None - assert test_custom_logger.standard_logging_object is not None - - validate_standard_logging_payload( - test_custom_logger.standard_logging_object, response, request_model - ) - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_openai_responses_api_returns_headers(sync_mode): - """ - Test that OpenAI responses API returns OpenAI headers in _hidden_params. - This ensures the proxy can forward these headers to clients. - - Related issue: LiteLLM responses API should return OpenAI headers like chat completions does - """ - litellm.turn_on_debug() - litellm.set_verbose = True - - if sync_mode: - response = litellm.responses( - model="gpt-5.5", - input="Say hello", - max_output_tokens=20, - ) - else: - response = await litellm.aresponses( - model="gpt-5.5", - input="Say hello", - max_output_tokens=20, - ) - - # Verify response is valid - assert response is not None - assert isinstance(response, ResponsesAPIResponse) - - # Verify _hidden_params exists - assert hasattr( - response, "_hidden_params" - ), "Response should have _hidden_params attribute" - assert response._hidden_params is not None, "_hidden_params should not be None" - - # Verify additional_headers exists in _hidden_params - assert ( - "additional_headers" in response._hidden_params - ), "_hidden_params should contain 'additional_headers' key" - - additional_headers = response._hidden_params["additional_headers"] - assert isinstance( - additional_headers, dict - ), "additional_headers should be a dictionary" - assert len(additional_headers) > 0, "additional_headers should not be empty" - - # Check for expected OpenAI rate limit headers - # These can be either direct (x-ratelimit-*) or prefixed (llm_provider-x-ratelimit-*) - rate_limit_headers = [ - "x-ratelimit-remaining-tokens", - "x-ratelimit-limit-tokens", - "x-ratelimit-remaining-requests", - "x-ratelimit-limit-requests", - ] - - found_headers = [] - for header_name in rate_limit_headers: - if header_name in additional_headers: - found_headers.append(header_name) - elif f"llm_provider-{header_name}" in additional_headers: - found_headers.append(f"llm_provider-{header_name}") - - assert ( - len(found_headers) > 0 - ), f"Should find at least one OpenAI rate limit header. Headers found: {list(additional_headers.keys())}" - - # Verify headers key also exists (raw headers) - assert ( - "headers" in response._hidden_params - ), "_hidden_params should contain 'headers' key with raw response headers" - - print( - f"✓ Successfully validated OpenAI headers in {'sync' if sync_mode else 'async'} mode" - ) - print(f" Found {len(additional_headers)} headers total") - print(f" Rate limit headers found: {found_headers}") - - -def validate_stream_event(event): - """ - Validate that a streaming event from litellm.responses() or litellm.aresponses() - with stream=True conforms to the expected structure based on its event type. - - Args: - event: The streaming event object to validate - - Raises: - AssertionError: If the event doesn't match the expected structure for its type - """ - # Common validation for all event types - assert hasattr(event, "type"), "Event should have a 'type' attribute" - - # Type-specific validation - if event.type == "response.created" or event.type == "response.in_progress": - assert hasattr( - event, "response" - ), f"{event.type} event should have a 'response' attribute" - validate_responses_api_response(event.response, final_chunk=False) - - elif event.type == "response.completed": - assert hasattr( - event, "response" - ), "response.completed event should have a 'response' attribute" - validate_responses_api_response(event.response, final_chunk=True) - # Usage is guaranteed only on the completed event - assert ( - "usage" in event.response - ), "response.completed event should have usage information" - print("Usage in event.response=", event.response["usage"]) - assert isinstance(event.response["usage"], ResponseAPIUsage) - elif event.type == "response.failed" or event.type == "response.incomplete": - assert hasattr( - event, "response" - ), f"{event.type} event should have a 'response' attribute" - - elif ( - event.type == "response.output_item.added" - or event.type == "response.output_item.done" - ): - assert hasattr( - event, "output_index" - ), f"{event.type} event should have an 'output_index' attribute" - assert hasattr( - event, "item" - ), f"{event.type} event should have an 'item' attribute" - - elif ( - event.type == "response.content_part.added" - or event.type == "response.content_part.done" - ): - assert hasattr( - event, "item_id" - ), f"{event.type} event should have an 'item_id' attribute" - assert hasattr( - event, "output_index" - ), f"{event.type} event should have an 'output_index' attribute" - assert hasattr( - event, "content_index" - ), f"{event.type} event should have a 'content_index' attribute" - assert hasattr( - event, "part" - ), f"{event.type} event should have a 'part' attribute" - - elif event.type == "response.output_text.delta": - assert hasattr( - event, "item_id" - ), f"{event.type} event should have an 'item_id' attribute" - assert hasattr( - event, "output_index" - ), f"{event.type} event should have an 'output_index' attribute" - assert hasattr( - event, "content_index" - ), f"{event.type} event should have a 'content_index' attribute" - assert hasattr( - event, "delta" - ), f"{event.type} event should have a 'delta' attribute" - - elif event.type == "response.output_text.annotation.added": - assert hasattr( - event, "item_id" - ), f"{event.type} event should have an 'item_id' attribute" - assert hasattr( - event, "output_index" - ), f"{event.type} event should have an 'output_index' attribute" - assert hasattr( - event, "content_index" - ), f"{event.type} event should have a 'content_index' attribute" - assert hasattr( - event, "annotation_index" - ), f"{event.type} event should have an 'annotation_index' attribute" - assert hasattr( - event, "annotation" - ), f"{event.type} event should have an 'annotation' attribute" - - elif event.type == "response.output_text.done": - assert hasattr( - event, "item_id" - ), f"{event.type} event should have an 'item_id' attribute" - assert hasattr( - event, "output_index" - ), f"{event.type} event should have an 'output_index' attribute" - assert hasattr( - event, "content_index" - ), f"{event.type} event should have a 'content_index' attribute" - assert hasattr( - event, "text" - ), f"{event.type} event should have a 'text' attribute" - - elif event.type == "response.refusal.delta": - assert hasattr( - event, "item_id" - ), f"{event.type} event should have an 'item_id' attribute" - assert hasattr( - event, "output_index" - ), f"{event.type} event should have an 'output_index' attribute" - assert hasattr( - event, "content_index" - ), f"{event.type} event should have a 'content_index' attribute" - assert hasattr( - event, "delta" - ), f"{event.type} event should have a 'delta' attribute" - - elif event.type == "response.refusal.done": - assert hasattr( - event, "item_id" - ), f"{event.type} event should have an 'item_id' attribute" - assert hasattr( - event, "output_index" - ), f"{event.type} event should have an 'output_index' attribute" - assert hasattr( - event, "content_index" - ), f"{event.type} event should have a 'content_index' attribute" - assert hasattr( - event, "refusal" - ), f"{event.type} event should have a 'refusal' attribute" - - elif event.type == "response.function_call_arguments.delta": - assert hasattr( - event, "item_id" - ), f"{event.type} event should have an 'item_id' attribute" - assert hasattr( - event, "output_index" - ), f"{event.type} event should have an 'output_index' attribute" - assert hasattr( - event, "delta" - ), f"{event.type} event should have a 'delta' attribute" - - elif event.type == "response.function_call_arguments.done": - assert hasattr( - event, "item_id" - ), f"{event.type} event should have an 'item_id' attribute" - assert hasattr( - event, "output_index" - ), f"{event.type} event should have an 'output_index' attribute" - assert hasattr( - event, "arguments" - ), f"{event.type} event should have an 'arguments' attribute" - - elif event.type in [ - "response.file_search_call.in_progress", - "response.file_search_call.searching", - "response.file_search_call.completed", - "response.web_search_call.in_progress", - "response.web_search_call.searching", - "response.web_search_call.completed", - ]: - assert hasattr( - event, "output_index" - ), f"{event.type} event should have an 'output_index' attribute" - assert hasattr( - event, "item_id" - ), f"{event.type} event should have an 'item_id' attribute" - - elif event.type == "error": - assert hasattr( - event, "message" - ), "Error event should have a 'message' attribute" - return True # Return True if validation passes - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_openai_responses_api_streaming_validation(sync_mode): - """Test that validates each streaming event from the responses API""" - litellm.turn_on_debug() - - event_types_seen = set() - - if sync_mode: - response = litellm.responses( - model="gpt-5.5", - input="Tell me about artificial intelligence in 3 sentences.", - stream=True, - ) - for event in response: - print(f"Validating event type: {event.type}") - validate_stream_event(event) - event_types_seen.add(event.type) - else: - response = await litellm.aresponses( - model="gpt-5.5", - input="Tell me about artificial intelligence in 3 sentences.", - stream=True, - ) - async for event in response: - print(f"Validating event type: {event.type}") - validate_stream_event(event) - event_types_seen.add(event.type) - - # At minimum, we should see these core event types - required_events = {"response.created", "response.completed"} - - missing_events = required_events - event_types_seen - assert not missing_events, f"Missing required event types: {missing_events}" - - print(f"Successfully validated all event types: {event_types_seen}") - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_openai_responses_litellm_router(sync_mode): - """ - Test the OpenAI responses API with LiteLLM Router in both sync and async modes - """ - litellm.turn_on_debug() - router = litellm.Router( - model_list=[ - { - "model_name": "gpt4o-special-alias", - "litellm_params": { - "model": "gpt-5.5", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - } - ] - ) - - # Call the handler - if sync_mode: - response = router.responses( - model="gpt4o-special-alias", - input="Hello, can you tell me a short joke?", - max_output_tokens=100, - ) - print("SYNC MODE RESPONSE=", response) - else: - response = await router.aresponses( - model="gpt4o-special-alias", - input="Hello, can you tell me a short joke?", - max_output_tokens=100, - ) - - print( - f"Router {'sync' if sync_mode else 'async'} response=", - json.dumps(response, indent=4, default=str), - ) - - # Use the helper function to validate the response - validate_responses_api_response(response, final_chunk=True) - - return response - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_openai_responses_litellm_router_streaming(sync_mode): - """ - Test the OpenAI responses API with streaming through LiteLLM Router - """ - litellm.turn_on_debug() - router = litellm.Router( - model_list=[ - { - "model_name": "gpt4o-special-alias", - "litellm_params": { - "model": "gpt-5.5", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - } - ] - ) - - event_types_seen = set() - - if sync_mode: - response = router.responses( - model="gpt4o-special-alias", - input="Tell me about artificial intelligence in 2 sentences.", - stream=True, - ) - for event in response: - print(f"Validating event type: {event.type}") - validate_stream_event(event) - event_types_seen.add(event.type) - else: - response = await router.aresponses( - model="gpt4o-special-alias", - input="Tell me about artificial intelligence in 2 sentences.", - stream=True, - ) - async for event in response: - print(f"Validating event type: {event.type}") - validate_stream_event(event) - event_types_seen.add(event.type) - - # At minimum, we should see these core event types - required_events = {"response.created", "response.completed"} - - missing_events = required_events - event_types_seen - assert not missing_events, f"Missing required event types: {missing_events}" - - print(f"Successfully validated all event types: {event_types_seen}") - - -def test_mcp_tools_with_responses_api(): - litellm.turn_on_debug() - MCP_TOOLS = [ - { - "type": "mcp", - "server_label": "zapier", - "server_url": "https://mcp.zapier.com/api/mcp/mcp", - "headers": { - "Authorization": f"Bearer {os.getenv('ZAPIER_CI_CD_MCP_TOKEN')}" - }, - } - ] - MODEL = "openai/gpt-4.1" - USER_QUERY = "how does tiktoken work?" - ######################################################### - # Step 1: OpenAI will use MCP LIST, and return a list of MCP calls for our approval - try: - response = litellm.responses(model=MODEL, tools=MCP_TOOLS, input=USER_QUERY) - print(response) - - response = cast(ResponsesAPIResponse, response) - - mcp_approval_id: Optional[str] = None - for output in response.output: - if output.type == "mcp_approval_request": - mcp_approval_id = output.id - break - - # Step 2: Send followup with approval for the MCP call - if mcp_approval_id: - response_with_mcp_call = litellm.responses( - model=MODEL, - tools=MCP_TOOLS, - input=[ - { - "type": "mcp_approval_response", - "approve": True, - "approval_request_id": mcp_approval_id, - } - ], - previous_response_id=response.id, - ) - print(response_with_mcp_call) - except litellm.APIError as e: - if ( - "424" in str(e) - or "Failed Dependency" in str(e) - or "external_connector_error" in str(e) - ): - pytest.skip(f"Skipping test due to external MCP server error: {e}") - else: - raise e - except litellm.InternalServerError as e: - if "500" in str(e) or "server_error" in str(e): - pytest.skip( - f"Skipping test due to OpenAI server error (likely MCP server unavailable): {e}" - ) - else: - raise e - - -@pytest.mark.asyncio -async def test_openai_responses_api_field_types(): - """Test that specific fields in the response have the correct types""" - litellm.turn_on_debug() - litellm.set_verbose = True - - # Test with store=True - response = await litellm.aresponses( - model="gpt-5.5", - input="hi", - ) - - # Verify created_at is an integer - assert isinstance(response.created_at, int), "created_at should be an integer" - - # Verify store field is present and matches input - assert hasattr(response, "store"), "store field should be present" - assert response.store is True, "store field should match input value" - - # Test without store parameter - response_without_store = await litellm.aresponses(model="gpt-5.5", input="hi") - - # Verify created_at is still an integer - assert isinstance( - response_without_store.created_at, int - ), "created_at should be an integer" - - # Verify store field is present but None when not specified - assert hasattr(response_without_store, "store"), "store field should be present" - - -@pytest.mark.asyncio -async def test_openai_responses_api_token_limit_error(): - """ - Relevant issue: https://github.com/BerriAI/litellm/issues/15785 - - Parsing the in-stream ErrorEvent must not raise - "pydantic_core._pydantic_core.ValidationError: 3 validation errors for ErrorEvent". - The iterator routes the event through litellm.exception_type, so it surfaces as - the typed 400 client error the non-streaming path raises (litellm.BadRequestError) - carrying the provider's message. invalid_request_error is a non-retriable client - error, so there is no MidStreamFallbackError wrapping. - """ - litellm.turn_on_debug() - - # Generate text with >400k tokens to trigger token limit error - oversized_text = "This is a test sentence. " * 50000 # ~400k tokens - - response = await litellm.aresponses( - model="gpt-5-mini", input=oversized_text, stream=True - ) - - async def _drain(): - async for event in response: - print(event) - - with pytest.raises(litellm.BadRequestError) as exc_info: - await _drain() - - assert exc_info.value.status_code == 400 - assert "exceeds the context window" in str(exc_info.value) - - -async def test_openai_streaming_logging(): - """Test that OpenAI Responses API streaming logging is working correctly.""" - litellm.turn_on_debug() - from litellm.integrations.custom_logger import CustomLogger - from litellm.types.utils import Usage - - class TestCustomLogger(CustomLogger): - validate_usage = False - - def __init__(self): - self.standard_logging_object: Optional[StandardLoggingPayload] = None - - async def async_log_success_event( - self, kwargs, response_obj, start_time, end_time - ): - print(f"response_obj: {response_obj.usage}") - assert isinstance( - response_obj.usage, (Usage, dict) - ), f"Expected response_obj.usage to be of type Usage or dict, but got {type(response_obj.usage)}" - # Verify it has the chat completion format fields - if isinstance(response_obj.usage, dict): - assert ( - "prompt_tokens" in response_obj.usage - ), "Usage dict should have prompt_tokens" - assert ( - "completion_tokens" in response_obj.usage - ), "Usage dict should have completion_tokens" - print("\n\nVALIDATED USAGE\n\n") - self.validate_usage = True - - tcl = TestCustomLogger() - litellm.callbacks = [tcl] - request_model = "gpt-5-mini" - response = await litellm.aresponses( - model=request_model, - input="What is the capital of France?", - stream=True, - ) - print("response=", json.dumps(response, indent=4, default=str)) - - async for event in response: - if event.type == "response.completed": - final_response = event - print("litellm response=", json.dumps(event, indent=4, default=str)) - - await asyncio.sleep(2) - assert tcl.validate_usage, "Usage should be validated" - - - - -@pytest.mark.asyncio -@pytest.mark.parametrize("sync_mode", [True, False]) -async def test_openai_compact_responses_api(sync_mode): - """ - Test the compact_responses API for OpenAI. - - This test verifies that the compact_responses endpoint works correctly - for compressing conversation history. - """ - litellm.turn_on_debug() - litellm.set_verbose = True - - input_messages = [ - {"role": "user", "content": "Hello, how are you?"}, - {"role": "assistant", "content": "I'm doing well, thank you for asking!"}, - {"role": "user", "content": "What is the weather like today?"}, - ] - - try: - if sync_mode: - response = litellm.compact_responses( - model="openai/gpt-5.5", - input=input_messages, - instructions="Be helpful and concise", - ) - else: - response = await litellm.acompact_responses( - model="openai/gpt-5.5", - input=input_messages, - instructions="Be helpful and concise", - ) - except litellm.InternalServerError: - pytest.skip("Skipping test due to InternalServerError") - except litellm.BadRequestError as e: - # compact_responses may not be available for all models/accounts - pytest.skip(f"Skipping test due to BadRequestError: {e}") - - print("compact_responses response=", json.dumps(response, indent=4, default=str)) - - # Validate response structure - assert response is not None - assert "id" in response, "Response should have an 'id' field" - assert "output" in response, "Response should have an 'output' field" - assert isinstance(response["output"], list), "Output should be a list" diff --git a/tests/llm_translation/Readme.md b/tests/llm_translation/Readme.md index f0a32f6c989..aef28d38782 100644 --- a/tests/llm_translation/Readme.md +++ b/tests/llm_translation/Readme.md @@ -26,7 +26,6 @@ provider APIs. The reusable conftest plumbing lives in `tests/_vcr_conftest_common.py` and is wired into: - `tests/llm_translation/` -- `tests/llm_responses_api_testing/` - `tests/audio_tests/` - `tests/batches_tests/` - `tests/guardrails_tests/` diff --git a/tests/llm_translation/base_audio_transcription_unit_tests.py b/tests/llm_translation/base_audio_transcription_unit_tests.py index 0234c05f853..bdf9f998d57 100644 --- a/tests/llm_translation/base_audio_transcription_unit_tests.py +++ b/tests/llm_translation/base_audio_transcription_unit_tests.py @@ -1,20 +1,11 @@ import httpx import json -import pytest from typing import Any, Dict, List from unittest.mock import MagicMock, Mock, patch import os from litellm._uuid import uuid import litellm -from litellm import transcription -from litellm.litellm_core_utils.get_supported_openai_params import ( - get_supported_openai_params, -) -from litellm.llms.base_llm.audio_transcription.transformation import ( - BaseAudioTranscriptionConfig, -) -from litellm.utils import ProviderConfigManager from abc import ABC, abstractmethod pwd = os.path.dirname(os.path.realpath(__file__)) @@ -24,7 +15,6 @@ file_path = os.path.join(pwd, "gettysburg.wav") audio_file = open(file_path, "rb") - class BaseLLMAudioTranscriptionTest(ABC): @abstractmethod def get_base_audio_transcription_call_args(self) -> dict: @@ -36,65 +26,3 @@ class BaseLLMAudioTranscriptionTest(ABC): """Must return the custom llm provider""" pass - def test_audio_transcription(self): - """ - Test that the audio transcription is translated correctly. - """ - litellm.set_verbose = True - transcription_call_args = self.get_base_audio_transcription_call_args() - transcript = transcription(**transcription_call_args, file=audio_file) - print(f"transcript: {transcript.model_dump()}") - print(f"transcript hidden params: {transcript._hidden_params}") - - assert transcript.text is not None - - @pytest.mark.asyncio - async def test_audio_transcription_async(self): - """ - Test that the audio transcription is translated correctly. - """ - - litellm.set_verbose = True - litellm.turn_on_debug() - AUDIO_FILE = open(file_path, "rb") - transcription_call_args = self.get_base_audio_transcription_call_args() - transcript = await litellm.atranscription( - **transcription_call_args, file=AUDIO_FILE - ) - print(f"transcript: {transcript.model_dump()}") - print(f"transcript hidden params: {transcript._hidden_params}") - - assert transcript.text is not None - - def test_audio_transcription_optional_params(self): - """ - Test that the audio transcription is translated correctly. - """ - transcription_args = self.get_base_audio_transcription_call_args() - model = transcription_args["model"] - custom_llm_provider = self.get_custom_llm_provider() - optional_params = get_supported_openai_params( - model=model, - custom_llm_provider=custom_llm_provider.value, - request_type="transcription", - ) - print(f"optional_params: {optional_params}") - assert optional_params is not None - assert ( - "max_completion_tokens" not in optional_params - ) # assert default chat completion response not returned - - def test_audio_transcription_config(self): - """ - Test that the audio transcription config is implemented and correctly instrumented. - """ - transcription_args = self.get_base_audio_transcription_call_args() - model = transcription_args["model"] - custom_llm_provider = self.get_custom_llm_provider() - config = ProviderConfigManager.get_provider_audio_transcription_config( - model=model, - provider=custom_llm_provider, - ) - print(f"config: {config}") - assert config is not None - assert isinstance(config, BaseAudioTranscriptionConfig) diff --git a/tests/llm_translation/base_embedding_unit_tests.py b/tests/llm_translation/base_embedding_unit_tests.py index 469416fc0cf..7efd08896c8 100644 --- a/tests/llm_translation/base_embedding_unit_tests.py +++ b/tests/llm_translation/base_embedding_unit_tests.py @@ -4,17 +4,14 @@ import json import pytest from typing import Any, Dict, List from unittest.mock import MagicMock, Mock, patch -import os import litellm -from litellm import embedding from litellm.exceptions import BadRequestError from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.utils import ( CustomStreamWrapper, get_supported_openai_params, get_optional_params, - get_optional_params_embeddings, ) import base64 from pathlib import Path @@ -26,7 +23,6 @@ file_data = (Path(__file__).parent.parent / "white_100x100.png").read_bytes() encoded_file = base64.b64encode(file_data).decode("utf-8") base64_image = f"data:image/png;base64,{encoded_file}" - class BaseLLMEmbeddingTest(ABC): """ Abstract base test class that enforces a common test across all test classes. @@ -66,23 +62,3 @@ class BaseLLMEmbeddingTest(ABC): CreateEmbeddingResponse.model_validate(response.model_dump()) - def test_embedding_optional_params_max_retries(self): - embedding_call_args = self.get_base_embedding_call_args() - optional_params = get_optional_params_embeddings( - **embedding_call_args, max_retries=20 - ) - assert optional_params["max_retries"] == 20 - - def test_image_embedding(self): - litellm.set_verbose = True - from litellm.utils import supports_embedding_image_input - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - base_embedding_call_args = self.get_base_embedding_call_args() - if not supports_embedding_image_input(base_embedding_call_args["model"], None): - print("Model does not support embedding image input") - pytest.skip("Model does not support embedding image input") - - embedding(**base_embedding_call_args, input=[base64_image]) diff --git a/tests/llm_translation/base_llm_unit_tests.py b/tests/llm_translation/base_llm_unit_tests.py index e744b6275a8..d8d53dcc45c 100644 --- a/tests/llm_translation/base_llm_unit_tests.py +++ b/tests/llm_translation/base_llm_unit_tests.py @@ -1,11 +1,9 @@ -import httpx import json import pytest import sys from typing import Any, Dict, List -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import MagicMock, Mock import os -import base64 import inspect import litellm @@ -14,10 +12,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.utils import ( CustomStreamWrapper, get_supported_openai_params, - get_optional_params, - ProviderConfigManager, ) -from litellm.main import stream_chunk_builder from typing import Union from litellm.types.utils import Usage, ModelResponse @@ -27,8 +22,6 @@ from openai import OpenAI sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))) - - class BaseLLMChatTest(ABC): """ Abstract base test class that enforces a common test across all test classes. @@ -61,155 +54,6 @@ class BaseLLMChatTest(ABC): except litellm.InternalServerError: pytest.skip("Model is overloaded") - def test_developer_role_translation(self): - """ - Test that the developer role is translated correctly for non-OpenAI providers. - - Translate `developer` role to `system` role for non-OpenAI providers. - """ - base_completion_call_args = self.get_base_completion_call_args() - messages = [ - { - "role": "developer", - "content": "Be a good bot!", - }, - { - "role": "user", - "content": [{"type": "text", "text": "Hello, how are you?"}], - }, - ] - try: - response = self.completion_function( - **base_completion_call_args, - messages=messages, - ) - assert response is not None - except litellm.InternalServerError: - pytest.skip("Model is overloaded") - - assert response.choices[0].message.content is not None - - def test_content_list_handling(self): - """Check if content list is supported by LLM API""" - base_completion_call_args = self.get_base_completion_call_args() - messages = [ - { - "role": "user", - "content": [{"type": "text", "text": "Hello, how are you?"}], - } - ] - try: - response = self.completion_function( - **base_completion_call_args, - messages=messages, - ) - assert response is not None - except litellm.InternalServerError: - pytest.skip("Model is overloaded") - - # for OpenAI the content contains the JSON schema, so we need to assert that the content is not None - assert response.choices[0].message.content is not None - - def test_tool_call_with_property_type_array(self): - litellm.turn_on_debug() - from litellm.utils import supports_function_calling - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - base_completion_call_args = self.get_base_completion_call_args() - if not supports_function_calling(base_completion_call_args["model"], None): - print("Model does not support function calling") - pytest.skip("Model does not support function calling") - base_completion_call_args = self.get_base_completion_call_args() - response = self.completion_function( - **base_completion_call_args, - messages=[ - { - "role": "user", - "content": "Tell me if the shoe brand Air Jordan has more models than the shoe brand Nike.", - } - ], - tools=[ - { - "type": "function", - "function": { - "name": "shoe_get_id", - "description": "Get information about a show by its ID or name", - "parameters": { - "type": "object", - "properties": { - "shoe_id": { - "type": ["string", "number"], - "description": "The shoe ID or name", - } - }, - "required": ["shoe_id"], - "additionalProperties": False, - "$schema": "http://json-schema.org/draft-07/schema#", - }, - }, - }, - ], - ) - print(response) - print(json.dumps(response, indent=4, default=str)) - - @pytest.mark.flaky(retries=3, delay=1) - def test_tool_call_with_empty_enum_property(self): - litellm.turn_on_debug() - from litellm.utils import supports_function_calling - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - base_completion_call_args = self.get_base_completion_call_args() - if not supports_function_calling(base_completion_call_args["model"], None): - print("Model does not support function calling") - pytest.skip("Model does not support function calling") - base_completion_call_args = self.get_base_completion_call_args() - response = self.completion_function( - **base_completion_call_args, - messages=[ - { - "role": "user", - "content": "Search for the latest iPhone models and tell me which storage options are available.", - } - ], - tools=[ - { - "type": "function", - "function": { - "name": "litellm_product_search", - "description": "Search for product information and specifications.\n\nSupports filtering by category, brand, price range, and availability.\nCan retrieve detailed product specifications, pricing, and stock information.\nSupports different search modes and result formatting options.\n", - "parameters": { - "properties": { - "search_mode": { - "default": "", - "description": "The search strategy to use for finding products.", - "enum": [ - "", - "product_search", - "product_search_with_filters", - "product_search_with_sorting", - "product_search_with_pagination", - "product_search_with_aggregation", - ], - "title": "Search Mode", - "type": "string", - }, - }, - "required": ["search_mode"], - "title": "product_search_arguments", - "type": "object", - }, - }, - } - ], - ) - print(response) - print(json.dumps(response, indent=4, default=str)) - def test_streaming(self): """Check if litellm handles streaming correctly""" from litellm.types.utils import ModelResponseStream @@ -252,16 +96,6 @@ class BaseLLMChatTest(ABC): # assert resp.usage.completion_tokens > 0 # assert resp.usage.total_tokens > 0 - def test_pydantic_model_input(self): - litellm.set_verbose = True - - from litellm import completion, Message - - base_completion_call_args = self.get_base_completion_call_args() - messages = [Message(content="Hello, how are you?", role="user")] - - self.completion_function(**base_completion_call_args, messages=messages) - def test_web_search(self): from litellm.utils import supports_web_search @@ -388,407 +222,6 @@ class BaseLLMChatTest(ABC): assert response is not None - def test_file_data_unit_test(self, pdf_messages): - from litellm.utils import supports_pdf_input, return_raw_request - from litellm.types.utils import CallTypes - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_anthropic_image_obj, - ) - - media_chunk = convert_to_anthropic_image_obj( - openai_image_url=pdf_messages, - format=None, - ) - - file_content = [ - {"type": "text", "text": "What's this file about?"}, - { - "type": "file", - "file": { - "file_data": pdf_messages, - }, - }, - ] - - image_messages = [{"role": "user", "content": file_content}] - - base_completion_call_args = self.get_base_completion_call_args() - - if not supports_pdf_input(base_completion_call_args["model"], None): - pytest.skip("Model does not support image input") - - raw_request = return_raw_request( - endpoint=CallTypes.completion, - kwargs={**base_completion_call_args, "messages": image_messages}, - ) - - print("RAW REQUEST", raw_request) - - assert media_chunk["data"] in json.dumps(raw_request) - - def test_message_with_name(self): - try: - litellm.set_verbose = True - base_completion_call_args = self.get_base_completion_call_args() - messages = [ - {"role": "user", "content": "Hello", "name": "test_name"}, - ] - response = self.completion_function( - **base_completion_call_args, messages=messages - ) - assert response is not None - except litellm.RateLimitError: - pass - - @pytest.mark.parametrize( - "response_format", - [ - {"type": "json_object"}, - {"type": "text"}, - ], - ) - @pytest.mark.flaky(retries=6, delay=1) - def test_json_response_format(self, response_format): - """ - Test that the JSON response format is supported by the LLM API - """ - from litellm.utils import supports_response_schema - - base_completion_call_args = self.get_base_completion_call_args() - litellm.set_verbose = True - - if not supports_response_schema(base_completion_call_args["model"], None): - pytest.skip("Model does not support response schema") - - messages = [ - { - "role": "system", - "content": "Your output should be a JSON object with no additional properties. ", - }, - { - "role": "user", - "content": "Respond with this in json. city=San Francisco, state=CA, weather=sunny, temp=60", - }, - ] - - response = self.completion_function( - **base_completion_call_args, - messages=messages, - response_format=response_format, - ) - - print(f"response={response}") - - # OpenAI guarantees that the JSON schema is returned in the content - # relevant issue: https://github.com/BerriAI/litellm/issues/6741 - assert response.choices[0].message.content is not None - - @pytest.mark.parametrize( - "response_format", - [ - {"type": "text"}, - ], - ) - @pytest.mark.flaky(retries=6, delay=1) - def test_response_format_type_text_with_tool_calls_no_tool_choice( - self, response_format - ): - base_completion_call_args = self.get_base_completion_call_args() - messages = [ - {"role": "user", "content": "What's the weather like in Boston today?"}, - ] - tools = [ - { - "type": "function", - "function": { - "name": "get_current_weather", - "description": "Get the current weather in a given location", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "The city and state, e.g. San Francisco, CA", - }, - "unit": { - "type": "string", - "enum": ["celsius", "fahrenheit"], - }, - }, - "required": ["location"], - }, - }, - } - ] - try: - print(f"MAKING LLM CALL") - response = self.completion_function( - **base_completion_call_args, - messages=messages, - response_format=response_format, - tools=tools, - drop_params=True, - ) - print(f"RESPONSE={response}") - except litellm.ContextWindowExceededError: - pytest.skip("Model exceeded context window") - assert response is not None - - def test_response_format_type_text(self): - """ - Test that the response format type text does not lead to tool calls - """ - from litellm import LlmProviders - - base_completion_call_args = self.get_base_completion_call_args() - litellm.set_verbose = True - - _, provider, _, _ = litellm.get_llm_provider( - model=base_completion_call_args["model"] - ) - - provider_config = ProviderConfigManager.get_provider_chat_config( - base_completion_call_args["model"], LlmProviders(provider) - ) - - print(f"provider_config={provider_config}") - - translated_params = provider_config.map_openai_params( - non_default_params={"response_format": {"type": "text"}}, - optional_params={}, - model=base_completion_call_args["model"], - drop_params=False, - ) - - assert "tool_choice" not in translated_params - assert ( - "tools" not in translated_params - ), f"Got tools={translated_params['tools']}, expected no tools" - - print(f"translated_params={translated_params}") - - @pytest.mark.flaky(retries=6, delay=1) - def test_json_response_pydantic_obj(self): - litellm.turn_on_debug() - from pydantic import BaseModel - from litellm.utils import supports_response_schema - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - class TestModel(BaseModel): - first_response: str - - base_completion_call_args = self.get_base_completion_call_args() - if not supports_response_schema(base_completion_call_args["model"], None): - pytest.skip("Model does not support response schema") - - try: - res = self.completion_function( - **base_completion_call_args, - messages=[ - {"role": "system", "content": "You are a helpful assistant."}, - { - "role": "user", - "content": "What is the capital of France?", - }, - ], - response_format=TestModel, - timeout=5, - ) - assert res is not None - - print(res.choices[0].message) - - assert res.choices[0].message.content is not None - assert res.choices[0].message.tool_calls is None - except litellm.Timeout: - pytest.skip("Model took too long to respond") - except litellm.InternalServerError: - pytest.skip("Model is overloaded") - - @pytest.mark.flaky(retries=6, delay=1) - def test_json_response_nested_pydantic_obj(self): - from pydantic import BaseModel - from litellm.utils import supports_response_schema - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - class CalendarEvent(BaseModel): - name: str - date: str - participants: list[str] - - class EventsList(BaseModel): - events: list[CalendarEvent] - - messages = [ - {"role": "user", "content": "List 5 important events in the XIX century"} - ] - - base_completion_call_args = self.get_base_completion_call_args() - if not supports_response_schema(base_completion_call_args["model"], None): - pytest.skip( - f"Model={base_completion_call_args['model']} does not support response schema" - ) - - try: - res = self.completion_function( - **base_completion_call_args, - messages=messages, - response_format=EventsList, - timeout=60, - ) - assert res is not None - - print(res.choices[0].message) - - assert res.choices[0].message.content is not None - assert res.choices[0].message.tool_calls is None - except litellm.Timeout: - pytest.skip("Model took too long to respond") - except litellm.InternalServerError: - pytest.skip("Model is overloaded") - - @pytest.mark.flaky(retries=6, delay=1) - def test_json_response_nested_json_schema(self): - """ - PROD Test: ensure nested json schema sent to proxy works as expected. - """ - litellm.turn_on_debug() - from pydantic import BaseModel - from litellm.utils import supports_response_schema - from litellm.llms.base_llm.base_utils import type_to_response_format_param - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - class CalendarEvent(BaseModel): - name: str - date: str - participants: list[str] - - class EventsList(BaseModel): - events: list[CalendarEvent] - - response_format = type_to_response_format_param(EventsList) - - messages = [ - {"role": "user", "content": "List 5 important events in the XIX century"} - ] - - base_completion_call_args = self.get_base_completion_call_args() - if not supports_response_schema(base_completion_call_args["model"], None): - pytest.skip( - f"Model={base_completion_call_args['model']} does not support response schema" - ) - - try: - res = self.completion_function( - **base_completion_call_args, - messages=messages, - response_format=response_format, - timeout=60, - ) - assert res is not None - - print(res.choices[0].message) - - assert res.choices[0].message.content is not None - assert res.choices[0].message.tool_calls is None - except litellm.Timeout: - pytest.skip("Model took too long to respond") - except litellm.InternalServerError: - pytest.skip("Model is overloaded") - - @pytest.mark.flaky(retries=6, delay=1) - def test_audio_input(self): - """ - Test that audio input is supported by the LLM API - """ - from litellm.utils import supports_audio_input - - litellm.turn_on_debug() - base_completion_call_args = self.get_base_completion_call_args() - if not supports_audio_input(base_completion_call_args["model"], None): - pytest.skip( - f"Model={base_completion_call_args['model']} does not support audio input" - ) - - url = "https://openaiassets.blob.core.windows.net/$web/API/docs/audio/alloy.wav" - response = httpx.get(url) - response.raise_for_status() - wav_data = response.content - encoded_string = base64.b64encode(wav_data).decode("utf-8") - - completion = self.completion_function( - **base_completion_call_args, - messages=[ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is in this recording?"}, - { - "type": "input_audio", - "input_audio": {"data": encoded_string, "format": "wav"}, - }, - ], - }, - ], - ) - - print(completion.choices[0].message) - - @pytest.mark.flaky(retries=6, delay=1) - def test_json_response_format_stream(self): - """ - Test that the JSON response format with streaming is supported by the LLM API - """ - from litellm.utils import supports_response_schema - - base_completion_call_args = self.get_base_completion_call_args() - litellm.set_verbose = True - - base_completion_call_args = self.get_base_completion_call_args() - if not supports_response_schema(base_completion_call_args["model"], None): - pytest.skip("Model does not support response schema") - - messages = [ - { - "role": "system", - "content": "Your output should be a JSON object with no additional properties. ", - }, - { - "role": "user", - "content": "Respond with this in json. city=San Francisco, state=CA, weather=sunny, temp=60", - }, - ] - - try: - response = self.completion_function( - **base_completion_call_args, - messages=messages, - response_format={"type": "json_object"}, - stream=True, - ) - except litellm.InternalServerError: - pytest.skip("Model is overloaded") - - print(response) - - content = "" - for chunk in response: - content += chunk.choices[0].delta.content or "" - - print(f"content={content}") - - # OpenAI guarantees that the JSON schema is returned in the content - # relevant issue: https://github.com/BerriAI/litellm/issues/6741 - # we need to assert that the JSON schema was returned in the content, (for Anthropic we were returning it as part of the tool call) - assert content is not None - assert len(content) > 0 - @pytest.fixture def tool_call_no_arguments(self): return { @@ -803,118 +236,6 @@ class BaseLLMChatTest(ABC): ], } - @pytest.mark.parametrize("detail", [None, "low", "high"]) - @pytest.mark.parametrize( - "image_url", - [ - # In-repo logo served via jsdelivr (sha-pinned, immutable). - # Bedrock fetches the URL and base64-embeds it in the - # Converse request body; using a multi-MB hosted product - # photo here previously bloated cassettes to ~60 MB each. - "https://cdn.jsdelivr.net/gh/BerriAI/litellm@d769e81c90d453240c61fc572cdb27fae06a89d0/ui/litellm-dashboard/public/assets/logos/litellm_logo.jpg", - "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png", - ], - ) - @pytest.mark.flaky(retries=4, delay=2) - def test_image_url(self, detail, image_url): - litellm.set_verbose = True - from litellm.utils import supports_vision - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - base_completion_call_args = self.get_base_completion_call_args() - if not supports_vision(base_completion_call_args["model"], None): - pytest.skip("Model does not support image input") - elif "http://" in image_url and ( - "fireworks_ai" in base_completion_call_args.get("model", "") - or "mistral" in base_completion_call_args.get("model", "") - ): - pytest.skip("Model does not support http:// input") - - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What's in this image?"}, - { - "type": "image_url", - "image_url": { - "url": image_url, - }, - }, - ], - } - ] - - if detail is not None: - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What's in this image?"}, - { - "type": "image_url", - "image_url": { - # sha-pinned in-repo logo via jsdelivr; gstatic's - # robots.txt blocks server-side fetchers (e.g. - # Anthropic), which 400s the request. - "url": "https://cdn.jsdelivr.net/gh/BerriAI/litellm@d769e81c90d453240c61fc572cdb27fae06a89d0/ui/litellm-dashboard/public/assets/logos/litellm_logo.jpg", - "detail": detail, - }, - }, - ], - } - ] - try: - response = self.completion_function( - **base_completion_call_args, messages=messages - ) - except litellm.InternalServerError: - pytest.skip("Model is overloaded") - - assert response is not None - - def test_image_url_string(self): - litellm.set_verbose = True - from litellm.utils import supports_vision - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - image_url = "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png" - - base_completion_call_args = self.get_base_completion_call_args() - if not supports_vision(base_completion_call_args["model"], None): - pytest.skip("Model does not support image input") - elif "http://" in image_url and "fireworks_ai" in base_completion_call_args.get( - "model" - ): - pytest.skip("Model does not support http:// input") - - image_url_param = image_url - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What's in this image?"}, - { - "type": "image_url", - "image_url": image_url_param, - }, - ], - } - ] - - try: - response = self.completion_function( - **base_completion_call_args, messages=messages - ) - except litellm.InternalServerError: - pytest.skip("Model is overloaded") - - assert response is not None - @pytest.fixture def pdf_messages(self): import base64 @@ -932,42 +253,6 @@ class BaseLLMChatTest(ABC): return url - @pytest.mark.flaky(retries=3, delay=1) - def test_empty_tools(self): - """ - Related Issue: https://github.com/BerriAI/litellm/issues/9080 - """ - try: - from litellm import completion, ModelResponse - - litellm.set_verbose = True - litellm.turn_on_debug() - from litellm.utils import supports_function_calling - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - base_completion_call_args = self.get_base_completion_call_args() - if not supports_function_calling(base_completion_call_args["model"], None): - print("Model does not support function calling") - pytest.skip("Model does not support function calling") - - response = completion( - **base_completion_call_args, - messages=[{"role": "user", "content": "Hello, how are you?"}], - tools=[], - ) # just make sure call doesn't fail - print("response: ", response) - assert response is not None - except litellm.ContentPolicyViolationError: - pass - except litellm.InternalServerError: - pytest.skip("Model is overloaded") - except litellm.RateLimitError: - pass - except Exception as e: - pytest.fail(f"Error occurred: {e}") - @pytest.mark.flaky(retries=3, delay=1) def test_basic_tool_calling(self): try: @@ -1091,94 +376,6 @@ class BaseLLMChatTest(ABC): except Exception as e: pytest.fail(f"Error occurred: {e}") - @pytest.mark.flaky(retries=3, delay=1) - @pytest.mark.asyncio - async def test_completion_cost(self): - from litellm import completion_cost - - litellm.turn_on_debug() - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - litellm.set_verbose = True - response = await self.async_completion_function( - **self.get_base_completion_call_args(), - messages=[{"role": "user", "content": "Hello, how are you?"}], - ) - print(response._hidden_params["response_cost"]) - - assert response._hidden_params["response_cost"] > 0 - - @pytest.mark.parametrize("input_type", ["input_audio", "audio_url"]) - @pytest.mark.parametrize("format_specified", [True]) - def test_supports_audio_input(self, input_type, format_specified): - from litellm.utils import return_raw_request, supports_audio_input - from litellm.types.utils import CallTypes - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - litellm.drop_params = True - base_completion_call_args = self.get_base_completion_call_args() - if not supports_audio_input(base_completion_call_args["model"], None): - print("Model does not support audio input") - pytest.skip("Model does not support audio input") - - url = "https://openaiassets.blob.core.windows.net/$web/API/docs/audio/alloy.wav" - response = httpx.get(url) - response.raise_for_status() - wav_data = response.content - audio_format = "wav" - encoded_string = base64.b64encode(wav_data).decode("utf-8") - - audio_content = [{"type": "text", "text": "What is in this recording?"}] - - test_file_id = "gs://bucket/file.wav" - - if input_type == "input_audio": - audio_content.append( - { - "type": "input_audio", - "input_audio": {"data": encoded_string, "format": audio_format}, - } - ) - elif input_type == "audio_url": - audio_content.append( - { - "type": "file", - "file": { - "file_id": test_file_id, - "filename": "my-sample-audio-file", - }, - } - ) - - raw_request = return_raw_request( - endpoint=CallTypes.completion, - kwargs={ - **base_completion_call_args, - "modalities": ["text", "audio"], - "audio": {"voice": "alloy", "format": audio_format}, - "messages": [ - { - "role": "user", - "content": audio_content, - }, - ], - }, - ) - print("raw_request: ", raw_request) - - if input_type == "input_audio": - assert encoded_string in json.dumps( - raw_request - ), "Audio data not sent to gemini" - elif input_type == "audio_url": - assert test_file_id in json.dumps( - raw_request - ), "Audio URL not sent to gemini" - def test_function_calling_with_tool_response(self): from litellm.utils import supports_function_calling from litellm import completion @@ -1277,53 +474,6 @@ class BaseLLMChatTest(ABC): except litellm.ServiceUnavailableError: pass - def test_reasoning_effort(self): - """Test that reasoning_effort is passed correctly to the model""" - from litellm.utils import supports_reasoning - from litellm import completion - - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - base_completion_call_args = ( - self.get_base_completion_call_args_with_reasoning_model() - ) - if len(base_completion_call_args) == 0: - print("base_completion_call_args is empty") - pytest.skip("Model does not support reasoning") - if not supports_reasoning(base_completion_call_args["model"], None): - print("Model does not support reasoning") - pytest.skip("Model does not support reasoning") - - _, provider, _, _ = litellm.get_llm_provider( - model=base_completion_call_args["model"] - ) - - ## CHECK PARAM MAPPING - optional_params = get_optional_params( - model=base_completion_call_args["model"], - custom_llm_provider=provider, - reasoning_effort="high", - ) - # either accepts reasoning effort or thinking budget - from litellm.constants import DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET - - assert "reasoning_effort" in optional_params or str( - DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET - ) in json.dumps(optional_params) - - try: - litellm.turn_on_debug() - response = completion( - **base_completion_call_args, - reasoning_effort="low", - messages=[{"role": "user", "content": "Hello!"}], - ) - print(f"response: {response}") - except Exception as e: - pytest.fail(f"Error: {e}") - - class BaseOSeriesModelsTest(ABC): # test across azure/openai @abstractmethod def get_base_completion_call_args(self): @@ -1333,105 +483,6 @@ class BaseOSeriesModelsTest(ABC): # test across azure/openai def get_client(self) -> OpenAI: pass - def test_reasoning_effort(self): - """Test that reasoning_effort is passed correctly to the model""" - - from litellm import completion - - client = self.get_client() - - completion_args = self.get_base_completion_call_args() - - with patch.object( - client.chat.completions.with_raw_response, "create" - ) as mock_client: - try: - completion( - **completion_args, - reasoning_effort="low", - messages=[{"role": "user", "content": "Hello!"}], - client=client, - ) - except Exception as e: - print(f"Error: {e}") - - mock_client.assert_called_once() - request_body = mock_client.call_args.kwargs - print("request_body: ", request_body) - assert request_body["reasoning_effort"] == "low" - - def test_developer_role_translation(self): - """Test that developer role is translated correctly to system role for non-OpenAI providers""" - from litellm import completion - - client = self.get_client() - - completion_args = self.get_base_completion_call_args() - - with patch.object( - client.chat.completions.with_raw_response, "create" - ) as mock_client: - try: - completion( - **completion_args, - reasoning_effort="low", - messages=[ - {"role": "developer", "content": "Be a good bot!"}, - {"role": "user", "content": "Hello!"}, - ], - client=client, - ) - except Exception as e: - print(f"Error: {e}") - - mock_client.assert_called_once() - request_body = mock_client.call_args.kwargs - print("request_body: ", request_body) - assert ( - request_body["messages"][0]["role"] == "developer" - ), "Got={} instead of system".format(request_body["messages"][0]["role"]) - assert request_body["messages"][0]["content"] == "Be a good bot!" - - def test_completion_o_series_models_temperature(self): - """ - Test that temperature is not passed to O-series models - """ - try: - from litellm import completion - - client = self.get_client() - - completion_args = self.get_base_completion_call_args() - - with patch.object( - client.chat.completions.with_raw_response, "create" - ) as mock_client: - try: - completion( - **completion_args, - temperature=0.0, - messages=[ - { - "role": "user", - "content": "Hello, world!", - } - ], - drop_params=True, - client=client, - ) - except Exception as e: - print(f"Error: {e}") - - mock_client.assert_called_once() - request_body = mock_client.call_args.kwargs - print("request_body: ", request_body) - assert ( - "temperature" not in request_body - ), "temperature should not be in the request body" - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - class BaseAnthropicChatTest(ABC): """ Ensures consistent result across anthropic model usage @@ -1451,197 +502,6 @@ class BaseAnthropicChatTest(ABC): def completion_function(self): return litellm.completion - def test_anthropic_response_format_streaming_vs_non_streaming(self): - args = { - "messages": [ - { - "content": "Your goal is to summarize the previous agent's thinking process into short descriptions to let user better understand the research progress. If no information is available, just say generic phrase like 'Doing some research...' with the given output format. Make sure to adhere to the output format no matter what, even if you don't have any information or you are not allowed to respond to the given input information (then just say generic phrase like 'Doing some research...').", - "role": "system", - }, - { - "role": "user", - "content": "Here is the input data (previous agent's output): \n\n Let's try to refine our search further, focusing more on the technical aspects of home automation and home energy system management:", - }, - ], - "response_format": { - "type": "json_schema", - "json_schema": { - "name": "final_output", - "strict": True, - "schema": { - "description": 'Progress report for the thinking process\n\nThis model represents a snapshot of the agent\'s current progress during\nthe thinking process, providing a brief description of the current activity.\n\nAttributes:\n agent_doing: Brief description of what the agent is currently doing.\n Should be kept under 10 words. Example: "Learning about home automation"', - "properties": { - "agent_doing": {"title": "Agent Doing", "type": "string"} - }, - "required": ["agent_doing"], - "title": "ThinkingStep", - "type": "object", - "additionalProperties": False, - }, - }, - }, - } - - base_completion_call_args = self.get_base_completion_call_args() - - response = self.completion_function( - **base_completion_call_args, **args, stream=True - ) - - chunks = [] - for chunk in response: - print(f"chunk: {chunk}") - chunks.append(chunk) - - print(f"chunks: {chunks}") - built_response = stream_chunk_builder(chunks=chunks) - - non_stream_response = self.completion_function( - **base_completion_call_args, **args, stream=False - ) - - print( - "built_response.choices[0].message.content", - built_response.choices[0].message.content, - ) - print( - "non_stream_response.choices[0].message.content", - non_stream_response.choices[0].message.content, - ) - assert ( - json.loads(built_response.choices[0].message.content).keys() - == json.loads(non_stream_response.choices[0].message.content).keys() - ), f"Got={json.loads(built_response.choices[0].message.content)}, Expected={json.loads(non_stream_response.choices[0].message.content)}" - - def test_completion_thinking_with_response_format(self): - from pydantic import BaseModel - - litellm.turn_on_debug() - - class RFormat(BaseModel): - question: str - answer: str - - base_completion_call_args = self.get_base_completion_call_args_with_thinking() - - messages = [{"role": "user", "content": "Generate 5 question + answer pairs"}] - response = self.completion_function( - **base_completion_call_args, - messages=messages, - response_format=RFormat, - ) - - print(response) - - def test_completion_thinking_with_max_tokens(self): - from pydantic import BaseModel - - litellm.turn_on_debug() - - base_completion_call_args = self.get_base_completion_call_args_with_thinking() - - messages = [{"role": "user", "content": "Generate 5 question + answer pairs"}] - response = self.completion_function( - **base_completion_call_args, - messages=messages, - max_completion_tokens=20000, - ) - - print(response) - - def test_completion_thinking_without_max_tokens(self): - from pydantic import BaseModel - - litellm.turn_on_debug() - - base_completion_call_args = self.get_base_completion_call_args_with_thinking() - - messages = [{"role": "user", "content": "Generate 5 question + answer pairs"}] - response = self.completion_function( - **base_completion_call_args, - messages=messages, - ) - - print(response) - - def test_completion_with_thinking_basic(self): - litellm.turn_on_debug() - base_completion_call_args = self.get_base_completion_call_args_with_thinking() - - messages = [{"role": "user", "content": "Generate 5 question + answer pairs"}] - response = self.completion_function( - **base_completion_call_args, - messages=messages, - ) - - print(f"response: {response}") - assert response.choices[0].message.reasoning_content is not None - assert isinstance(response.choices[0].message.reasoning_content, str) - assert response.choices[0].message.thinking_blocks is not None - assert isinstance(response.choices[0].message.thinking_blocks, list) - assert len(response.choices[0].message.thinking_blocks) > 0 - - assert response.choices[0].message.thinking_blocks[0]["signature"] is not None - - def test_anthropic_thinking_output_stream(self): - # litellm.set_verbose = True - try: - base_completion_call_args = ( - self.get_base_completion_call_args_with_thinking() - ) - resp = litellm.completion( - **base_completion_call_args, - messages=[{"role": "user", "content": "Tell me a joke."}], - stream=True, - timeout=10, - ) - - reasoning_content_exists = False - signature_block_exists = False - tool_call_exists = False - for chunk in resp: - print(f"chunk 2: {chunk}") - if chunk.choices[0].delta.tool_calls: - tool_call_exists = True - if ( - hasattr(chunk.choices[0].delta, "thinking_blocks") - and chunk.choices[0].delta.thinking_blocks is not None - and chunk.choices[0].delta.reasoning_content is not None - and isinstance(chunk.choices[0].delta.thinking_blocks, list) - and len(chunk.choices[0].delta.thinking_blocks) > 0 - and isinstance(chunk.choices[0].delta.reasoning_content, str) - ): - reasoning_content_exists = True - print(chunk.choices[0].delta.thinking_blocks[0]) - if chunk.choices[0].delta.thinking_blocks[0].get("signature"): - signature_block_exists = True - assert not tool_call_exists - assert reasoning_content_exists - assert signature_block_exists - except litellm.Timeout: - pytest.skip("Model is timing out") - - def test_anthropic_reasoning_effort_thinking_translation(self): - base_completion_call_args = self.get_base_completion_call_args_with_thinking() - _, provider, _, _ = litellm.get_llm_provider( - model=base_completion_call_args["model"] - ) - - from litellm.constants import DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET - - optional_params = get_optional_params( - model=base_completion_call_args.get("model"), - custom_llm_provider=provider, - reasoning_effort="high", - ) - assert optional_params["thinking"] == { - "type": "enabled", - "budget_tokens": DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, - } - - assert "reasoning_effort" not in optional_params - - class BaseReasoningLLMTests(ABC): """ Base class for testing reasoning llms diff --git a/tests/llm_translation/base_rerank_unit_tests.py b/tests/llm_translation/base_rerank_unit_tests.py index e3c66dcc0be..dbf7d0fd1ff 100644 --- a/tests/llm_translation/base_rerank_unit_tests.py +++ b/tests/llm_translation/base_rerank_unit_tests.py @@ -1,10 +1,8 @@ import asyncio import httpx import json -import pytest from typing import Any, Dict, List from unittest.mock import MagicMock, Mock, patch -import os import litellm from litellm.exceptions import BadRequestError @@ -18,7 +16,6 @@ from litellm.utils import ( # test_example.py from abc import ABC, abstractmethod - def assert_response_shape(response, custom_llm_provider): expected_response_shape = {"id": str, "results": list, "meta": dict} @@ -63,7 +60,6 @@ def assert_response_shape(response, custom_llm_provider): expected_billed_units_shape["search_units"], ) - class BaseLLMRerankTest(ABC): """ Abstract base test class that enforces a common test across all test classes. @@ -87,57 +83,3 @@ class BaseLLMRerankTest(ABC): """ return None - @pytest.mark.asyncio() - @pytest.mark.parametrize("sync_mode", [True, False]) - async def test_basic_rerank(self, sync_mode): - litellm.turn_on_debug() - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - rerank_call_args = self.get_base_rerank_call_args() - custom_llm_provider = self.get_custom_llm_provider() - if sync_mode is True: - response = litellm.rerank( - **rerank_call_args, - query="hello", - documents=["hello", "world"], - top_n=2, - ) - - print("re rank response: ", response) - - assert response.id is not None - assert response.results is not None - - assert response._hidden_params["response_cost"] is not None - - # Check expected cost - expected_cost = self.get_expected_cost() - if expected_cost is not None: - # If expected cost is specified, check exact match or >= for 0 - if expected_cost == 0.0: - assert response._hidden_params["response_cost"] >= 0 - else: - assert response._hidden_params["response_cost"] == expected_cost - else: - # Default behavior: cost should be greater than 0 - assert response._hidden_params["response_cost"] > 0 - - assert_response_shape( - response=response, custom_llm_provider=custom_llm_provider.value - ) - else: - response = await litellm.arerank( - **rerank_call_args, - query="hello", - documents=["hello", "world"], - top_n=2, - ) - - print("async re rank response: ", response) - - assert response.id is not None - assert response.results is not None - - assert_response_shape( - response=response, custom_llm_provider=custom_llm_provider.value - ) diff --git a/tests/llm_translation/interactions/base_interactions_test.py b/tests/llm_translation/interactions/base_interactions_test.py index 22ecce3a57b..ed09ecb2b7b 100644 --- a/tests/llm_translation/interactions/base_interactions_test.py +++ b/tests/llm_translation/interactions/base_interactions_test.py @@ -12,7 +12,6 @@ import pytest import litellm.interactions as interactions - class BaseInteractionsTest(ABC): """Abstract base class for interactions API tests. @@ -102,17 +101,3 @@ class BaseInteractionsTest(ABC): assert len(chunks) > 0 - @pytest.mark.asyncio - async def test_acreate_simple(self): - """Test async interaction creation.""" - api_key = self.get_api_key() - if not api_key: - pytest.skip(f"API key not set for {self.__class__.__name__}") - - response = await interactions.acreate( - model=self.get_model(), - input="What is the speed of light?", - api_key=api_key, - ) - assert response is not None - assert response.id is not None or response.status is not None diff --git a/tests/llm_translation/realtime/test_openai_realtime.py b/tests/llm_translation/realtime/test_openai_realtime.py deleted file mode 100644 index f137e9df49b..00000000000 --- a/tests/llm_translation/realtime/test_openai_realtime.py +++ /dev/null @@ -1,138 +0,0 @@ -import os - -import pytest -from websockets.exceptions import ConnectionClosedError, ConnectionClosedOK - -import litellm - - -@pytest.mark.asyncio -@pytest.mark.skipif( - os.environ.get("OPENAI_API_KEY", None) is None, - reason="No OpenAI API key provided", -) -async def test_openai_realtime_direct_call_no_intent(): - """ - End-to-end test calling the actual OpenAI realtime endpoint via LiteLLM SDK - without intent parameter. This should succeed without "Invalid intent" error. - Uses real websocket connection to OpenAI. - """ - import asyncio - import json - - class RealTimeWebSocketClient: - def __init__(self): - self.messages_sent = [] - self.messages_received = [] - self.received_session_created = False - self.connection_successful = False - self._receive_called = False - self.close_code = None - self.close_reason = None - - async def accept(self): - pass - - async def send_text(self, message): - self.messages_sent.append(message) - try: - if isinstance(message, bytes): - message_str = message.decode("utf-8") - else: - message_str = message - - msg_data = json.loads(message_str) - msg_type = msg_data.get("type", "unknown") - - if msg_type == "error": - error_info = msg_data.get("error", {}) - error_code = error_info.get("code", "unknown") - error_message = error_info.get("message", "unknown") - # Don't fail on error, just record it - some errors are expected - self.messages_received.append(msg_data) - return - - if msg_type == "session.created" and not self.received_session_created: - self.messages_received.append(msg_data) - self.received_session_created = True - self.connection_successful = True - except (json.JSONDecodeError, UnicodeDecodeError): - # Non-JSON messages are acceptable - pass - - async def receive_text(self): - if not self._receive_called: - self._receive_called = True - max_wait = 60.0 - check_interval = 0.1 - waited = 0.0 - - while waited < max_wait: - if self.connection_successful: - break - await asyncio.sleep(check_interval) - waited += check_interval - - if not self.connection_successful: - await asyncio.sleep(3.0) - - raise ConnectionClosedOK(None, None) - - async def close(self, code=1000, reason=""): - self.close_code = code - self.close_reason = reason - - @property - def headers(self): - return {} - - websocket_client = RealTimeWebSocketClient() - caught_exception = None - - try: - await litellm._arealtime( - # OpenAI shut down the gpt-4o-realtime-preview family (incl. the - # undated alias) on 2026-05-07; gpt-realtime is the GA successor. - model="openai/gpt-realtime", - websocket=websocket_client, - api_key=os.environ.get("OPENAI_API_KEY"), - timeout=60, - ) - except (ConnectionClosedOK, ConnectionClosedError): - pass - except Exception as e: - caught_exception = e - if "invalid_intent" in str(e).lower(): - pytest.fail(f"Still getting invalid intent error: {e}") - # Other exceptions are recorded but don't fail immediately - - # Build detailed error message for debugging - error_details = [] - error_details.append(f"messages_sent count: {len(websocket_client.messages_sent)}") - error_details.append( - f"messages_received count: {len(websocket_client.messages_received)}" - ) - error_details.append(f"close_code: {websocket_client.close_code}") - error_details.append(f"close_reason: {websocket_client.close_reason}") - if caught_exception: - error_details.append( - f"exception: {type(caught_exception).__name__}: {caught_exception}" - ) - - assert ( - websocket_client.connection_successful - ), f"Failed to establish connection. Debug info: {'; '.join(error_details)}" - assert ( - websocket_client.received_session_created - ), "Did not receive session.created response" - assert len(websocket_client.messages_received) > 0, "No messages received" - - session_message = websocket_client.messages_received[0] - assert ( - session_message["type"] == "session.created" - ), f"Expected session.created, got {session_message.get('type')}" - assert ( - "session" in session_message - ), "session.created response missing session object" - assert "id" in session_message["session"], "Session object missing id field" - assert "model" in session_message["session"], "Session object missing model field" diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py index 5f17b44fd3d..73872728a59 100644 --- a/tests/llm_translation/test_anthropic_completion.py +++ b/tests/llm_translation/test_anthropic_completion.py @@ -32,7 +32,6 @@ from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion from httpx import Headers from base_llm_unit_tests import BaseLLMChatTest, BaseAnthropicChatTest - def streaming_format_tests(chunk: dict, idx: int): """ 1st chunk - chunk.get("type") == "message_start" @@ -46,7 +45,6 @@ def streaming_format_tests(chunk: dict, idx: int): elif idx == 2: assert chunk.get("type") == "content_block_delta" - anthropic_chunk_list = [ { "type": "content_block_start", @@ -234,17 +232,6 @@ anthropic_chunk_list = [ {"type": "message_stop"}, ] - - - - - - - - - - - @pytest.mark.parametrize( "tool_type, tool_config, message_content", [ @@ -295,16 +282,8 @@ def test_anthropic_tool_use(tool_type, tool_config, message_content): except litellm.InternalServerError: pass - - - - - - - from litellm import completion - class TestAnthropicCompletion(BaseLLMChatTest, BaseAnthropicChatTest): def get_base_completion_call_args(self) -> dict: return {"model": "anthropic/claude-sonnet-4-5-20250929"} @@ -380,33 +359,10 @@ class TestAnthropicCompletion(BaseLLMChatTest, BaseAnthropicChatTest): @pytest.mark.asyncio async def test_pdf_handling(self, pdf_messages, sync_mode): await super().test_pdf_handling(pdf_messages, sync_mode) - test_content_list_handling = None - test_image_url = None - test_image_url_string = None test_web_search = None - - - - - - - - - - - - - - - from litellm.constants import RESPONSE_FORMAT_TOOL_NAME - - - - - def test_anthropic_citations_api(): """ Test the citations API @@ -454,7 +410,6 @@ def test_anthropic_citations_api(): assert "start_char_index" in citation assert "end_char_index" in citation - def test_anthropic_citations_api_streaming(): resp = completion( @@ -493,35 +448,6 @@ def test_anthropic_citations_api_streaming(): assert has_citations - -@pytest.mark.parametrize( - "model", - [ - "anthropic/claude-sonnet-4-5-20250929", - "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - ], -) -def test_anthropic_thinking_output(model): - - litellm.turn_on_debug() - - resp = completion( - model=model, - messages=[{"role": "user", "content": "What is the capital of France?"}], - thinking={"type": "enabled", "budget_tokens": 1024}, - ) - - print(resp) - assert resp.choices[0].message.reasoning_content is not None - assert isinstance(resp.choices[0].message.reasoning_content, str) - assert resp.choices[0].message.thinking_blocks is not None - assert isinstance(resp.choices[0].message.thinking_blocks, list) - assert len(resp.choices[0].message.thinking_blocks) > 0 - - assert resp.choices[0].message.thinking_blocks[0]["type"] == "thinking" - assert resp.choices[0].message.thinking_blocks[0]["signature"] is not None - - @pytest.mark.parametrize( "model", [ @@ -566,7 +492,6 @@ def test_anthropic_thinking_output_stream(model): except litellm.Timeout: pytest.skip("Model is timing out") - def test_anthropic_custom_headers(): from litellm.llms.custom_httpx.http_handler import HTTPHandler @@ -604,9 +529,6 @@ def test_anthropic_custom_headers(): headers = mock_post.call_args[1]["headers"] assert "computer-use-2025-01-24" in headers["anthropic-beta"] - - - @pytest.mark.parametrize( "optional_params", [ @@ -645,7 +567,6 @@ def test_anthropic_websearch(optional_params: dict): assert response.usage.server_tool_use is not None assert response.usage.server_tool_use.web_search_requests >= 1 - def test_anthropic_text_editor(): litellm.turn_on_debug() params = { @@ -668,7 +589,6 @@ def test_anthropic_text_editor(): assert response is not None - @pytest.mark.parametrize("spec", ["anthropic", "openai"]) @pytest.mark.skipif( os.getenv("ZAPIER_CI_CD_MCP_TOKEN") is None, reason="ZAPIER_CI_CD_MCP_TOKEN not set" @@ -710,7 +630,6 @@ def test_anthropic_mcp_server_tool_use(spec: str): except litellm.InternalServerError as e: pytest.skip(f"Skipping test due to internal server error: {e}") - @pytest.mark.parametrize( "model", ["openai/gpt-4.1", "anthropic/claude-sonnet-4-5-20250929"] ) @@ -742,7 +661,6 @@ def test_anthropic_mcp_server_responses_api(model: str): assert response is not None - def test_anthropic_prefix_prompt(): params = { "model": "anthropic/claude-sonnet-4-5-20250929", @@ -757,7 +675,6 @@ def test_anthropic_prefix_prompt(): assert response is not None assert response.choices[0].message.content.startswith("Argentina") - @pytest.mark.asyncio async def test_claude_tool_use_with_anthropic_acreate(): response = await litellm.anthropic.messages.acreate( @@ -782,9 +699,6 @@ async def test_claude_tool_use_with_anthropic_acreate(): async for chunk in response: print(chunk) - - - def test_anthropic_streaming(): request_data = { @@ -839,7 +753,6 @@ def test_anthropic_streaming(): assert role_set_count == 1 - def test_anthropic_via_responses_api(): from litellm.types.llms.openai import ResponsesAPIStreamEvents @@ -962,11 +875,6 @@ def test_anthropic_via_responses_api(): print(f"✓ All {len(events_seen)} events matched expected structure") print(f"✓ Received {text_delta_count} text delta chunks") - - - - - def _make_transform_request(optional_params: dict, litellm_params: dict) -> dict: from litellm.llms.anthropic.chat.transformation import AnthropicConfig @@ -978,19 +886,6 @@ def _make_transform_request(optional_params: dict, litellm_params: dict) -> dict headers={}, ) - - - - - - - - - - - - - def test_anthropic_basic_completion_replay(): response = litellm.completion( model="anthropic/claude-sonnet-4-5-20250929", @@ -1004,7 +899,6 @@ def test_anthropic_basic_completion_replay(): assert response.usage.completion_tokens > 0 assert response.choices[0].finish_reason in {"stop", "length"} - def test_anthropic_streaming_completion_replay(): stream = litellm.completion( model="anthropic/claude-sonnet-4-5-20250929", diff --git a/tests/llm_translation/test_azure_o_series.py b/tests/llm_translation/test_azure_o_series.py index 4d5ac43ab6b..ef774d5f7a3 100644 --- a/tests/llm_translation/test_azure_o_series.py +++ b/tests/llm_translation/test_azure_o_series.py @@ -1,15 +1,11 @@ import os - import pytest import litellm from base_llm_unit_tests import BaseLLMChatTest, BaseOSeriesModelsTest - class TestAzureOpenAIO3Mini(BaseOSeriesModelsTest, BaseLLMChatTest): - test_content_list_handling = None - test_empty_tools = None test_function_calling_with_tool_response = None def get_base_completion_call_args(self): @@ -34,8 +30,6 @@ class TestAzureOpenAIO3Mini(BaseOSeriesModelsTest, BaseLLMChatTest): def test_basic_tool_calling(self): pass - - class TestAzureOpenAIO3(BaseOSeriesModelsTest): def get_base_completion_call_args(self): return { diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index 6d8e8a431db..3c8e7ff416d 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -35,7 +35,6 @@ litellm.success_callback = [] user_message = "Write a short poem about the sky" messages = [{"content": user_message, "role": "user"}] - @pytest.fixture(autouse=True) def reset_callbacks(): print("\npytest fixture - resetting callbacks") @@ -44,7 +43,6 @@ def reset_callbacks(): litellm.failure_callback = [] litellm.callbacks = [] - def test_completion_bedrock_claude_completion_auth(monkeypatch): print("calling bedrock claude completion params auth") @@ -72,16 +70,13 @@ def test_completion_bedrock_claude_completion_auth(monkeypatch): except Exception as e: pytest.fail(f"Error occurred: {e}") - # test_completion_bedrock_claude_completion_auth() - @pytest.mark.parametrize("streaming", [True, False]) def test_completion_bedrock_guardrails(streaming): litellm.set_verbose = True - # verbose_logger.setLevel(logging.DEBUG) try: if streaming is False: @@ -146,10 +141,8 @@ def test_completion_bedrock_guardrails(streaming): except Exception as e: pytest.fail(f"Error occurred: {e}") - # test_completion_bedrock_claude_2_1_completion_auth() - def test_completion_bedrock_claude_external_client_auth(monkeypatch): print("\ncalling bedrock claude external client auth") @@ -187,21 +180,10 @@ def test_completion_bedrock_claude_external_client_auth(monkeypatch): except Exception as e: pytest.fail(f"Error occurred: {e}") - # test_completion_bedrock_claude_external_client_auth() - - - - - - - - - # test_completion_bedrock_claude_sts_client_auth() - @pytest.mark.parametrize( "stop", [""], @@ -240,7 +222,6 @@ def test_bedrock_stop_value(stop, model): except Exception as e: pytest.fail(f"Error occurred: {e}") - @pytest.mark.parametrize( "system", ["You are an AI", [{"type": "text", "text": "You are an AI"}], ""], @@ -279,89 +260,9 @@ def test_bedrock_system_prompt(system, model): pytest.fail(f"Error occurred: {e}") -def test_bedrock_claude_3_tool_calling(): - try: - litellm.set_verbose = True - litellm.turn_on_debug() - tools = [ - { - "type": "function", - "function": { - "name": "get_current_weather", - "description": "Get the current weather in a given location", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "The city and state, e.g. San Francisco, CA", - }, - "unit": { - "type": "string", - "enum": ["celsius", "fahrenheit"], - }, - }, - "required": ["location"], - }, - }, - } - ] - messages = [ - { - "role": "user", - "content": "What's the weather like in Boston today in fahrenheit?", - } - ] - response: ModelResponse = completion( - model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - messages=messages, - tools=tools, - tool_choice="auto", - ) # type: ignore - print(f"response: {response}") - # Add any assertions here to check the response - assert isinstance(response.choices[0].message.tool_calls[0].function.name, str) - assert isinstance( - response.choices[0].message.tool_calls[0].function.arguments, str - ) - messages.append( - response.choices[0].message.model_dump() - ) # Add assistant tool invokes - tool_result = ( - '{"location": "Boston", "temperature": "72", "unit": "fahrenheit"}' - ) - # Add user submitted tool results in the OpenAI format - messages.append( - { - "tool_call_id": response.choices[0].message.tool_calls[0].id, - "role": "tool", - "name": response.choices[0].message.tool_calls[0].function.name, - "content": tool_result, - } - ) - # In the second response, Claude should deduce answer from tool results - second_response = completion( - model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - messages=messages, - tools=tools, - tool_choice="auto", - ) - print(f"second response: {second_response}") - assert isinstance(second_response.choices[0].message.content, str) - except RateLimitError: - pass - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - - - - - def test_completion_bedrock_mistral_completion_auth(): print("calling bedrock mistral completion params auth") - litellm.turn_on_debug() # aws_access_key_id = os.environ["AWS_ACCESS_KEY_ID"] @@ -390,10 +291,8 @@ def test_completion_bedrock_mistral_completion_auth(): except Exception as e: pytest.fail(f"Error occurred: {e}") - # test_completion_bedrock_mistral_completion_auth() - def test_bedrock_ptu(): """ Check if a url with 'modelId' passed in, is created correctly @@ -425,7 +324,6 @@ def test_bedrock_ptu(): ) mock_client_post.assert_called_once() - @pytest.mark.asyncio async def test_bedrock_custom_api_base(): """ @@ -463,7 +361,6 @@ async def test_bedrock_custom_api_base(): ) mock_client_post.assert_called_once() - @pytest.mark.parametrize( "model", [ @@ -500,7 +397,6 @@ async def test_bedrock_extra_headers(model): ) mock_client_post.assert_called_once() - @pytest.mark.asyncio async def test_bedrock_custom_prompt_template(): """ @@ -545,7 +441,6 @@ async def test_bedrock_custom_prompt_template(): assert prompt == "<|im_start|>user\nWhat's AWS?<|im_end|>" mock_client_post.assert_called_once() - def test_completion_bedrock_external_client_region(monkeypatch): print("\ncalling bedrock claude external client auth") @@ -594,32 +489,15 @@ def test_completion_bedrock_external_client_region(monkeypatch): except Exception as e: pytest.fail(f"Error occurred: {e}") - - - - - - - - - - - - - from litellm.litellm_core_utils.prompt_templates.factory import ( _bedrock_converse_messages_pt, ) - - - def test_base_aws_llm_get_credentials(): import time import boto3 - start_time = time.time() session = boto3.Session( aws_access_key_id="test", @@ -650,15 +528,6 @@ def test_base_aws_llm_get_credentials(): ) ) - - - - - - - - - def test_bedrock_converse_route(): litellm.set_verbose = True try: @@ -672,7 +541,6 @@ def test_bedrock_converse_route(): else: raise - def test_bedrock_mapped_converse_models(): litellm.set_verbose = True os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" @@ -683,23 +551,8 @@ def test_bedrock_mapped_converse_models(): messages=[{"role": "user", "content": "Hello, world!"}], ) - - - - - - - - - class TestBedrockConverseChatCrossRegion(BaseLLMChatTest): - test_content_list_handling = None - test_developer_role_translation = None test_function_calling_with_tool_response = None - test_image_url = None - test_json_response_format_stream = None - test_tool_call_with_empty_enum_property = None - test_tool_call_with_property_type_array = None def get_base_completion_call_args(self) -> dict: os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" @@ -728,10 +581,7 @@ class TestBedrockConverseChatCrossRegion(BaseLLMChatTest): assert cost > 0 - class TestBedrockConverseAnthropicUnitTests(BaseAnthropicChatTest): - test_completion_thinking_with_max_tokens = None - test_completion_thinking_without_max_tokens = None def get_base_completion_call_args(self) -> dict: return { @@ -744,12 +594,8 @@ class TestBedrockConverseAnthropicUnitTests(BaseAnthropicChatTest): "thinking": {"type": "enabled", "budget_tokens": 16000}, } - class TestBedrockConverseChatNormal(BaseLLMChatTest): - test_content_list_handling = None - test_empty_tools = None test_function_calling_with_tool_response = None - test_image_url = None def get_base_completion_call_args(self) -> dict: os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" @@ -760,12 +606,8 @@ class TestBedrockConverseChatNormal(BaseLLMChatTest): "aws_region_name": "us-east-1", } - - class TestBedrockConverseNovaTestSuite(BaseLLMChatTest): - test_content_list_handling = None test_function_calling_with_tool_response = None - test_image_url = None def get_base_completion_call_args(self) -> dict: os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" @@ -776,9 +618,6 @@ class TestBedrockConverseNovaTestSuite(BaseLLMChatTest): "aws_region_name": "us-east-1", } - - - class TestBedrockRerank(BaseLLMRerankTest): def get_custom_llm_provider(self) -> litellm.LlmProviders: return litellm.LlmProviders.BEDROCK @@ -788,7 +627,6 @@ class TestBedrockRerank(BaseLLMRerankTest): "model": "bedrock/arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0", } - class TestBedrockCohereRerank(BaseLLMRerankTest): def get_custom_llm_provider(self) -> litellm.LlmProviders: return litellm.LlmProviders.BEDROCK @@ -798,13 +636,6 @@ class TestBedrockCohereRerank(BaseLLMRerankTest): "model": "bedrock/arn:aws:bedrock:us-west-2::foundation-model/cohere.rerank-v3-5:0", } - - - - - - - @pytest.mark.parametrize("top_k_param", ["top_k", "topK"]) def test_bedrock_nova_topk(top_k_param): litellm.set_verbose = True @@ -830,7 +661,6 @@ def test_bedrock_nova_topk(top_k_param): assert "inferenceConfig" in captured_data["additionalModelRequestFields"] assert captured_data["additionalModelRequestFields"]["inferenceConfig"]["topK"] == 10 - def test_bedrock_cross_region_inference(monkeypatch): from litellm.llms.custom_httpx.http_handler import HTTPHandler @@ -856,7 +686,6 @@ def test_bedrock_cross_region_inference(monkeypatch): == "https://bedrock-runtime.us-west-2.amazonaws.com/model/us.meta.llama3-3-70b-instruct-v1%3A0/converse" ) - def test_bedrock_empty_content_real_call(): completion( model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", @@ -874,11 +703,6 @@ def test_bedrock_empty_content_real_call(): ], ) - - - - - class TestBedrockEmbedding(BaseLLMEmbeddingTest): def get_base_embedding_call_args(self) -> dict: return { @@ -888,8 +712,6 @@ class TestBedrockEmbedding(BaseLLMEmbeddingTest): def get_custom_llm_provider(self) -> litellm.LlmProviders: return litellm.LlmProviders.BEDROCK - - @pytest.mark.asyncio async def test_bedrock_image_url_sync_client(): import logging @@ -928,9 +750,6 @@ async def test_bedrock_image_url_sync_client(): print(e) mock_post.assert_called_once() - - - def test_bedrock_custom_proxy(): from litellm.llms.custom_httpx.http_handler import HTTPHandler @@ -953,7 +772,6 @@ def test_bedrock_custom_proxy(): assert mock_post.call_args.kwargs["headers"]["Authorization"] == "Bearer Token" - def test_bedrock_custom_deepseek(): import json @@ -1003,13 +821,6 @@ def test_bedrock_custom_deepseek(): print(f"Error: {str(e)}") raise e - - - - - - - def test_bedrock_description_param(): from litellm import completion from litellm.llms.custom_httpx.http_handler import HTTPHandler @@ -1048,7 +859,6 @@ def test_bedrock_description_param(): "Find the meaning inside a poem" in request_body_str ) # assert description is passed - @pytest.mark.parametrize( "sync_mode", [ @@ -1108,7 +918,6 @@ async def test_bedrock_thinking_in_assistant_message(sync_mode): in json_data ) - @pytest.mark.asyncio async def test_bedrock_stream_thinking_content_openwebui(): """ @@ -1177,7 +986,6 @@ async def test_bedrock_stream_thinking_content_openwebui(): len(response_content) > 0 ), "There should be non-empty content after thinking tags" - def test_bedrock_application_inference_profile(): from litellm.llms.custom_httpx.http_handler import HTTPHandler @@ -1242,7 +1050,6 @@ def test_bedrock_application_inference_profile(): ) assert mock_post2.call_args.kwargs["url"] == mock_post.call_args.kwargs["url"] - def return_mocked_response(model: str): if model == "bedrock/mistral.mistral-large-2407-v1:0": return { @@ -1257,7 +1064,6 @@ def return_mocked_response(model: str): "usage": {"inputTokens": 5, "outputTokens": 10, "totalTokens": 15}, } - @pytest.mark.parametrize( "model", [ @@ -1301,9 +1107,6 @@ async def test_bedrock_max_completion_tokens(model: str): "inferenceConfig": {"maxTokens": 10}, } - - - @pytest.mark.asyncio async def test_bedrock_passthrough_router(): """ @@ -1357,7 +1160,6 @@ async def test_bedrock_passthrough_router(): assert response.status_code == 200 - @pytest.mark.asyncio async def test_bedrock_converse__streaming_passthrough(monkeypatch): import asyncio @@ -1411,7 +1213,6 @@ async def test_bedrock_converse__streaming_passthrough(monkeypatch): assert response_cost is not None and response_cost > 0 assert "standard_logging_object" in mock_callback.call_args.kwargs["kwargs"] - @pytest.mark.asyncio async def test_bedrock_streaming_passthrough_test2(monkeypatch): import asyncio @@ -1462,7 +1263,6 @@ async def test_bedrock_streaming_passthrough_test2(monkeypatch): assert "standard_logging_object" in mock_callback.call_args.kwargs["kwargs"] assert "response_cost" in mock_callback.call_args.kwargs["kwargs"] - def test_bedrock_openai_imported_model(): """ Test that Bedrock imported models using OpenAI format work correctly. @@ -1564,25 +1364,6 @@ def test_bedrock_openai_imported_model(): assert request_body["max_tokens"] == 300 assert request_body["temperature"] == 0.5 - - - - - - - - - - - - - - - - - - - def test_bedrock_openai_multiple_message_types(): """ Test that various message content types are handled correctly. @@ -1631,14 +1412,10 @@ def test_bedrock_openai_multiple_message_types(): print("✓ Multiple message types handled correctly") - - - # ============================================================================ # Nova Grounding (web_search_options) Unit Tests (Mocked) # ============================================================================ - def test_bedrock_nova_grounding_web_search_options_non_streaming(): """ Unit test for Nova grounding using web_search_options parameter (non-streaming). @@ -1697,7 +1474,6 @@ def test_bedrock_nova_grounding_web_search_options_non_streaming(): f"✓ web_search_options correctly transformed to systemTool (non-streaming)" ) - def test_bedrock_nova_grounding_with_function_tools(): """ Unit test for Nova grounding combined with regular function tools. @@ -1782,7 +1558,6 @@ def test_bedrock_nova_grounding_with_function_tools(): assert system_tool_found, "systemTool (nova_grounding) should be present" print(f"✓ Both function tools and web_search_options correctly combined") - @pytest.mark.asyncio async def test_bedrock_nova_grounding_async(): """ @@ -1836,9 +1611,6 @@ async def test_bedrock_nova_grounding_async(): assert system_tool_found, "systemTool with nova_grounding should be present" print(f"✓ Async web_search_options correctly transformed to systemTool") - - - def test_bedrock_nova_grounding_request_transformation(): """ Unit test to verify that web_search_options transforms to systemTool in the request. diff --git a/tests/llm_translation/test_bedrock_gpt_oss.py b/tests/llm_translation/test_bedrock_gpt_oss.py index e0d7b1b904f..b5252847805 100644 --- a/tests/llm_translation/test_bedrock_gpt_oss.py +++ b/tests/llm_translation/test_bedrock_gpt_oss.py @@ -1,8 +1,6 @@ from base_llm_unit_tests import BaseLLMChatTest - class TestBedrockGPTOSS(BaseLLMChatTest): - test_json_response_format = None def get_base_completion_call_args(self) -> dict: return { @@ -19,8 +17,3 @@ class TestBedrockGPTOSS(BaseLLMChatTest): """ pass - async def test_completion_cost(self): - """ - Bedrock GPT-OSS models are flaky and occasionally report 0 token counts in api response - """ - pass diff --git a/tests/llm_translation/test_bedrock_invoke_tests.py b/tests/llm_translation/test_bedrock_invoke_tests.py index 584b0ef341f..8f61a359f6a 100644 --- a/tests/llm_translation/test_bedrock_invoke_tests.py +++ b/tests/llm_translation/test_bedrock_invoke_tests.py @@ -5,7 +5,6 @@ import os import litellm from litellm.types.llms.bedrock import BedrockInvokeNovaRequest - _LITELLM_LOGO_IMAGE_URL = ( "https://cdn.jsdelivr.net/gh/BerriAI/litellm@d769e81c90d453240c61fc572cdb27fae06a89d0/" "ui/litellm-dashboard/public/assets/logos/litellm_logo.jpg" @@ -15,7 +14,6 @@ _AWSMP_LOGO_IMAGE_URL = ( "c233c9ade2ccb5491072ae232c814942.png" ) - @pytest.mark.flaky(retries=3, delay=5) class TestBedrockInvokeClaudeJson(BaseLLMChatTest): def get_base_completion_call_args(self) -> dict: @@ -24,26 +22,9 @@ class TestBedrockInvokeClaudeJson(BaseLLMChatTest): "model": "bedrock/invoke/us.anthropic.claude-haiku-4-5-20251001-v1:0", } - @pytest.mark.parametrize( - "image_url, detail", - [ - (_LITELLM_LOGO_IMAGE_URL, None), - (_LITELLM_LOGO_IMAGE_URL, "low"), - (_LITELLM_LOGO_IMAGE_URL, "high"), - (_AWSMP_LOGO_IMAGE_URL, "low"), - (_AWSMP_LOGO_IMAGE_URL, "high"), - ], - ) - @pytest.mark.flaky(retries=4, delay=2) - def test_image_url(self, image_url, detail): - super().test_image_url(detail=detail, image_url=image_url) - test_content_list_handling = None - test_image_url_string = None test_pdf_handling = None - class TestBedrockInvokeNovaJson(BaseLLMChatTest): - test_json_response_format = None def get_base_completion_call_args(self) -> dict: return { @@ -57,11 +38,3 @@ class TestBedrockInvokeNovaJson(BaseLLMChatTest): f"Skipping non-JSON test: {request.function.__name__} does not contain 'json'" ) - def test_json_response_pydantic_obj(self): - if os.environ.get("LITELLM_RUN_LIVE_BEDROCK_NOVA_JSON_TESTS") != "1": - pytest.skip("Live Bedrock Nova response-schema E2E tests are opt-in") - if os.environ.get("CASSETTE_REDIS_URL"): - pytest.skip( - "Live Bedrock Nova response-schema E2E tests cannot run under VCR replay" - ) - super().test_json_response_pydantic_obj() diff --git a/tests/llm_translation/test_bedrock_llama.py b/tests/llm_translation/test_bedrock_llama.py index 9c83c28bbe0..3a3f0064c52 100644 --- a/tests/llm_translation/test_bedrock_llama.py +++ b/tests/llm_translation/test_bedrock_llama.py @@ -3,10 +3,7 @@ import pytest import litellm - class TestBedrockTestSuite(BaseLLMChatTest): - test_content_list_handling = None - test_empty_tools = None test_function_calling_with_tool_response = None def get_base_completion_call_args(self) -> dict: diff --git a/tests/llm_translation/test_bedrock_moonshot.py b/tests/llm_translation/test_bedrock_moonshot.py index cf5e554e91b..b2026b9ed1f 100644 --- a/tests/llm_translation/test_bedrock_moonshot.py +++ b/tests/llm_translation/test_bedrock_moonshot.py @@ -15,20 +15,12 @@ from base_llm_unit_tests import BaseLLMChatTest import litellm - class TestBedrockMoonshotInvoke(BaseLLMChatTest): """ Test suite for Bedrock Moonshot via invoke route. Inherits all standard LLM tests from BaseLLMChatTest. """ - test_json_response_format_stream = None - test_completion_cost = None - test_content_list_handling = None - test_developer_role_translation = None - test_message_with_name = None - test_pydantic_model_input = None - test_response_format_type_text_with_tool_calls_no_tool_choice = None test_streaming = None def get_base_completion_call_args(self) -> dict: diff --git a/tests/llm_translation/test_bedrock_nova_json.py b/tests/llm_translation/test_bedrock_nova_json.py index 8adfef50618..7c1023e0671 100644 --- a/tests/llm_translation/test_bedrock_nova_json.py +++ b/tests/llm_translation/test_bedrock_nova_json.py @@ -3,15 +3,8 @@ import pytest import litellm - class TestBedrockNovaJson(BaseLLMChatTest): - test_content_list_handling = None - test_developer_role_translation = None - test_empty_tools = None test_function_calling_with_tool_response = None - test_json_response_format_stream = None - test_tool_call_with_empty_enum_property = None - test_tool_call_with_property_type_array = None def get_base_completion_call_args(self) -> dict: litellm.turn_on_debug() @@ -19,14 +12,6 @@ class TestBedrockNovaJson(BaseLLMChatTest): "model": "bedrock/converse/us.amazon.nova-micro-v1:0", } - def test_json_response_nested_pydantic_obj(self): - pass - - def test_json_response_nested_json_schema(self): - pass - - - # @pytest.fixture(autouse=True) # def skip_non_json_tests(self, request): # if not "json" in request.function.__name__.lower(): diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index d05afe92626..acbad4fb700 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -2,7 +2,6 @@ import os import pytest - from base_llm_unit_tests import BaseLLMChatTest from litellm.llms.vertex_ai.context_caching.transformation import ( separate_cached_messages, @@ -12,18 +11,9 @@ import litellm from litellm import completion import json - - - class TestGoogleAIStudioGemini(BaseLLMChatTest): test_async_pdf_handling_with_file_id = None - test_content_list_handling = None - test_developer_role_translation = None test_function_calling_with_tool_response = None - test_image_url = None - test_json_response_nested_json_schema = None - test_json_response_nested_pydantic_obj = None - test_json_response_pydantic_obj = None test_web_search = None def get_base_completion_call_args(self) -> dict: @@ -32,7 +22,6 @@ class TestGoogleAIStudioGemini(BaseLLMChatTest): def get_base_completion_call_args_with_reasoning_model(self) -> dict: return {"model": "gemini/gemini-2.5-flash"} - @pytest.mark.flaky(retries=3, delay=2) def test_url_context(self): from litellm.utils import supports_url_context @@ -64,11 +53,6 @@ class TestGoogleAIStudioGemini(BaseLLMChatTest): ), "URL context metadata should be present" print(f"response={response}") - - - - - def test_gemini_image_generation(): # litellm.turn_on_debug() response = completion( @@ -90,23 +74,6 @@ def test_gemini_image_generation(): .startswith("data:image/png;base64,") ) - - - - - - - - - - - - - - - - - def test_gemini_thinking(): litellm.turn_on_debug() from litellm.types.utils import Message, CallTypes @@ -146,9 +113,6 @@ def test_gemini_thinking(): print(response.choices[0].message) assert response.choices[0].message.content is not None - - - def test_gemini_finish_reason(): import os from litellm import completion @@ -163,7 +127,6 @@ def test_gemini_finish_reason(): assert response.choices[0].finish_reason is not None assert response.choices[0].finish_reason == "length" - @pytest.mark.flaky(retries=3, delay=2) def test_gemini_url_context(): from litellm import completion @@ -189,7 +152,6 @@ def test_gemini_url_context(): assert urlMetadata["retrievedUrl"] == URL1 assert urlMetadata["urlRetrievalStatus"] == "URL_RETRIEVAL_STATUS_SUCCESS" - @pytest.mark.flaky(retries=3, delay=2) def test_gemini_with_grounding(): from litellm import completion, Usage, stream_chunk_builder @@ -226,7 +188,6 @@ def test_gemini_with_grounding(): assert usage.prompt_tokens_details.web_search_requests is not None assert usage.prompt_tokens_details.web_search_requests > 0 - def test_gemini_with_empty_function_call_arguments(): from litellm import completion @@ -248,9 +209,6 @@ def test_gemini_with_empty_function_call_arguments(): print(response) assert response.choices[0].message.content is not None - - - def test_gemini_tool_use(): data = { "max_tokens": 8192, @@ -299,7 +257,6 @@ def test_gemini_tool_use(): assert stop_reason is not None assert stop_reason == "tool_calls" - @pytest.mark.asyncio async def test_gemini_image_generation_async(): litellm.turn_on_debug() @@ -332,7 +289,6 @@ async def test_gemini_image_generation_async(): assert IMAGE_URL["url"] is not None, "IMAGE_URL['url'] is not None" assert IMAGE_URL["url"].startswith("data:image/png;base64,") - @pytest.mark.asyncio async def test_gemini_image_generation_async_stream(): # litellm.turn_on_debug() @@ -367,7 +323,6 @@ async def test_gemini_image_generation_async_stream(): assert model_response_image is not None assert model_response_image["url"].startswith("data:image/png;base64,") - def test_system_message_with_no_user_message(): """ Test that the system message is translated correctly for non-OpenAI providers. @@ -387,7 +342,6 @@ def test_system_message_with_no_user_message(): assert response.choices[0].message.content is not None - def get_current_weather(location, unit="fahrenheit"): """Get the current weather in a given location""" if "tokyo" in location.lower(): @@ -401,7 +355,6 @@ def get_current_weather(location, unit="fahrenheit"): else: return json.dumps({"location": location, "temperature": "unknown"}) - def test_gemini_with_thinking(): from litellm import completion @@ -493,11 +446,6 @@ def test_gemini_with_thinking(): ) # get a new response from the model where it can see the function response print("second response\n", second_response) - - - - - @pytest.mark.parametrize( "status_code,expected_exception", [ @@ -582,7 +530,6 @@ def l(status_code, expected_exception): "VertexAIException" not in error_message ), f"Should not contain 'VertexAIException' for status {status_code}, got: {error_message}" - def test_gemini_embedding(): litellm.turn_on_debug() response = litellm.embedding( @@ -592,23 +539,6 @@ def test_gemini_embedding(): print("response: ", response) assert response is not None - - - - - - - - - - - - - - - - - @pytest.mark.asyncio async def test_gemini_openai_web_search_tool_to_google_search(): """ diff --git a/tests/llm_translation/test_groq.py b/tests/llm_translation/test_groq.py index 55b1b9bc7eb..ea633b5cc84 100644 --- a/tests/llm_translation/test_groq.py +++ b/tests/llm_translation/test_groq.py @@ -1,6 +1,5 @@ - # sys.path.insert( # 0, os.path.abspath("../..") # ) # noqa @@ -8,10 +7,7 @@ from base_llm_unit_tests import BaseLLMChatTest - class TestGroq(BaseLLMChatTest): - test_content_list_handling = None - test_empty_tools = None test_web_search = None def get_base_completion_call_args(self) -> dict: @@ -19,5 +15,3 @@ class TestGroq(BaseLLMChatTest): "model": "groq/openai/gpt-oss-120b", } - def test_tool_call_with_empty_enum_property(self): - pass diff --git a/tests/llm_translation/test_huggingface_chat_completion.py b/tests/llm_translation/test_huggingface_chat_completion.py index 57c2d4df71c..6e076e6de63 100644 --- a/tests/llm_translation/test_huggingface_chat_completion.py +++ b/tests/llm_translation/test_huggingface_chat_completion.py @@ -7,7 +7,6 @@ from unittest.mock import AsyncMock, MagicMock, patch from base_llm_unit_tests import BaseLLMChatTest - import pytest import litellm @@ -126,7 +125,6 @@ MOCK_STREAMING_CHUNKS = [ }, ] - PROVIDER_MAPPING_RESPONSE = { "fireworks-ai": { "status": "live", @@ -145,14 +143,12 @@ PROVIDER_MAPPING_RESPONSE = { }, } - @pytest.fixture def mock_provider_mapping(): with patch("litellm.llms.huggingface.chat.transformation.fetch_inference_provider_mapping") as mock: mock.return_value = PROVIDER_MAPPING_RESPONSE yield mock - @pytest.fixture(autouse=True) def clear_lru_cache(): from litellm.llms.huggingface.common_utils import fetch_inference_provider_mapping @@ -161,7 +157,6 @@ def clear_lru_cache(): yield fetch_inference_provider_mapping.cache_clear() - @pytest.fixture def mock_http_handler(): """Fixture to mock the HTTP handler""" @@ -188,7 +183,6 @@ def mock_http_handler(): mock.side_effect = mock_side_effect yield mock - @pytest.fixture def mock_http_async_handler(): """Fixture to mock the async HTTP handler""" @@ -220,7 +214,6 @@ def mock_http_async_handler(): mock.side_effect = mock_side_effect yield mock - class TestHuggingFace(BaseLLMChatTest): @pytest.fixture(autouse=True) def setup(self, mock_provider_mapping, mock_http_handler, mock_http_async_handler): @@ -355,8 +348,6 @@ class TestHuggingFace(BaseLLMChatTest): == tool_call_no_arguments["tool_calls"][0]["function"]["arguments"] ) - - def test_completion_with_api_base(self): messages = [{"role": "user", "content": "This is a test message"}] api_base = "https://abcd123.us-east-1.aws.endpoints.huggingface.cloud" @@ -421,9 +412,3 @@ class TestHuggingFace(BaseLLMChatTest): called_url = call_args[1]["url"] assert called_url == f"{api_base}/v1/chat/completions" - - - - @pytest.mark.asyncio - async def test_completion_cost(self): - pass diff --git a/tests/llm_translation/test_mistral_audio_transcription_transformation.py b/tests/llm_translation/test_mistral_audio_transcription_transformation.py index db77eabba23..6328e52cbed 100644 --- a/tests/llm_translation/test_mistral_audio_transcription_transformation.py +++ b/tests/llm_translation/test_mistral_audio_transcription_transformation.py @@ -8,7 +8,6 @@ from tests.llm_translation.base_audio_transcription_unit_tests import ( BaseLLMAudioTranscriptionTest, ) - @pytest.mark.skipif( not os.getenv("MISTRAL_API_KEY"), reason="MISTRAL_API_KEY not set, skipping Mistral audio transcription tests", @@ -22,8 +21,3 @@ class TestMistralAudioTranscription(BaseLLMAudioTranscriptionTest): def get_custom_llm_provider(self) -> litellm.LlmProviders: return litellm.LlmProviders.MISTRAL - def test_audio_transcription_async(self): # type: ignore[override] - pytest.skip( - "Async audio transcription test for Mistral is skipped in this suite; " - "async test plugins (e.g. pytest-asyncio/anyio) are not configured here." - ) diff --git a/tests/llm_translation/test_nvidia_nim.py b/tests/llm_translation/test_nvidia_nim.py index 2de1d0f3b84..402035444b1 100644 --- a/tests/llm_translation/test_nvidia_nim.py +++ b/tests/llm_translation/test_nvidia_nim.py @@ -3,7 +3,6 @@ from datetime import datetime from typing import Final from unittest.mock import AsyncMock - import httpx import pytest from openai.types import CreateEmbeddingResponse, Embedding @@ -16,7 +15,6 @@ from litellm import completion from base_rerank_unit_tests import BaseLLMRerankTest from tests.capturing_transport import CapturingTransport - def test_completion_nvidia_nim(): from openai import OpenAI @@ -59,7 +57,6 @@ def test_completion_nvidia_nim(): assert request_body["frequency_penalty"] == 0.1 assert request_body["presence_penalty"] == 0.5 - class TestNvidiaNim(BaseLLMRerankTest): def get_custom_llm_provider(self) -> litellm.LlmProviders: return litellm.LlmProviders.NVIDIA_NIM @@ -73,43 +70,3 @@ class TestNvidiaNim(BaseLLMRerankTest): """Nvidia NIM rerank models are free (cost = 0.0)""" return 0.0 - @pytest.mark.asyncio() - @pytest.mark.parametrize("sync_mode", [True, False]) - async def test_basic_rerank(self, sync_mode, monkeypatch): - """ - Override the base live rerank test with a mocked HTTP layer. - - NVIDIA reached end-of-life for the hosted - nvidia/llama-3.2-nv-rerankqa-1b-v2 rerank API on 2026-05-18 and - published no replacement model, so a live call now returns HTTP 410 - ("Gone"). NVIDIA's hosted catalog rotates on a schedule, so pointing - at another live model would only defer the same failure. Mock the - transport instead (same pattern as - test_nvidia_nim_rerank_ranking_endpoint above) so the request/response - transformation and cost calculation stay covered offline. - """ - monkeypatch.setenv("NVIDIA_NIM_API_KEY", "fake-api-key") - - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.headers = {} - mock_response.text = "" - mock_response.json.return_value = { - "rankings": [ - {"index": 0, "logit": 0.95}, - {"index": 1, "logit": 0.75}, - ], - "usage": {"total_tokens": 7}, - } - - with ( - patch( - "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", - return_value=mock_response, - ), - patch( - "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", - return_value=mock_response, - ), - ): - await super().test_basic_rerank(sync_mode=sync_mode) diff --git a/tests/llm_translation/test_openai_o1.py b/tests/llm_translation/test_openai_o1.py index 61b17dfecf1..6b597b016fc 100644 --- a/tests/llm_translation/test_openai_o1.py +++ b/tests/llm_translation/test_openai_o1.py @@ -1,18 +1,13 @@ import os from unittest.mock import patch - import pytest import litellm from litellm import ModelResponse from base_llm_unit_tests import BaseLLMChatTest, BaseOSeriesModelsTest - class TestOpenAIO1(BaseOSeriesModelsTest, BaseLLMChatTest): - test_empty_tools = None - test_tool_call_with_empty_enum_property = None - test_tool_call_with_property_type_array = None def get_base_completion_call_args(self): return { @@ -24,9 +19,6 @@ class TestOpenAIO1(BaseOSeriesModelsTest, BaseLLMChatTest): return OpenAI(api_key="fake-api-key") - - - class TestOpenAIO3(BaseOSeriesModelsTest, BaseLLMChatTest): test_basic_tool_calling = None test_function_calling_with_tool_response = None @@ -41,9 +33,6 @@ class TestOpenAIO3(BaseOSeriesModelsTest, BaseLLMChatTest): return OpenAI(api_key="fake-api-key") - - - def test_o3_reasoning_effort(): resp = litellm.completion( model="o3-mini", diff --git a/tests/llm_translation/test_ovhcloud_audio_transcription_transformation.py b/tests/llm_translation/test_ovhcloud_audio_transcription_transformation.py index 8cc46dc98d0..a1e03ae5e6f 100644 --- a/tests/llm_translation/test_ovhcloud_audio_transcription_transformation.py +++ b/tests/llm_translation/test_ovhcloud_audio_transcription_transformation.py @@ -12,7 +12,6 @@ from tests.llm_translation.base_audio_transcription_unit_tests import ( BaseLLMAudioTranscriptionTest, ) - @pytest.mark.skipif( not os.getenv("OVHCLOUD_API_KEY"), reason="OVHCLOUD_API_KEY not set, skipping OVHCloud audio transcription tests", @@ -29,12 +28,6 @@ class TestOVHCloudAudioTranscription(BaseLLMAudioTranscriptionTest): # Override the async base test with a sync no-op to avoid # 'async def functions are not natively supported' failures when # running this file in isolation without pytest-asyncio. - def test_audio_transcription_async(self): # type: ignore[override] - pytest.skip( - "Async audio transcription test for OVHCloud is skipped in this suite; " - "async test plugins (e.g. pytest-asyncio/anyio) are not configured here." - ) - @pytest.mark.skipif( not os.getenv("OVHCLOUD_API_KEY"), diff --git a/tests/llm_translation/test_together_ai.py b/tests/llm_translation/test_together_ai.py index a203c9edcfe..b78fbeca9e8 100644 --- a/tests/llm_translation/test_together_ai.py +++ b/tests/llm_translation/test_together_ai.py @@ -8,21 +8,12 @@ import json from datetime import datetime from unittest.mock import AsyncMock - import litellm import pytest - class TestTogetherAI(BaseLLMChatTest): test_basic_tool_calling = None - test_empty_tools = None test_function_calling_with_tool_response = None - test_json_response_format = None - test_json_response_nested_json_schema = None - test_json_response_nested_pydantic_obj = None - test_json_response_pydantic_obj = None - test_tool_call_with_empty_enum_property = None - test_tool_call_with_property_type_array = None def get_base_completion_call_args(self) -> dict: litellm.set_verbose = True diff --git a/tests/local_testing/test_caching.py b/tests/local_testing/test_caching.py index b8c7a4c189b..f4c5382b8a2 100644 --- a/tests/local_testing/test_caching.py +++ b/tests/local_testing/test_caching.py @@ -2277,56 +2277,3 @@ def test_caching_thinking_args_hit(): # test in memory cache except Exception as e: print(f"error occurred: {traceback.format_exc()}") pytest.fail(f"Error occurred: {e}") - - -@pytest.mark.asyncio -async def test_cache_key_in_hidden_params_acompletion(): - """ - Test that cache_key is present in _hidden_params on cache hits for acompletion. - - Validates fix for missing x-litellm-cache-key header on proxy cache hits. - """ - litellm.cache = Cache( - type="redis", - host=os.environ["REDIS_HOST"], - port=os.environ["REDIS_PORT"], - password=os.environ["REDIS_PASSWORD"], - ) - - unique_content = f"test cache key hidden params {uuid.uuid4()}" - messages = [{"role": "user", "content": unique_content}] - - # First call - cache miss - response1 = await litellm.acompletion( - model="gpt-3.5-turbo", - messages=messages, - mock_response="test response", - caching=True, - ) - - print(f"Response 1 _hidden_params: {response1._hidden_params}") - assert response1._hidden_params.get("cache_hit") is not True - - await asyncio.sleep(0.5) - - # Second call - cache hit - response2 = await litellm.acompletion( - model="gpt-3.5-turbo", - messages=messages, - mock_response="test response", - caching=True, - ) - - print(f"Response 2 _hidden_params: {response2._hidden_params}") - - # Verify cache hit occurred - assert response2._hidden_params.get("cache_hit") is True - - # Verify cache_key is present in _hidden_params - assert "cache_key" in response2._hidden_params - assert response2._hidden_params["cache_key"] is not None - - # Verify both responses have same ID (cache hit) - assert response1.id == response2.id - - litellm.cache = None diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py index 78e185c3b37..b5703224eb1 100644 --- a/tests/local_testing/test_completion.py +++ b/tests/local_testing/test_completion.py @@ -1413,28 +1413,6 @@ def test_replicate_custom_prompt_dict(): # test_completion_together_ai_mixtral() -def test_completion_together_ai_llama(): - litellm.set_verbose = True - model_name = "together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo" - try: - messages = [ - {"role": "user", "content": "What llm are you?"}, - ] - response = completion(model=model_name, messages=messages, max_tokens=5) - # Add any assertions here to check the response - print(response) - cost = completion_cost(completion_response=response) - assert cost > 0.0 - print( - "Cost for completion call together-computer/llama-2-70b: ", - f"${float(cost):.10f}", - ) - except litellm.Timeout as e: - pass - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - # test_completion_together_ai_yi_chat() @@ -1903,25 +1881,6 @@ def test_langfuse_completion(monkeypatch): ) - - - - -def test_deepseek_reasoning_content_completion(): - try: - litellm.set_verbose = True - litellm.turn_on_debug() - resp = litellm.completion( - timeout=5, - model="deepseek/deepseek-reasoner", - messages=[{"role": "user", "content": "Tell me a joke."}], - ) - - assert resp.choices[0].message.reasoning_content is not None - except litellm.Timeout: - pytest.skip("Model is timing out") - - def test_qwen_text_completion(): # litellm.turn_on_debug() resp = litellm.completion( diff --git a/tests/local_testing/test_custom_logger.py b/tests/local_testing/test_custom_logger.py index 192a73a68ab..b721b45a8bf 100644 --- a/tests/local_testing/test_custom_logger.py +++ b/tests/local_testing/test_custom_logger.py @@ -365,54 +365,6 @@ async def test_async_custom_handler_embedding_optional_param(): # asyncio.run(test_async_custom_handler_embedding_optional_param()) - - -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=1) -async def test_cost_tracking_with_caching(): - """ - Important Test - This tests if that cost is 0 for cached responses - """ - from litellm import Cache - - litellm.set_verbose = True - litellm.cache = Cache( - type="redis", - host=os.environ["REDIS_HOST"], - port=os.environ["REDIS_PORT"], - password=os.environ["REDIS_PASSWORD"], - ) - customHandler_optional_params = MyCustomHandler() - litellm.callbacks = [customHandler_optional_params] - messages = [ - { - "role": "user", - "content": f"write a one sentence poem about: {time.time()}", - } - ] - response1 = await litellm.acompletion( - model="gpt-3.5-turbo", - messages=messages, - max_tokens=40, - temperature=0.2, - caching=True, - mock_response="Hey, i'm doing well!", - ) - await asyncio.sleep(3) # success callback is async - response_cost = customHandler_optional_params.response_cost - assert response_cost > 0 - response2 = await litellm.acompletion( - model="gpt-3.5-turbo", - messages=messages, - max_tokens=40, - temperature=0.2, - caching=True, - ) - await asyncio.sleep(1) # success callback is async - response_cost_2 = customHandler_optional_params.response_cost - assert response_cost_2 == 0 - - @pytest.mark.flaky(retries=3, delay=3) def test_redis_cache_completion_stream(): # Important Test - This tests if we can add to streaming cache, when custom callbacks are set diff --git a/tests/local_testing/test_lunary.py b/tests/local_testing/test_lunary.py deleted file mode 100644 index 4312abcc0a4..00000000000 --- a/tests/local_testing/test_lunary.py +++ /dev/null @@ -1,57 +0,0 @@ -import io - - -import litellm - -litellm.failure_callback = ["lunary"] -litellm.success_callback = ["lunary"] -litellm.set_verbose = True - - - - - - - - -def test_lunary_with_tools(): - import litellm - - messages = [ - { - "role": "user", - "content": "What's the weather like in San Francisco, Tokyo, and Paris?", - } - ] - tools = [ - { - "type": "function", - "function": { - "name": "get_current_weather", - "description": "Get the current weather in a given location", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "The city and state, e.g. San Francisco, CA", - }, - "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]}, - }, - "required": ["location"], - }, - }, - } - ] - - response = litellm.completion( - model="gpt-6-luna", - messages=messages, - tools=tools, - tool_choice="auto", # auto is default, but we'll be explicit - ) - - response_message = response.choices[0].message - assert response.choices[0].message.tool_calls - assert all(call.function.name == "get_current_weather" for call in response.choices[0].message.tool_calls) - print("\nLLM Response:\n", response.choices[0].message) diff --git a/tests/local_testing/test_text_completion.py b/tests/local_testing/test_text_completion.py index d62b141d636..a117ca975aa 100644 --- a/tests/local_testing/test_text_completion.py +++ b/tests/local_testing/test_text_completion.py @@ -2729,19 +2729,6 @@ def test_completion_openai_engine() -> None: # test_completion_openai_engine() -def test_completion_chatgpt_prompt(): - try: - print("\n gpt3.5 test\n") - response = text_completion(model="openai/gpt-3.5-turbo", prompt="What's the weather in SF?") - print(response) - response_str = response["choices"][0]["text"] - print("\n", response.choices) - print("\n", response.choices[0]) - # print(response.choices[0].text) - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - # test_completion_chatgpt_prompt() diff --git a/tests/logging_callback_tests/test_alerting.py b/tests/logging_callback_tests/test_alerting.py deleted file mode 100644 index 691bb58b998..00000000000 --- a/tests/logging_callback_tests/test_alerting.py +++ /dev/null @@ -1,338 +0,0 @@ -# What is this? -## Tests slack alerting on proxy logging object - -import asyncio - -# import logging -# logging.basicConfig(level=logging.DEBUG) -from datetime import datetime -from unittest.mock import AsyncMock, patch - -import pytest - -import litellm -from litellm.caching.caching import DualCache -from litellm.integrations.SlackAlerting.slack_alerting import ( - SlackAlerting, -) -from litellm.proxy._types import CallInfo, Litellm_EntityType -from litellm.proxy.utils import ProxyLogging -from litellm.types.integrations.slack_alerting import AlertType - - -@pytest.mark.asyncio -async def test_get_api_base(): - _pl = ProxyLogging(user_api_key_cache=DualCache()) - _pl.update_values(alerting=["slack"], alerting_threshold=100, redis_cache=None) - model = "chatgpt-v-3" - messages = [{"role": "user", "content": "Hey how's it going?"}] - litellm_params = { - "acompletion": True, - "api_key": None, - "api_base": "https://openai-gpt-4-test-v-1.openai.azure.com/", - "force_timeout": 600, - "logger_fn": None, - "verbose": False, - "custom_llm_provider": "azure", - "litellm_call_id": "68f46d2d-714d-4ad8-8137-69600ec8755c", - "model_alias_map": {}, - "completion_call_id": None, - "metadata": None, - "model_info": None, - "proxy_server_request": None, - "preset_cache_key": None, - "no-log": False, - "stream_response": {}, - } - start_time = datetime.now() - end_time = datetime.now() - - time_difference_float, model, api_base, messages = ( - _pl.slack_alerting_instance._response_taking_too_long_callback_helper( - kwargs={ - "model": model, - "messages": messages, - "litellm_params": litellm_params, - }, - start_time=start_time, - end_time=end_time, - ) - ) - - assert api_base is not None - assert isinstance(api_base, str) - assert len(api_base) > 0 - request_info = ( - f"\nRequest Model: `{model}`\nAPI Base: `{api_base}`\nMessages: `{messages}`" - ) - slow_message = f"`Responses are slow - {round(time_difference_float,2)}s response time > Alerting threshold: {100}s`" - await _pl.alerting_handler( - message=slow_message + request_info, - level="Low", - alert_type=AlertType.llm_too_slow, - ) - print("passed test_get_api_base") - - -# Create a mock environment for testing -@pytest.fixture -def mock_env(monkeypatch): - monkeypatch.setenv("SLACK_WEBHOOK_URL", "https://example.com/webhook") - monkeypatch.setenv("LANGFUSE_HOST", "https://cloud.langfuse.com") - monkeypatch.setenv("LANGFUSE_PROJECT_ID", "test-project-id") - - -# Test the __init__ method - - -@pytest.fixture -def slack_alerting(): - return SlackAlerting( - alerting_threshold=1, internal_usage_cache=DualCache(), alerting=["slack"] - ) - - -# Test for slow LLM responses - - - - -# Test for budget crossed - - -# Test for budget crossed again (should not fire alert 2nd time) - - -# Test for send_alert - should be called once -@pytest.mark.asyncio -async def test_send_alert(slack_alerting): - import logging - - from litellm._logging import verbose_logger - - asyncio.create_task(slack_alerting.periodic_flush()) - verbose_logger.setLevel(level=logging.DEBUG) - with patch.object( - slack_alerting.async_http_handler, "post", new=AsyncMock() - ) as mock_post: - mock_post.return_value.status_code = 200 - await slack_alerting.send_alert( - "Test message", "Low", "budget_alerts", alerting_metadata={} - ) - - await asyncio.sleep(6) - mock_post.assert_awaited_once() - - - - -@pytest.mark.asyncio -async def test_daily_reports_completion(slack_alerting): - with patch.object(slack_alerting, "send_alert", new=AsyncMock()) as mock_send_alert: - litellm.callbacks = [slack_alerting] - - # on async success - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-5.5", - "litellm_params": { - "model": "gpt-5-mini", - }, - } - ] - ) - - await router.acompletion( - model="gpt-5-mini", - messages=[{"role": "user", "content": "Hey, how's it going?"}], - ) - - await asyncio.sleep(3) - response_val = await slack_alerting.send_daily_reports(router=router) - - assert response_val is True - - mock_send_alert.assert_awaited_once() - - # on async failure - router = litellm.Router( - model_list=[ - { - "model_name": "gpt-5.5", - "litellm_params": {"model": "gpt-5-mini", "api_key": "bad_key"}, - } - ] - ) - - try: - await router.acompletion( - model="gpt-5-mini", - messages=[{"role": "user", "content": "Hey, how's it going?"}], - ) - except Exception as e: - pass - - await asyncio.sleep(3) - response_val = await slack_alerting.send_daily_reports(router=router) - - assert response_val is True - - mock_send_alert.assert_awaited() - - - - -# test models with 0 metrics are ignored - - -# test no alert is sent if all None or 0 metrics - - -# test user budget crossed alert sent only once, even if user makes multiple calls - - - - -# @pytest.mark.asyncio -# async def test_webhook_customer_spend_event(): -# """ -# Test if customer spend is working as expected -# """ -# slack_alerting = SlackAlerting(alerting=["webhook"]) - -# with patch.object( -# slack_alerting, "send_webhook_alert", new=AsyncMock() -# ) as mock_send_alert: -# user_info = { -# "token": "sk-test-mock-token-606", -# "spend": 1, -# "max_budget": 0, -# "user_id": "ishaan@berri.ai", -# "user_email": "ishaan@berri.ai", -# "key_alias": "my-test-key", -# "projected_exceeded_date": "10/20/2024", -# "projected_spend": 200, -# } - -# user_info = CallInfo(**user_info) -# for _ in range(50): -# await slack_alerting.budget_alerts( -# type=alerting_type, -# user_info=user_info, -# ) -# mock_send_alert.assert_awaited_once() - - - - - - -@pytest.mark.asyncio -async def test_langfuse_trace_id(): - """ - - Unit test for `_add_langfuse_trace_id_to_alert` function in slack_alerting.py - """ - from litellm.integrations.SlackAlerting.utils import add_langfuse_trace_id_to_alert - from litellm.litellm_core_utils.litellm_logging import Logging - - litellm.success_callback = ["langfuse"] - - litellm_logging_obj = Logging( - model="gpt-5-mini", - messages=[{"role": "user", "content": "hi"}], - stream=False, - call_type="acompletion", - litellm_call_id="1234", - start_time=datetime.now(), - function_id="1234", - ) - - litellm.completion( - model="gpt-5-mini", - messages=[{"role": "user", "content": "Hey how's it going?"}], - mock_response="Hey!", - litellm_logging_obj=litellm_logging_obj, - ) - - await asyncio.sleep(3) - - assert litellm_logging_obj.get_trace_id(service_name="langfuse") is not None - - slack_alerting = SlackAlerting( - alerting_threshold=32, - alerting=["slack"], - alert_types=[AlertType.llm_exceptions], - internal_usage_cache=DualCache(), - ) - - trace_url = await add_langfuse_trace_id_to_alert( - request_data={"litellm_logging_obj": litellm_logging_obj} - ) - - assert trace_url is not None - - returned_trace_id = trace_url.split("/")[-1] - - assert returned_trace_id == litellm_logging_obj.get_trace_id( - service_name="langfuse" - ) - - - - -@pytest.mark.parametrize("report_type", ["weekly", "monthly"]) -@pytest.mark.asyncio -async def test_spend_report_cache(report_type): - """ - Test that spend reports are only sent once within their period - """ - # Mock prisma client response - mock_spend_data = [ - {"team_alias": "team1", "total_spend": 100.0}, - {"team_alias": "team2", "total_spend": 200.0}, - ] - - mock_tag_data = [ - {"individual_request_tag": "tag1", "total_spend": 150.0}, - {"individual_request_tag": "tag2", "total_spend": 150.0}, - ] - - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: - # Setup mock for database query - mock_prisma.db.query_raw = AsyncMock( - side_effect=[mock_spend_data, mock_tag_data] - ) - - slack_alerting = SlackAlerting( - alerting=["webhook"], internal_usage_cache=DualCache() - ) - - user_info = CallInfo( - token="test_token", - spend=100, - max_budget=1000, - user_id="test@test.com", - user_email="test@test.com", - key_alias="test-key", - event_group=Litellm_EntityType.KEY, - ) - - with patch.object( - slack_alerting, "send_alert", new=AsyncMock() - ) as mock_send_alert: - # First call should send alert - if report_type == "weekly": - await slack_alerting.send_weekly_spend_report() - else: - await slack_alerting.send_monthly_spend_report() - - mock_send_alert.assert_called_once() - mock_send_alert.reset_mock() - - # Second call should not send alert (cached) - if report_type == "weekly": - await slack_alerting.send_weekly_spend_report() - else: - await slack_alerting.send_monthly_spend_report() - mock_send_alert.assert_not_called() diff --git a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py index bf6294baa00..8c17ee132d8 100644 --- a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py +++ b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py @@ -6,24 +6,15 @@ import litellm import litellm.vector_stores.main import json from typing import Optional -from unittest.mock import AsyncMock, patch, Mock import pytest import litellm -from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( - VectorStorePreCallHook, -) from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.integrations.custom_logger import CustomLogger from litellm.types.utils import ( StandardLoggingPayload, ) -from litellm.types.vector_stores import ( - VectorStoreSearchResponse, - VectorStoreResultContent, - VectorStoreSearchResult, -) class MockCustomLogger(CustomLogger): @@ -63,113 +54,6 @@ def setup_vector_store_registry(): ) -@pytest.mark.asyncio -async def test_vector_store_hook_routes_search_through_proxy_router( - setup_vector_store_registry, -): - proxy_router = Mock() - proxy_router.avector_store_search = AsyncMock( - return_value=VectorStoreSearchResponse( - object="vector_store.search_results.page", - search_query="what is litellm?", - data=[ - VectorStoreSearchResult( - score=1.0, - content=[VectorStoreResultContent(text="routed context", type="text")], - ) - ], - ) - ) - logging_obj = Mock() - logging_obj.model_call_details = { - "litellm_params": {"metadata": {"user_api_key_team_id": "team-a"}} - } - - with patch("litellm.proxy.proxy_server.llm_router", proxy_router): - _, messages, _ = await VectorStorePreCallHook().async_get_chat_completion_prompt( - model="chat-model", - messages=[{"role": "user", "content": "what is litellm?"}], - non_default_params={"vector_store_ids": ["T37J8R4WTM"]}, - prompt_id=None, - prompt_variables=None, - dynamic_callback_params={}, - litellm_logging_obj=logging_obj, - ) - - proxy_router.avector_store_search.assert_awaited_once_with( - vector_store_id="T37J8R4WTM", - query="what is litellm?", - custom_llm_provider="bedrock", - metadata={"user_api_key_team_id": "team-a"}, - ) - assert messages[0]["content"] == "Context:\n\nrouted context\n\n" - - -@pytest.mark.asyncio -async def test_e2e_bedrock_knowledgebase_retrieval_with_completion( - setup_vector_store_registry, -): - litellm.turn_on_debug() - client = AsyncHTTPHandler() - print("value of litellm.vector_store_registry:", litellm.vector_store_registry) - - with patch.object(client, "post") as mock_post: - # Mock the response for the LLM call - mock_response = Mock() - mock_response.status_code = 200 - mock_response.headers = {"Content-Type": "application/json"} - # Provide proper JSON response content - mock_response.text = json.dumps( - { - "id": "msg_01ABC123", - "type": "message", - "role": "assistant", - "content": [ - { - "type": "text", - "text": "LiteLLM is a library that simplifies LLM API access.", - } - ], - "model": "claude-3.5-sonnet", - "stop_reason": "end_turn", - "stop_sequence": None, - "usage": {"input_tokens": 100, "output_tokens": 50}, - } - ) - mock_response.json = lambda: json.loads(mock_response.text) - mock_post.return_value = mock_response - - try: - response = await litellm.acompletion( - model="anthropic/claude-3.5-sonnet", - messages=[{"role": "user", "content": "what is litellm?"}], - vector_store_ids=["T37J8R4WTM"], - client=client, - ) - except Exception as e: - print(f"Error: {e}") - - # Verify the LLM request was made - mock_post.assert_called_once() - - # Verify the request body - print("call args:", mock_post.call_args) - request_body = mock_post.call_args.kwargs["json"] - print("Request body:", json.dumps(request_body, indent=4, default=str)) - - # Assert content from the knowedge base was applied to the request - - # 1. we should have 2 content blocks, the first is the context from the knowledge base, the second is the user message - content = request_body["messages"][0]["content"] - assert len(content) == 2 - assert content[0]["type"] == "text" - assert content[1]["type"] == "text" - - # 2. the first content block should have the bedrock knowledge base prefix string - # this helps confirm that the context from the knowledge base was applied to the request - assert VectorStorePreCallHook.CONTENT_PREFIX_STRING in content[0]["text"] - - @pytest.mark.asyncio async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call( setup_vector_store_registry, @@ -214,65 +98,6 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call( print(f"First search result has {len(first_search_result['data'])} items") -@pytest.mark.asyncio -async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_streaming( - setup_vector_store_registry, -): - """ - Test that the Bedrock Knowledge Base Hook works with streaming and returns search_results in chunks. - """ - - # Init client - # litellm.turn_on_debug() - async_client = AsyncHTTPHandler() - response = await litellm.acompletion( - model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}", - messages=[{"role": "user", "content": "what is litellm?"}], - vector_store_ids=["T37J8R4WTM"], - stream=True, - client=async_client, - ) - - # Collect chunks - chunks = [] - search_results_found = False - async for chunk in response: - chunks.append(chunk) - print(f"Chunk: {chunk}") - - # Check if this chunk has search_results in provider_specific_fields - if hasattr(chunk, "choices") and chunk.choices: - for choice in chunk.choices: - if hasattr(choice, "delta") and choice.delta: - provider_fields = getattr( - choice.delta, "provider_specific_fields", None - ) - if provider_fields and "search_results" in provider_fields: - search_results = provider_fields["search_results"] - print( - f"Found search_results in streaming chunk: {len(search_results)} results" - ) - - # Verify structure - assert search_results is not None - assert len(search_results) > 0 - - first_search_result = search_results[0] - assert "object" in first_search_result - assert ( - first_search_result["object"] - == "vector_store.search_results.page" - ) - assert "data" in first_search_result - assert len(first_search_result["data"]) > 0 - - search_results_found = True - - print(f"Total chunks received: {len(chunks)}") - assert len(chunks) > 0 - assert search_results_found, "search_results should be present in streaming chunks" - - @pytest.mark.asyncio async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_with_tools( setup_vector_store_registry, @@ -345,328 +170,6 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_with_tools_ print(f" Search was performed and {len(search_results)} result(s) returned") -@pytest.mark.asyncio -async def test_bedrock_kb_request_body_has_transformed_filters( - setup_vector_store_registry, -): - """ - Validate that the Bedrock Knowledge Base request body contains the transformed filters. - """ - captured_request_body: dict = {} - - async def fake_async_vector_store_search_handler( - vector_store_id, - query, - vector_store_search_optional_params, - vector_store_provider_config, - custom_llm_provider, - litellm_params, - logging_obj, - embedding_executor=None, - extra_headers=None, - extra_body=None, - timeout=None, - client=None, - _is_async=False, - ): - litellm_params_dict = ( - litellm_params.model_dump(exclude_none=False) - if hasattr(litellm_params, "model_dump") - else dict(litellm_params) - ) - api_base = vector_store_provider_config.get_complete_url( - api_base=litellm_params_dict.get("api_base"), - litellm_params=litellm_params_dict, - ) - - url, request_body = ( - vector_store_provider_config.transform_search_vector_store_request( - vector_store_id=vector_store_id, - query=query, - vector_store_search_optional_params=vector_store_search_optional_params, - api_base=api_base, - litellm_logging_obj=logging_obj, - litellm_params=litellm_params_dict, - extra_body=None, - ) - ) - captured_request_body["url"] = url - captured_request_body["body"] = request_body - - return VectorStoreSearchResponse( - object="vector_store.search_results.page", - search_query=query if isinstance(query, str) else " ".join(query), - data=[ - VectorStoreSearchResult( - score=0.9, - content=[ - VectorStoreResultContent( - text="LiteLLM is a library", type="text" - ) - ], - ) - ], - ) - - with patch.object( - litellm.vector_stores.main.base_llm_http_handler, - "async_vector_store_search_handler", - new=AsyncMock(side_effect=fake_async_vector_store_search_handler), - ): - response = await litellm.acompletion( - model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}", - messages=[{"role": "user", "content": "what is litellm?"}], - max_tokens=10, - tools=[ - { - "type": "file_search", - "vector_store_ids": ["T37J8R4WTM"], - "filters": { - "key": "user_id", - "value": "fake-user-id", - "operator": "eq", - }, - } - ], - ) - - assert response is not None - print( - "captured_request_body:", - json.dumps(captured_request_body, indent=4, default=str), - ) - assert "body" in captured_request_body, "Bedrock KB request body was not captured" - - vector_search = captured_request_body["body"]["retrievalConfiguration"][ - "vectorSearchConfiguration" - ] - aws_filter = vector_search["filter"] - assert "equals" in aws_filter, f"Expected 'equals' in AWS format, got: {aws_filter}" - assert aws_filter["equals"]["key"] == "user_id" - assert aws_filter["equals"]["value"] == "fake-user-id" - - print("✅ Filters transformed correctly: OpenAI format -> AWS Bedrock format") - - -@pytest.mark.asyncio -async def test_openai_with_knowledge_base_mock_openai(setup_vector_store_registry): - """ - Tests that knowledge base content is correctly passed to the OpenAI API call - """ - litellm.set_verbose = True - from openai import AsyncOpenAI - - client = AsyncOpenAI(api_key="fake-api-key") - - # Variable to capture the request - captured_request = {} - - with patch.object( - client.chat.completions.with_raw_response, "create" - ) as mock_client: - # Create async mock that returns proper structure - async def mock_create(**kwargs): - mock_response = Mock() - mock_response.choices = [ - Mock( - message=Mock(content="Mock response from OpenAI", role="assistant") - ) - ] - mock_response.usage = Mock( - prompt_tokens=100, completion_tokens=50, total_tokens=150 - ) - mock_response.id = "chatcmpl-123" - mock_response.object = "chat.completion" - mock_response.created = 1234567890 - mock_response.model = "gpt-5.5" - - # Store the request for verification - captured_request.update(kwargs) - - # Return wrapper with parse method - wrapper = Mock() - wrapper.parse.return_value = mock_response - return wrapper - - mock_client.side_effect = mock_create - - try: - await litellm.acompletion( - model="gpt-5.5", - messages=[{"role": "user", "content": "what is litellm?"}], - vector_store_ids=["T37J8R4WTM"], - client=client, - ) - except Exception as e: - print(f"Error: {e}") - - # Verify the API was called - mock_client.assert_called_once() - request_body = captured_request - - # Verify the request contains messages with knowledge base context - assert "messages" in request_body - messages = request_body["messages"] - - # We expect at least 2 messages: - # 1. User message with the knowledge base context - # 2. User message with the question - assert len(messages) >= 2 - - print("request messages:", json.dumps(messages, indent=4, default=str)) - - # assert message[0] is the user message with the knowledge base context - assert messages[0]["role"] == "user" - assert VectorStorePreCallHook.CONTENT_PREFIX_STRING in messages[0]["content"] - - -@pytest.mark.asyncio -async def test_openai_with_vector_store_ids_in_tool_call_mock_openai( - setup_vector_store_registry, -): - """ - Tests that vector store ids can be passed as tools - - This is the OpenAI format - """ - litellm.set_verbose = True - from openai import AsyncOpenAI - - client = AsyncOpenAI(api_key="fake-api-key") - - # Variable to capture the request - captured_request = {} - - with patch.object( - client.chat.completions.with_raw_response, "create" - ) as mock_client: - # Create async mock that returns proper structure - async def mock_create(**kwargs): - mock_response = Mock() - mock_response.choices = [ - Mock( - message=Mock(content="Mock response from OpenAI", role="assistant") - ) - ] - mock_response.usage = Mock( - prompt_tokens=100, completion_tokens=50, total_tokens=150 - ) - mock_response.id = "chatcmpl-123" - mock_response.object = "chat.completion" - mock_response.created = 1234567890 - mock_response.model = "gpt-5.5" - - # Store the request for verification - captured_request.update(kwargs) - - # Return wrapper with parse method - wrapper = Mock() - wrapper.parse.return_value = mock_response - return wrapper - - mock_client.side_effect = mock_create - - try: - await litellm.acompletion( - model="gpt-5.5", - messages=[{"role": "user", "content": "what is litellm?"}], - tools=[{"type": "file_search", "vector_store_ids": ["T37J8R4WTM"]}], - client=client, - ) - except Exception as e: - print(f"Error: {e}") - - # Verify the API was called - mock_client.assert_called_once() - request_body = captured_request - print("request body:", json.dumps(request_body, indent=4, default=str)) - - # Verify the request contains messages with knowledge base context - assert "messages" in request_body - messages = request_body["messages"] - - # We expect at least 2 messages: - # 1. User message with the knowledge base context - # 2. User message with the question - assert len(messages) >= 2 - - print("request messages:", json.dumps(messages, indent=4, default=str)) - - # assert message[0] is the user message with the knowledge base context - assert messages[0]["role"] == "user" - assert VectorStorePreCallHook.CONTENT_PREFIX_STRING in messages[0]["content"] - - # assert that the tool call was not sent to the upstream llm API if it's a litellm vector store - assert "tools" not in request_body - - -@pytest.mark.asyncio -async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_registry): - """Ensure unrecognized vector store tools are forwarded to the provider""" - from openai import AsyncOpenAI - - client = AsyncOpenAI(api_key="fake-api-key") - - # Variable to capture the request - captured_request = {} - - with patch.object( - client.chat.completions.with_raw_response, "create" - ) as mock_client: - # Create async mock that returns proper structure - async def mock_create(**kwargs): - mock_response = Mock() - mock_response.choices = [ - Mock( - message=Mock(content="Mock response from OpenAI", role="assistant") - ) - ] - mock_response.usage = Mock( - prompt_tokens=100, completion_tokens=50, total_tokens=150 - ) - mock_response.id = "chatcmpl-123" - mock_response.object = "chat.completion" - mock_response.created = 1234567890 - mock_response.model = "gpt-5.5" - - # Store the request for verification - captured_request.update(kwargs) - - # Return wrapper with parse method - wrapper = Mock() - wrapper.parse.return_value = mock_response - return wrapper - - mock_client.side_effect = mock_create - - try: - await litellm.acompletion( - model="gpt-5.5", - messages=[{"role": "user", "content": "what is litellm?"}], - tools=[ - {"type": "file_search", "vector_store_ids": ["T37J8R4WTM"]}, - {"type": "file_search", "vector_store_ids": ["unknownVS"]}, - ], - client=client, - ) - except Exception as e: - print(f"Error: {e}") - - mock_client.assert_called_once() - request_body = captured_request - - assert "messages" in request_body - messages = request_body["messages"] - assert len(messages) >= 2 - assert messages[0]["role"] == "user" - assert VectorStorePreCallHook.CONTENT_PREFIX_STRING in messages[0]["content"] - - assert "tools" in request_body - tools = request_body["tools"] - assert len(tools) == 1 - assert tools[0]["vector_store_ids"] == ["unknownVS"] - - # @pytest.mark.asyncio # async def test_logging_with_knowledge_base_hook(setup_vector_store_registry): # """ @@ -723,118 +226,3 @@ async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_regist -@pytest.mark.asyncio -async def test_provider_specific_fields_in_proxy_http_response( - setup_vector_store_registry, -): - """ - Test that provider_specific_fields (like search_results) are included - in the proxy HTTP JSON response, not just in Python SDK objects. - - This test catches serialization bugs where exclude=True would strip - provider_specific_fields from the HTTP response. - """ - from fastapi.testclient import TestClient - from litellm.proxy.proxy_server import app, initialize - from unittest.mock import patch as mock_patch - - # Initialize proxy - await initialize( - model="gpt-5-mini", - alias=None, - api_base=None, - debug=False, - temperature=None, - max_tokens=None, - request_timeout=600, - max_budget=None, - drop_params=True, - add_function_to_prompt=False, - headers=None, - save=False, - use_queue=False, - config=None, - ) - - # Create test client - client = TestClient(app) - - # Create mock response with provider_specific_fields - mock_response = litellm.ModelResponse( - id="test-123", - model="gpt-5-mini", - created=1234567890, - object="chat.completion", - ) - - # Create message with provider_specific_fields - mock_message = litellm.Message( - content="LiteLLM is a tool that simplifies working with multiple LLMs.", - role="assistant", - provider_specific_fields={ - "search_results": [ - { - "object": "vector_store.search_results.page", - "search_query": "what is litellm?", - "data": [ - { - "score": 0.95, - "content": [{"text": "Test content", "type": "text"}], - "file_id": "test-file", - "filename": "test.txt", - } - ], - } - ] - }, - ) - - mock_choice = litellm.Choices(finish_reason="stop", index=0, message=mock_message) - - mock_response.choices = [mock_choice] - mock_response.usage = litellm.Usage( - prompt_tokens=10, completion_tokens=20, total_tokens=30 - ) - - # Patch the completion call at the proxy level - with mock_patch("litellm.acompletion", new=AsyncMock(return_value=mock_response)): - # Make HTTP request to proxy - response = client.post( - "/v1/chat/completions", - json={ - "model": "gpt-5-mini", - "messages": [{"role": "user", "content": "What is litellm?"}], - }, - ) - - # Check HTTP response - assert response.status_code == 200 - result = response.json() - - print("HTTP Response JSON:", json.dumps(result, indent=2)) - - # THE KEY ASSERTIONS - These would FAIL with exclude=True! - assert "choices" in result - assert len(result["choices"]) > 0 - - choice = result["choices"][0] - assert "message" in choice - - message = choice["message"] - - # Verify provider_specific_fields is in the JSON response - assert ( - "provider_specific_fields" in message - ), "provider_specific_fields missing from HTTP JSON response! This means exclude=True is preventing serialization." - - assert "search_results" in message["provider_specific_fields"] - search_results = message["provider_specific_fields"]["search_results"] - assert len(search_results) > 0 - - # Verify search result structure - first_result = search_results[0] - assert first_result["object"] == "vector_store.search_results.page" - assert "data" in first_result - assert len(first_result["data"]) > 0 - - print("✅ provider_specific_fields successfully serialized in HTTP response") diff --git a/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py b/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py deleted file mode 100644 index 2d3933c3dad..00000000000 --- a/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py +++ /dev/null @@ -1,160 +0,0 @@ -import traceback -from litellm._uuid import uuid -import pytest -from dotenv import load_dotenv -from fastapi import Request -from fastapi.routing import APIRoute - -load_dotenv() -import io -import time -import json - -# this file is to test litellm/proxy - -import litellm -import asyncio -from typing import Optional -from litellm.types.utils import StandardLoggingPayload, Usage, ModelInfoBase -from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import ( - StandardBuiltInToolCostTracking, -) - - -class TestCustomLogger(CustomLogger): - def __init__(self): - self.recorded_usage: Optional[Usage] = None - self.standard_logging_payload: Optional[StandardLoggingPayload] = None - - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - standard_logging_payload = kwargs.get("standard_logging_object") - self.standard_logging_payload = standard_logging_payload - print( - "standard_logging_payload", - json.dumps(standard_logging_payload, indent=4, default=str), - ) - - self.recorded_usage = Usage( - prompt_tokens=standard_logging_payload.get("prompt_tokens"), - completion_tokens=standard_logging_payload.get("completion_tokens"), - total_tokens=standard_logging_payload.get("total_tokens"), - ) - pass - - -async def _setup_web_search_test(): - """Helper function to setup common test requirements""" - litellm.turn_on_debug() - test_custom_logger = TestCustomLogger() - litellm.callbacks = [test_custom_logger] - return test_custom_logger - - -async def _verify_web_search_cost(test_custom_logger, expected_context_size): - """Helper function to verify web search costs""" - await asyncio.sleep(1) - - standard_logging_payload = test_custom_logger.standard_logging_payload - response = standard_logging_payload.get("response") - response_cost = standard_logging_payload.get("response_cost") - assert response_cost is not None - - # Calculate token cost - model_map_information = standard_logging_payload["model_map_information"] - model_map_value: ModelInfoBase = model_map_information["model_map_value"] - total_token_cost = ( - standard_logging_payload["prompt_tokens"] - * model_map_value["input_cost_per_token"] - ) + ( - standard_logging_payload["completion_tokens"] - * model_map_value["output_cost_per_token"] - ) - - # Verify total cost - if StandardBuiltInToolCostTracking.response_object_includes_web_search_call( - response - ): - assert ( - response_cost - == total_token_cost - + model_map_value["search_context_cost_per_query"][expected_context_size] - ) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "web_search_options,expected_context_size", - [ - (None, "search_context_size_medium"), - ({"search_context_size": "low"}, "search_context_size_low"), - ({"search_context_size": "high"}, "search_context_size_high"), - ], -) -async def test_openai_web_search_logging_cost_tracking( - web_search_options, expected_context_size -): - """Test web search cost tracking with different search context sizes""" - test_custom_logger = await _setup_web_search_test() - - request_kwargs = { - "model": "openai/gpt-5-search-api", - "messages": [ - { - "role": "user", - "content": f"What was a positive news story from today? {uuid.uuid4()}", - } - ], - } - if web_search_options is not None: - request_kwargs["web_search_options"] = web_search_options - - response = await litellm.acompletion(**request_kwargs) - - await _verify_web_search_cost(test_custom_logger, expected_context_size) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "tools_config,expected_context_size,stream", - [ - ( - [{"type": "web_search_preview", "search_context_size": "low"}], - "search_context_size_low", - True, - ), - ( - [{"type": "web_search_preview", "search_context_size": "low"}], - "search_context_size_low", - False, - ), - ([{"type": "web_search_preview"}], "search_context_size_medium", True), - ([{"type": "web_search_preview"}], "search_context_size_medium", False), - ], -) -async def test_openai_responses_api_web_search_cost_tracking( - tools_config, expected_context_size, stream -): - """Test web search cost tracking with different search context sizes and streaming options""" - test_custom_logger = await _setup_web_search_test() - - response = await litellm.aresponses( - model="openai/gpt-4o", - input=[ - {"role": "user", "content": "What was a positive news story from today?"} - ], - tools=tools_config, - stream=stream, - ) - if stream is True: - async for chunk in response: - print("chunk", chunk) - else: - print("response", response) - - await asyncio.sleep(1) - - if StandardBuiltInToolCostTracking.response_object_includes_web_search_call( - test_custom_logger.standard_logging_payload.get("response") - ): - await _verify_web_search_cost(test_custom_logger, expected_context_size) diff --git a/tests/logging_callback_tests/test_custom_callback_router.py b/tests/logging_callback_tests/test_custom_callback_router.py deleted file mode 100644 index 4a12f8d536d..00000000000 --- a/tests/logging_callback_tests/test_custom_callback_router.py +++ /dev/null @@ -1,754 +0,0 @@ -### What this tests #### -## This test asserts the type of data passed into each method of the custom callback handler -import asyncio -import inspect -import os -import time -import traceback -from datetime import datetime - -import pytest - -from typing import List, Literal, Optional - -import litellm -from litellm import Cache, Router -from litellm.integrations.custom_logger import CustomLogger - -# Test Scenarios (test across completion, streaming, embedding) -## 1: Pre-API-Call -## 2: Post-API-Call -## 3: On LiteLLM Call success -## 4: On LiteLLM Call failure -## fallbacks -## retries - -# Test cases -## 1. Simple Azure OpenAI acompletion + streaming call -## 2. Simple Azure OpenAI aembedding call -## 3. Azure OpenAI acompletion + streaming call with retries -## 4. Azure OpenAI aembedding call with retries -## 5. Azure OpenAI acompletion + streaming call with fallbacks -## 6. Azure OpenAI aembedding call with fallbacks - -## Test interfaces -## 1. router.completion() + router.embeddings() -## 2. proxy.completions + proxy.embeddings - -litellm.num_retries = 0 - - -class CompletionCustomHandler( - CustomLogger -): # https://docs.litellm.ai/docs/observability/custom_callback#callback-class - """ - The set of expected inputs to a custom handler for a - """ - - # Class variables or attributes - def __init__(self): - self.errors = [] - self.states: Optional[ - List[ - Literal[ - "sync_pre_api_call", - "async_pre_api_call", - "post_api_call", - "sync_stream", - "async_stream", - "sync_success", - "async_success", - "sync_failure", - "async_failure", - ] - ] - ] = [] - - def log_pre_api_call(self, model, messages, kwargs): - try: - print(f"received kwargs in pre-input: {kwargs}") - self.states.append("sync_pre_api_call") - ## MODEL - assert isinstance(model, str) - ## MESSAGES - assert isinstance(messages, list) - ## KWARGS - assert isinstance(kwargs["model"], str) - assert isinstance(kwargs["messages"], list) - assert isinstance(kwargs["optional_params"], dict) - assert isinstance(kwargs["litellm_params"], dict) - assert isinstance(kwargs["start_time"], (datetime, type(None))) - assert isinstance(kwargs["stream"], bool) - assert isinstance(kwargs["user"], (str, type(None))) - ### ROUTER-SPECIFIC KWARGS - assert isinstance(kwargs["litellm_params"]["metadata"], dict) - assert isinstance(kwargs["litellm_params"]["metadata"]["model_group"], str) - assert isinstance(kwargs["litellm_params"]["metadata"]["deployment"], str) - assert isinstance(kwargs["litellm_params"]["model_info"], dict) - assert isinstance(kwargs["litellm_params"]["model_info"]["id"], str) - assert isinstance( - kwargs["litellm_params"]["proxy_server_request"], (str, type(None)) - ) - assert isinstance( - kwargs["litellm_params"]["preset_cache_key"], (str, type(None)) - ) - assert isinstance(kwargs["litellm_params"]["stream_response"], dict) - except Exception as e: - print(f"Assertion Error: {traceback.format_exc()}") - self.errors.append(traceback.format_exc()) - - def log_post_api_call(self, kwargs, response_obj, start_time, end_time): - try: - self.states.append("post_api_call") - ## START TIME - assert isinstance(start_time, datetime) - ## END TIME - assert end_time == None - ## RESPONSE OBJECT - assert response_obj == None - ## KWARGS - assert isinstance(kwargs["model"], str) - assert isinstance(kwargs["messages"], list) - assert isinstance(kwargs["optional_params"], dict) - assert isinstance(kwargs["litellm_params"], dict) - assert isinstance(kwargs["start_time"], (datetime, type(None))) - assert isinstance(kwargs["stream"], bool) - assert isinstance(kwargs["user"], (str, type(None))) - assert isinstance(kwargs["input"], (list, dict, str)) - assert isinstance(kwargs["api_key"], (str, type(None))) - assert ( - isinstance( - kwargs["original_response"], (str, litellm.CustomStreamWrapper) - ) - or inspect.iscoroutine(kwargs["original_response"]) - or inspect.isasyncgen(kwargs["original_response"]) - ) - assert isinstance(kwargs["additional_args"], (dict, type(None))) - assert isinstance(kwargs["log_event_type"], str) - ### ROUTER-SPECIFIC KWARGS - assert isinstance(kwargs["litellm_params"]["metadata"], dict) - assert isinstance(kwargs["litellm_params"]["metadata"]["model_group"], str) - assert isinstance(kwargs["litellm_params"]["metadata"]["deployment"], str) - assert isinstance(kwargs["litellm_params"]["model_info"], dict) - assert isinstance(kwargs["litellm_params"]["model_info"]["id"], str) - assert isinstance( - kwargs["litellm_params"]["proxy_server_request"], (str, type(None)) - ) - assert isinstance( - kwargs["litellm_params"]["preset_cache_key"], (str, type(None)) - ) - assert isinstance(kwargs["litellm_params"]["stream_response"], dict) - except Exception: - print(f"Assertion Error: {traceback.format_exc()}") - self.errors.append(traceback.format_exc()) - - async def async_log_stream_event(self, kwargs, response_obj, start_time, end_time): - try: - self.states.append("async_stream") - ## START TIME - assert isinstance(start_time, datetime) - ## END TIME - assert isinstance(end_time, datetime) - ## RESPONSE OBJECT - assert isinstance(response_obj, litellm.ModelResponseStream) - ## KWARGS - assert isinstance(kwargs["model"], str) - assert isinstance(kwargs["messages"], list) and isinstance( - kwargs["messages"][0], dict - ) - assert isinstance(kwargs["optional_params"], dict) - assert isinstance(kwargs["litellm_params"], dict) - assert isinstance(kwargs["start_time"], (datetime, type(None))) - assert isinstance(kwargs["stream"], bool) - assert isinstance(kwargs["user"], (str, type(None))) - assert ( - isinstance(kwargs["input"], list) - and isinstance(kwargs["input"][0], dict) - ) or isinstance(kwargs["input"], (dict, str)) - assert isinstance(kwargs["api_key"], (str, type(None))) - assert ( - isinstance( - kwargs["original_response"], (str, litellm.CustomStreamWrapper) - ) - or inspect.isasyncgen(kwargs["original_response"]) - or inspect.iscoroutine(kwargs["original_response"]) - ) - assert isinstance(kwargs["additional_args"], (dict, type(None))) - assert isinstance(kwargs["log_event_type"], str) - except Exception: - print(f"Assertion Error: {traceback.format_exc()}") - self.errors.append(traceback.format_exc()) - - def log_success_event(self, kwargs, response_obj, start_time, end_time): - try: - self.states.append("sync_success") - ## START TIME - assert isinstance(start_time, datetime) - ## END TIME - assert isinstance(end_time, datetime) - ## RESPONSE OBJECT - assert isinstance(response_obj, litellm.ModelResponse) - ## KWARGS - assert isinstance(kwargs["model"], str) - assert isinstance(kwargs["messages"], list) and isinstance( - kwargs["messages"][0], dict - ) - assert isinstance(kwargs["optional_params"], dict) - assert isinstance(kwargs["litellm_params"], dict) - assert isinstance(kwargs["start_time"], (datetime, type(None))) - assert isinstance(kwargs["stream"], bool) - assert isinstance(kwargs["user"], (str, type(None))) - assert ( - isinstance(kwargs["input"], list) - and isinstance(kwargs["input"][0], dict) - ) or isinstance(kwargs["input"], (dict, str)) - assert isinstance(kwargs["api_key"], (str, type(None))) - assert isinstance( - kwargs["original_response"], (str, litellm.CustomStreamWrapper) - ) - assert isinstance(kwargs["additional_args"], (dict, type(None))) - assert isinstance(kwargs["log_event_type"], str) - assert kwargs["cache_hit"] is None or isinstance(kwargs["cache_hit"], bool) - except Exception: - print(f"Assertion Error: {traceback.format_exc()}") - self.errors.append(traceback.format_exc()) - - def log_failure_event(self, kwargs, response_obj, start_time, end_time): - try: - self.states.append("sync_failure") - ## START TIME - assert isinstance(start_time, datetime) - ## END TIME - assert isinstance(end_time, datetime) - ## RESPONSE OBJECT - assert response_obj == None - ## KWARGS - assert isinstance(kwargs["model"], str) - assert isinstance(kwargs["messages"], list) and isinstance( - kwargs["messages"][0], dict - ) - assert isinstance(kwargs["optional_params"], dict) - assert isinstance(kwargs["litellm_params"], dict) - assert isinstance(kwargs["start_time"], (datetime, type(None))) - assert isinstance(kwargs["stream"], bool) - assert isinstance(kwargs["user"], (str, type(None))) - assert ( - isinstance(kwargs["input"], list) - and isinstance(kwargs["input"][0], dict) - ) or isinstance(kwargs["input"], (dict, str)) - assert isinstance(kwargs["api_key"], (str, type(None))) - assert ( - isinstance( - kwargs["original_response"], (str, litellm.CustomStreamWrapper) - ) - or kwargs["original_response"] == None - ) - assert isinstance(kwargs["additional_args"], (dict, type(None))) - assert isinstance(kwargs["log_event_type"], str) - except Exception: - print(f"Assertion Error: {traceback.format_exc()}") - self.errors.append(traceback.format_exc()) - - async def async_log_pre_api_call(self, model, messages, kwargs): - try: - """ - No-op. - Not implemented yet. - """ - pass - except Exception as e: - print(f"Assertion Error: {traceback.format_exc()}") - self.errors.append(traceback.format_exc()) - - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - try: - print("CompletionCustomHandler.async_log_success_event, kwargs: ", kwargs) - self.states.append("async_success") - print( - "############### CompletionCustomHandler async success, kwargs: ", - kwargs, - ) - ## START TIME - assert isinstance(start_time, datetime) - ## END TIME - assert isinstance(end_time, datetime) - ## RESPONSE OBJECT - assert isinstance( - response_obj, (litellm.ModelResponse, litellm.EmbeddingResponse) - ) - ## KWARGS - assert isinstance(kwargs["model"], str) - - # checking we use base_model for azure cost calculation - base_model = litellm.utils.get_base_model_from_metadata( - model_call_details=kwargs - ) - - if ( - kwargs["model"] == "chatgpt-v-3" - and base_model is not None - and kwargs["stream"] != True - ): - # when base_model is set for azure, we should use pricing for the base_model - # this checks response_cost == litellm.cost_per_token(model=base_model) - assert isinstance(kwargs["response_cost"], float) - response_cost = kwargs["response_cost"] - print( - f"response_cost: {response_cost}, for model: {kwargs['model']} and base_model: {base_model}" - ) - prompt_tokens = response_obj.usage.prompt_tokens - completion_tokens = response_obj.usage.completion_tokens - # ensure the pricing is based on the base_model here - prompt_price, completion_price = litellm.cost_per_token( - model=base_model, - prompt_tokens=prompt_tokens, - completion_tokens=completion_tokens, - ) - expected_price = prompt_price + completion_price - print(f"expected price: {expected_price}") - assert ( - response_cost == expected_price - ), f"response_cost: {response_cost} != expected_price: {expected_price}. For model: {kwargs['model']} and base_model: {base_model}. should have used base_model for price" - - assert isinstance(kwargs["messages"], list) - assert isinstance(kwargs["optional_params"], dict) - assert isinstance(kwargs["litellm_params"], dict) - assert isinstance(kwargs["start_time"], (datetime, type(None))) - assert isinstance(kwargs["stream"], bool) - assert isinstance(kwargs["user"], (str, type(None))) - assert isinstance(kwargs["input"], (list, dict, str)) - assert isinstance(kwargs["api_key"], (str, type(None))) - assert ( - isinstance( - kwargs["original_response"], (str, litellm.CustomStreamWrapper) - ) - or inspect.isasyncgen(kwargs["original_response"]) - or inspect.iscoroutine(kwargs["original_response"]) - ) - assert isinstance(kwargs["additional_args"], (dict, type(None))) - assert isinstance(kwargs["log_event_type"], str) - assert kwargs["cache_hit"] is None or isinstance(kwargs["cache_hit"], bool) - ### ROUTER-SPECIFIC KWARGS - assert isinstance(kwargs["litellm_params"]["metadata"], dict) - assert isinstance(kwargs["litellm_params"]["metadata"]["model_group"], str) - assert isinstance(kwargs["litellm_params"]["metadata"]["deployment"], str) - assert isinstance(kwargs["litellm_params"]["model_info"], dict) - assert isinstance(kwargs["litellm_params"]["model_info"]["id"], str) - assert isinstance( - kwargs["litellm_params"]["proxy_server_request"], (str, type(None)) - ) - assert isinstance( - kwargs["litellm_params"]["preset_cache_key"], (str, type(None)) - ) - assert isinstance(kwargs["litellm_params"]["stream_response"], dict) - except Exception: - print(f"Assertion Error: {traceback.format_exc()}") - self.errors.append(traceback.format_exc()) - - async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): - try: - print(f"received original response: {kwargs['original_response']}") - self.states.append("async_failure") - ## START TIME - assert isinstance(start_time, datetime) - ## END TIME - assert isinstance(end_time, datetime) - ## RESPONSE OBJECT - assert response_obj == None - ## KWARGS - assert isinstance(kwargs["model"], str) - assert isinstance(kwargs["messages"], list) - assert isinstance(kwargs["optional_params"], dict) - assert isinstance(kwargs["litellm_params"], dict) - assert isinstance(kwargs["start_time"], (datetime, type(None))) - assert isinstance(kwargs["stream"], bool) - assert isinstance(kwargs["user"], (str, type(None))) - assert isinstance(kwargs["input"], (list, str, dict)) - assert isinstance(kwargs["api_key"], (str, type(None))) - assert ( - isinstance( - kwargs["original_response"], (str, litellm.CustomStreamWrapper) - ) - or inspect.isasyncgen(kwargs["original_response"]) - or inspect.iscoroutine(kwargs["original_response"]) - or kwargs["original_response"] == None - ) - assert isinstance(kwargs["additional_args"], (dict, type(None))) - assert isinstance(kwargs["log_event_type"], str) - except Exception: - print(f"Assertion Error: {traceback.format_exc()}") - self.errors.append(traceback.format_exc()) - - -# Simple Azure OpenAI call -## COMPLETION -# @pytest.mark.flaky(retries=5, delay=1) -@pytest.mark.asyncio -async def test_async_chat_azure(): - try: - customHandler_completion_azure_router = CompletionCustomHandler() - customHandler_streaming_azure_router = CompletionCustomHandler() - customHandler_failure = CompletionCustomHandler() - litellm.callbacks = [customHandler_completion_azure_router] - litellm.set_verbose = True - model_list = [ - { - "model_name": "gpt-4.1-nano", # openai model name - "litellm_params": { # params for litellm completion/embedding call - "model": "azure/gpt-4.1-mini", - "api_key": os.getenv("AZURE_AI_API_KEY"), - "api_version": os.getenv("AZURE_API_VERSION"), - "api_base": os.getenv("AZURE_AI_API_BASE"), - }, - "model_info": {"base_model": "azure/gpt-4.1-mini"}, - "tpm": 240000, - "rpm": 1800, - }, - ] - router = Router(model_list=model_list, num_retries=0) # type: ignore - response = await router.acompletion( - model="gpt-4.1-nano", - messages=[{"role": "user", "content": "Hi 👋 - i'm openai"}], - ) - print("got response, sleeping 5 seconds....") - await asyncio.sleep(5) - assert len(customHandler_completion_azure_router.errors) == 0 - assert ( - len(customHandler_completion_azure_router.states) == 3 - ) # pre, post, success - # streaming - - litellm.logging_callback_manager._reset_all_callbacks() - litellm.callbacks = [customHandler_streaming_azure_router] - router2 = Router(model_list=model_list, num_retries=0) # type: ignore - response = await router2.acompletion( - model="gpt-4.1-nano", - messages=[{"role": "user", "content": "Hi 👋 - i'm openai"}], - stream=True, - ) - async for chunk in response: - print(f"async azure router chunk: {chunk}") - continue - await asyncio.sleep(5) - print(f"customHandler.states: {customHandler_streaming_azure_router.states}") - assert len(customHandler_streaming_azure_router.errors) == 0 - assert ( - len(customHandler_streaming_azure_router.states) >= 3 - ) # pre, post, stream (multiple times), success - # failure - model_list = [ - { - "model_name": "gpt-5-mini", # openai model name - "litellm_params": { # params for litellm completion/embedding call - "model": "azure/gpt-4o-new-test", - "api_key": "my-bad-key", - "api_version": os.getenv("AZURE_API_VERSION"), - "api_base": os.getenv("AZURE_AI_API_BASE"), - }, - "tpm": 240000, - "rpm": 1800, - }, - ] - - litellm.logging_callback_manager._reset_all_callbacks() - litellm.callbacks = [customHandler_failure] - router3 = Router(model_list=model_list, num_retries=0) # type: ignore - try: - response = await router3.acompletion( - model="gpt-5-mini", - messages=[{"role": "user", "content": "Hi 👋 - i'm openai"}], - ) - print(f"response in router3 acompletion: {response}") - except Exception: - pass - await asyncio.sleep(5) - print(f"customHandler.states: {customHandler_failure.states}") - assert len(customHandler_failure.errors) == 0 - assert len(customHandler_failure.states) == 3 # pre, post, failure - assert "async_failure" in customHandler_failure.states - except Exception as e: - print(f"Assertion Error: {traceback.format_exc()}") - pytest.fail(f"An exception occurred - {str(e)}") - - -## EMBEDDING -@pytest.mark.asyncio -async def test_async_embedding_azure(): - try: - customHandler = CompletionCustomHandler() - customHandler_failure = CompletionCustomHandler() - litellm.callbacks = [customHandler] - model_list = [ - { - "model_name": "azure-embedding-model", # openai model name - "litellm_params": { # params for litellm completion/embedding call - "model": "azure/text-embedding-ada-002", - "api_key": os.getenv("AZURE_AI_API_KEY"), - "api_version": os.getenv("AZURE_API_VERSION"), - "api_base": os.getenv("AZURE_AI_API_BASE"), - }, - "tpm": 240000, - "rpm": 1800, - }, - ] - router = Router(model_list=model_list) # type: ignore - response = await router.aembedding( - model="azure-embedding-model", input=["hello from litellm!"] - ) - await asyncio.sleep(2) - assert len(customHandler.errors) == 0 - assert len(customHandler.states) == 3 # pre, post, success - # failure - model_list = [ - { - "model_name": "azure-embedding-model", # openai model name - "litellm_params": { # params for litellm completion/embedding call - "model": "azure/text-embedding-ada-002", - "api_key": "my-bad-key", - "api_version": os.getenv("AZURE_API_VERSION"), - "api_base": os.getenv("AZURE_AI_API_BASE"), - }, - "tpm": 240000, - "rpm": 1800, - }, - ] - litellm.logging_callback_manager._reset_all_callbacks() - litellm.callbacks = [customHandler_failure] - router3 = Router(model_list=model_list, num_retries=0) # type: ignore - try: - response = await router3.aembedding( - model="azure-embedding-model", input=["hello from litellm!"] - ) - print(f"response in router3 aembedding: {response}") - except Exception: - pass - await asyncio.sleep(1) - print(f"customHandler.states: {customHandler_failure.states}") - assert len(customHandler_failure.errors) == 0 - assert len(customHandler_failure.states) == 3 # pre, post, failure - assert "async_failure" in customHandler_failure.states - except Exception as e: - print(f"Assertion Error: {traceback.format_exc()}") - pytest.fail(f"An exception occurred - {str(e)}") - - -# asyncio.run(test_async_embedding_azure()) -# Azure OpenAI call w/ Fallbacks -## COMPLETION -@pytest.mark.asyncio -async def test_async_chat_azure_with_fallbacks(): - try: - customHandler_fallbacks = CompletionCustomHandler() - litellm.callbacks = [customHandler_fallbacks] - litellm.set_verbose = True - # with fallbacks - model_list = [ - { - "model_name": "gpt-5-mini", # openai model name - "litellm_params": { # params for litellm completion/embedding call - "model": "azure/gpt-4.1-mini", - "api_key": "my-bad-key", - "api_version": os.getenv("AZURE_API_VERSION"), - "api_base": os.getenv("AZURE_AI_API_BASE"), - }, - "tpm": 240000, - "rpm": 1800, - }, - { - "model_name": "gpt-3.5-turbo-16k", - "litellm_params": { - "model": "gpt-3.5-turbo-16k", - }, - "tpm": 240000, - "rpm": 1800, - }, - ] - router = Router( - model_list=model_list, - fallbacks=[{"gpt-5-mini": ["gpt-3.5-turbo-16k"]}], - retry_policy=litellm.router.RetryPolicy( - AuthenticationErrorRetries=0, - ), - ) # type: ignore - response = await router.acompletion( - model="gpt-5-mini", - messages=[{"role": "user", "content": "Hi 👋 - i'm openai"}], - ) - await asyncio.sleep(2) - print(f"customHandler_fallbacks.states: {customHandler_fallbacks.states}") - assert len(customHandler_fallbacks.errors) == 0 - assert ( - len(customHandler_fallbacks.states) == 6 - ) # pre, post, failure, pre, post, success - litellm.callbacks = [] - except Exception as e: - print(f"Assertion Error: {traceback.format_exc()}") - pytest.fail(f"An exception occurred - {str(e)}") - - -# asyncio.run(test_async_chat_azure_with_fallbacks()) - - -# CACHING -## Test Azure - completion, embedding -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=1) -async def test_async_completion_azure_caching(): - customHandler_caching = CompletionCustomHandler() - litellm.cache = Cache( - type="redis", - host=os.environ["REDIS_HOST"], - port=os.environ["REDIS_PORT"], - password=os.environ["REDIS_PASSWORD"], - ) - litellm.callbacks = [customHandler_caching] - unique_time = time.time() - model_list = [ - { - "model_name": "gpt-4.1-nano", # openai model name - "litellm_params": { # params for litellm completion/embedding call - "model": "azure/gpt-4.1-mini", - "api_key": os.getenv("AZURE_AI_API_KEY"), - "api_version": os.getenv("AZURE_API_VERSION"), - "api_base": os.getenv("AZURE_AI_API_BASE"), - }, - "tpm": 240000, - "rpm": 1800, - }, - { - "model_name": "gpt-3.5-turbo-16k", - "litellm_params": { - "model": "gpt-3.5-turbo-16k", - }, - "tpm": 240000, - "rpm": 1800, - }, - ] - router = Router(model_list=model_list) # type: ignore - response1 = await router.acompletion( - model="gpt-4.1-nano", - messages=[ - {"role": "user", "content": f"Hi 👋 - i'm async azure {unique_time}"} - ], - caching=True, - ) - await asyncio.sleep(1) - print(f"customHandler_caching.states pre-cache hit: {customHandler_caching.states}") - response2 = await router.acompletion( - model="gpt-4.1-nano", - messages=[ - {"role": "user", "content": f"Hi 👋 - i'm async azure {unique_time}"} - ], - caching=True, - ) - await asyncio.sleep(1) # success callbacks are done in parallel - print( - f"customHandler_caching.states post-cache hit: {customHandler_caching.states}" - ) - assert len(customHandler_caching.errors) == 0 - assert len(customHandler_caching.states) == 4 # pre, post, success, success - - -@pytest.mark.asyncio -async def test_async_completion_azure_caching_streaming(): - import uuid - - litellm.set_verbose = True - customHandler_caching = CompletionCustomHandler() - litellm.cache = Cache( - type="redis", - host=os.environ["REDIS_HOST"], - port=os.environ["REDIS_PORT"], - password=os.environ["REDIS_PASSWORD"], - ) - litellm.callbacks = [customHandler_caching] - unique_time = uuid.uuid4() - - # Use Router instead of direct litellm.acompletion to get router-specific metadata - model_list = [ - { - "model_name": "gpt-4.1-nano", - "litellm_params": { - "model": "azure/gpt-4.1-mini", - "api_key": os.getenv("AZURE_AI_API_KEY"), - "api_version": os.getenv("AZURE_API_VERSION"), - "api_base": os.getenv("AZURE_AI_API_BASE"), - }, - "tpm": 240000, - "rpm": 1800, - }, - ] - router = Router(model_list=model_list) - - response1 = await router.acompletion( - model="gpt-4.1-nano", - messages=[ - {"role": "user", "content": f"Hi 👋 - i'm async azure {unique_time}"} - ], - caching=True, - stream=True, - ) - async for chunk in response1: - print(f"chunk in response1: {chunk}") - await asyncio.sleep(1) - initial_customhandler_caching_states = len(customHandler_caching.states) - print(f"customHandler_caching.states pre-cache hit: {customHandler_caching.states}") - response2 = await router.acompletion( - model="gpt-4.1-nano", - messages=[ - {"role": "user", "content": f"Hi 👋 - i'm async azure {unique_time}"} - ], - caching=True, - stream=True, - ) - async for chunk in response2: - print(f"chunk in response2: {chunk}") - await asyncio.sleep(1) # success callbacks are done in parallel - print( - f"customHandler_caching.states post-cache hit: {customHandler_caching.states}" - ) - assert len(customHandler_caching.errors) == 0 - assert ( - len(customHandler_caching.states) > initial_customhandler_caching_states - ) # pre, post, streaming .., success, success - - -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=2) -async def test_async_embedding_azure_caching(): - print("Testing custom callback input - Azure Caching") - customHandler_caching = CompletionCustomHandler() - litellm.cache = Cache( - type="redis", - host=os.environ["REDIS_HOST"], - port=os.environ["REDIS_PORT"], - password=os.environ["REDIS_PASSWORD"], - ) - router = Router( - model_list=[ - { - "model_name": "text-embedding-3-small", - "litellm_params": { - "model": "openai/text-embedding-3-small", - }, - } - ] - ) - litellm.callbacks = [customHandler_caching] - unique_time = time.time() - response1 = await router.aembedding( - model="text-embedding-3-small", - input=[f"good morning from litellm1 {unique_time}"], - caching=True, - ) - await asyncio.sleep(1) # set cache is async for aembedding() - response2 = await router.aembedding( - model="text-embedding-3-small", - input=[f"good morning from litellm1 {unique_time}"], - caching=True, - ) - await asyncio.sleep(1) # success callbacks are done in parallel - print(customHandler_caching.states) - print(customHandler_caching.errors) - assert len(customHandler_caching.errors) == 0 - assert len(customHandler_caching.states) == 4 # pre, post, success, success diff --git a/tests/logging_callback_tests/test_datadog.py b/tests/logging_callback_tests/test_datadog.py deleted file mode 100644 index 35b46a58fd4..00000000000 --- a/tests/logging_callback_tests/test_datadog.py +++ /dev/null @@ -1,226 +0,0 @@ -import asyncio -import gzip -import io -import json -import logging -import os -from datetime import datetime as datetime_class -from unittest.mock import AsyncMock - -import pytest - -import litellm -from litellm._logging import verbose_logger -from litellm.integrations.datadog.datadog import * -from litellm.types.utils import ( - StandardLoggingHiddenParams, - StandardLoggingMetadata, - StandardLoggingModelInformation, - StandardLoggingPayload, -) - -verbose_logger.setLevel(logging.DEBUG) - - -def create_standard_logging_payload() -> StandardLoggingPayload: - return StandardLoggingPayload( - id="test_id", - call_type="completion", - response_cost=0.1, - response_cost_failure_debug_info=None, - status="success", - total_tokens=30, - prompt_tokens=20, - completion_tokens=10, - startTime=1234567890.0, - endTime=1234567891.0, - completionStartTime=1234567890.5, - model_map_information=StandardLoggingModelInformation( - model_map_key="gpt-4.1-mini", model_map_value=None - ), - model="gpt-4.1-mini", - model_id="model-123", - model_group="openai-gpt", - api_base="https://api.openai.com", - metadata=StandardLoggingMetadata( - user_api_key_hash="test_hash", - user_api_key_org_id=None, - user_api_key_alias="test_alias", - user_api_key_team_id="test_team", - user_api_key_user_id="test_user", - user_api_key_team_alias="test_team_alias", - spend_logs_metadata=None, - requester_ip_address="127.0.0.1", - requester_metadata=None, - ), - cache_hit=False, - cache_key=None, - saved_cache_cost=0.0, - request_tags=[], - end_user=None, - requester_ip_address="127.0.0.1", - messages=[{"role": "user", "content": "Hello, world!"}], - response={"choices": [{"message": {"content": "Hi there!"}}]}, - error_str=None, - model_parameters={"stream": True}, - hidden_params=StandardLoggingHiddenParams( - model_id="model-123", - cache_key=None, - api_base="https://api.openai.com", - response_cost="0.1", - additional_headers=None, - ), - ) - - - - - - -@pytest.mark.asyncio -async def test_create_datadog_logging_payload(): - """Test creating a DataDog logging payload from a standard logging object""" - dd_logger = DataDogLogger() - standard_payload = create_standard_logging_payload() - - # Create mock kwargs with the standard logging object - kwargs = {"standard_logging_object": standard_payload} - - # Test payload creation - dd_payload = dd_logger.create_datadog_logging_payload( - kwargs=kwargs, - response_obj=None, - start_time=datetime_class.now(), - end_time=datetime_class.now(), - ) - - # Verify payload structure - assert dd_payload["ddsource"] == os.getenv("DD_SOURCE", "litellm") - assert dd_payload["service"] == "litellm-server" - assert dd_payload["status"] == DataDogStatus.INFO - - # verify the message field == standard_payload - dict_payload = json.loads(dd_payload["message"]) - assert dict_payload == standard_payload - - -@pytest.mark.asyncio -async def test_datadog_failure_logging(): - """Test logging a failure event to DataDog""" - dd_logger = DataDogLogger() - standard_payload = create_standard_logging_payload() - standard_payload["status"] = "failure" # Set status to failure - standard_payload["error_str"] = "Test error" - - kwargs = {"standard_logging_object": standard_payload} - - dd_payload = dd_logger.create_datadog_logging_payload( - kwargs=kwargs, - response_obj=None, - start_time=datetime_class.now(), - end_time=datetime_class.now(), - ) - - assert ( - dd_payload["status"] == DataDogStatus.ERROR - ) # Verify failure maps to warning status - - # verify the message field == standard_payload - dict_payload = json.loads(dd_payload["message"]) - assert dict_payload == standard_payload - - # verify error_str is in the message field - assert "error_str" in dict_payload - assert dict_payload["error_str"] == "Test error" - - - - - - - - - - - - -@pytest.mark.asyncio -async def test_datadog_log_redis_failures(): - """ - Test that poorly configured Redis is logged as Warning on DataDog - """ - try: - from litellm.caching.caching import Cache - from litellm.integrations.datadog.datadog import DataDogLogger - - litellm.cache = Cache( - type="redis", host="badhost", port="6379", password="badpassword" - ) - - os.environ["DD_SITE"] = "https://fake.datadoghq.com" - os.environ["DD_API_KEY"] = "anything" - dd_logger = DataDogLogger() - - litellm.callbacks = [dd_logger] - litellm.service_callback = ["datadog"] - - litellm.set_verbose = True - - # Create a mock for the async_client's post method - mock_post = AsyncMock() - mock_post.return_value.status_code = 202 - mock_post.return_value.text = "Accepted" - dd_logger.async_client.post = mock_post - - # Make the completion call - for _ in range(3): - response = await litellm.acompletion( - model="gpt-4.1-mini", - messages=[{"role": "user", "content": "what llm are u"}], - max_tokens=10, - temperature=0.2, - mock_response="Accepted", - ) - print(response) - - # Wait for 5 seconds - await asyncio.sleep(6) - - # Assert that the mock was called - assert mock_post.called, "HTTP request was not made" - - # Get the arguments of the last call - args, kwargs = mock_post.call_args - print("CAll args and kwargs", args, kwargs) - - # For example, checking if the URL is correct - assert kwargs["url"].endswith("/api/v2/logs"), "Incorrect DataDog endpoint" - - body = kwargs["data"] - - # use gzip to unzip the body - with gzip.open(io.BytesIO(body), "rb") as f: - body = f.read().decode("utf-8") - print(body) - - # body is string parse it to dict - body = json.loads(body) - print(body) - - failure_events = [log for log in body if log["status"] == "warning"] - assert len(failure_events) > 0, "No failure events logged" - - print("ALL FAILURE/WARN EVENTS", failure_events) - - for event in failure_events: - message = json.loads(event["message"]) - assert ( - event["status"] == "warning" - ), f"Event status is not 'warning': {event['status']}" - assert ( - message["service"] == "redis" - ), f"Service is not 'redis': {message['service']}" - assert "error" in message, "No 'error' field in the message" - assert message["error"], "Error field is empty" - except Exception as e: - pytest.fail(f"Test failed with exception: {str(e)}") diff --git a/tests/logging_callback_tests/test_log_db_redis_services.py b/tests/logging_callback_tests/test_log_db_redis_services.py index ba7b333e097..aad404a5525 100644 --- a/tests/logging_callback_tests/test_log_db_redis_services.py +++ b/tests/logging_callback_tests/test_log_db_redis_services.py @@ -1,11 +1,9 @@ import io -import asyncio import gzip import json import logging -import time from unittest.mock import AsyncMock, patch import pytest @@ -13,215 +11,7 @@ import pytest import litellm from litellm import completion from litellm._logging import verbose_logger -from litellm.proxy.utils import log_db_metrics, ServiceTypes -from litellm.proxy.db.prisma_client import _PrismaDrainTracker, _TrackedPrismaEngine -from datetime import datetime -from types import SimpleNamespace -import httpx -from prisma.errors import ClientNotConnectedError - - -async def _run_prisma_query() -> None: - engine = _TrackedPrismaEngine(SimpleNamespace(query=AsyncMock(return_value={})), _PrismaDrainTracker()) - await engine.query("{}", tx_id=None) - - -# Test async function to decorate -@log_db_metrics -async def sample_db_function(*args, **kwargs): - await _run_prisma_query() - return "success" - - -@log_db_metrics -async def sample_proxy_function(*args, **kwargs): - return "success" - - -@pytest.mark.asyncio -async def test_log_db_metrics_success(): - # Mock the proxy_logging_obj - with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: - # Setup mock - mock_proxy_logging.service_logging_obj.async_service_success_hook = AsyncMock() - - # Call the decorated function - result = await sample_db_function(parent_otel_span="test_span") - - # Assertions - assert result == "success" - mock_proxy_logging.service_logging_obj.async_service_success_hook.assert_called_once() - call_args = ( - mock_proxy_logging.service_logging_obj.async_service_success_hook.call_args[ - 1 - ] - ) - - assert call_args["service"] == ServiceTypes.DB - assert call_args["call_type"] == "sample_db_function" - assert call_args["parent_otel_span"] == "test_span" - assert isinstance(call_args["duration"], float) - assert isinstance(call_args["start_time"], datetime) - assert isinstance(call_args["end_time"], datetime) - assert call_args["event_metadata"] is None - - -@pytest.mark.asyncio -async def test_log_db_metrics_event_metadata_is_safe(): - """event_metadata must surface only the table name, never the raw - kwargs/args which carry live clients (Prisma, OTel spans) and secrets. - - Regression guard for #28909: a previous version dumped function_kwargs and - function_args onto the span. - """ - with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: - mock_proxy_logging.service_logging_obj.async_service_success_hook = AsyncMock() - - @log_db_metrics - async def db_call(**kwargs): - await _run_prisma_query() - return "success" - - await db_call( - parent_otel_span="test_span", - table_name="LiteLLM_SpendLogs", - token="sk-secret-should-not-leak", - prisma_client=object(), - ) - await asyncio.sleep(0) - - call_args = ( - mock_proxy_logging.service_logging_obj.async_service_success_hook.call_args[ - 1 - ] - ) - assert call_args["event_metadata"] == {"table_name": "LiteLLM_SpendLogs"} - - -@pytest.mark.asyncio -async def test_log_db_metrics_duration(): - # Mock the proxy_logging_obj - with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: - # Setup mock - mock_proxy_logging.service_logging_obj.async_service_success_hook = AsyncMock() - - # Add a delay to the function to test duration - @log_db_metrics - async def delayed_function(**kwargs): - await _run_prisma_query() - await asyncio.sleep(1) # 1 second delay - return "success" - - # Call the decorated function - start = time.time() - result = await delayed_function(parent_otel_span="test_span") - end = time.time() - - # Get the actual duration - actual_duration = end - start - - # Get the logged duration from the mock call - call_args = ( - mock_proxy_logging.service_logging_obj.async_service_success_hook.call_args[ - 1 - ] - ) - logged_duration = call_args["duration"] - - # Assert the logged duration is approximately equal to actual duration (within 0.1 seconds) - assert abs(logged_duration - actual_duration) < 0.1 - assert result == "success" - - -@pytest.mark.asyncio -async def test_log_db_metrics_failure(): - """ - should log a failure if a prisma error is raised - """ - # Mock the proxy_logging_obj - from prisma.errors import ClientNotConnectedError - - with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: - # Setup mock - mock_proxy_logging.service_logging_obj.async_service_failure_hook = AsyncMock() - - # Create a failing function - @log_db_metrics - async def failing_function(**kwargs): - raise ClientNotConnectedError() - - # Call the decorated function and expect it to raise - with pytest.raises(ClientNotConnectedError) as exc_info: - await failing_function(parent_otel_span="test_span") - - # Assertions - assert "Client is not connected to the query engine" in str(exc_info.value) - mock_proxy_logging.service_logging_obj.async_service_failure_hook.assert_called_once() - call_args = ( - mock_proxy_logging.service_logging_obj.async_service_failure_hook.call_args[ - 1 - ] - ) - - assert call_args["service"] == ServiceTypes.DB - assert call_args["call_type"] == "failing_function" - assert call_args["parent_otel_span"] == "test_span" - assert isinstance(call_args["duration"], float) - assert isinstance(call_args["start_time"], datetime) - assert isinstance(call_args["end_time"], datetime) - assert isinstance(call_args["error"], ClientNotConnectedError) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "exception,should_log", - [ - (ValueError("Generic error"), False), - (KeyError("Missing key"), False), - (TypeError("Type error"), False), - (httpx.ConnectError("Failed to connect"), True), - (httpx.TimeoutException("Request timed out"), True), - (ClientNotConnectedError(), True), # Prisma error - ], -) -async def test_log_db_metrics_failure_error_types(exception, should_log): - """ - Why Test? - Users were seeing that non-DB errors were being logged as DB Service Failures - Example a failure to read a value from cache was being logged as a DB Service Failure - - - Parameterized test to verify: - - DB-related errors (Prisma, httpx) are logged as service failures - - Non-DB errors (ValueError, KeyError, etc.) are not logged - """ - with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: - mock_proxy_logging.service_logging_obj.async_service_failure_hook = AsyncMock() - - @log_db_metrics - async def failing_function(**kwargs): - raise exception - - # Call the function and expect it to raise the exception - with pytest.raises(type(exception)): - await failing_function(parent_otel_span="test_span") - - if should_log: - # Assert failure was logged for DB-related errors - mock_proxy_logging.service_logging_obj.async_service_failure_hook.assert_called_once() - call_args = mock_proxy_logging.service_logging_obj.async_service_failure_hook.call_args[ - 1 - ] - assert call_args["service"] == ServiceTypes.DB - assert call_args["call_type"] == "failing_function" - assert call_args["parent_otel_span"] == "test_span" - assert isinstance(call_args["duration"], float) - assert isinstance(call_args["start_time"], datetime) - assert isinstance(call_args["end_time"], datetime) - assert isinstance(call_args["error"], type(exception)) - else: - # Assert failure was NOT logged for non-DB errors - mock_proxy_logging.service_logging_obj.async_service_failure_hook.assert_not_called() +from litellm.proxy.utils import ServiceTypes @pytest.mark.asyncio diff --git a/tests/logging_callback_tests/test_moderations_api_logging.py b/tests/logging_callback_tests/test_moderations_api_logging.py deleted file mode 100644 index a2a356d3665..00000000000 --- a/tests/logging_callback_tests/test_moderations_api_logging.py +++ /dev/null @@ -1,100 +0,0 @@ -import traceback -from litellm._uuid import uuid -import pytest -from dotenv import load_dotenv -from fastapi import Request -from fastapi.routing import APIRoute - -load_dotenv() -import io -import time -import json - -import litellm -from litellm.router import Router -import asyncio -from typing import Optional -from litellm.types.utils import StandardLoggingPayload, Usage, ModelInfoBase -from litellm.integrations.custom_logger import CustomLogger - - -class TestCustomLogger(CustomLogger): - def __init__(self): - self.recorded_usage: Optional[Usage] = None - self.standard_logging_payload: Optional[StandardLoggingPayload] = None - - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - standard_logging_payload = kwargs.get("standard_logging_object") - self.standard_logging_payload = standard_logging_payload - print( - "standard_logging_payload", - json.dumps(standard_logging_payload, indent=4, default=str), - ) - - pass - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "model", [None, "omni-moderation-latest", "router-internal-moderation-model"] -) -async def test_moderations_api_logging(model): - """ - When moderations API is called, it should log the event on standard_logging_payload - """ - custom_logger = TestCustomLogger() - litellm.logging_callback_manager.add_litellm_callback(custom_logger) - - MODEL_GROUP = "internal-moderation-model" - router = Router( - model_list=[ - { - "model_name": MODEL_GROUP, - "litellm_params": { - "model": "openai/omni-moderation-latest", - }, - } - ] - ) - - input_content = "Hello, how are you?" - if model == "router-internal-moderation-model": - response = await router.amoderation( - input=input_content, - model=MODEL_GROUP, - ) - else: - response = await litellm.amoderation( - input=input_content, - model=model, - ) - - print("response", json.dumps(response, indent=4, default=str)) - - await asyncio.sleep(2) - - assert custom_logger.standard_logging_payload is not None - - # validate the standard_logging_payload - standard_logging_payload: StandardLoggingPayload = ( - custom_logger.standard_logging_payload - ) - assert ( - standard_logging_payload["call_type"] - == litellm.utils.CallTypes.amoderation.value - ) - assert standard_logging_payload["status"] == "success" - assert ( - standard_logging_payload["custom_llm_provider"] - == litellm.LlmProviders.OPENAI.value - ) - - # assert the logged input == input - assert standard_logging_payload["messages"][0]["content"] == input_content - - # assert the logged response == response user received client side - assert dict(standard_logging_payload["response"]) == response.model_dump() - - # if router used, validate model_group is logged as expected - if model == "router-internal-moderation-model": - assert standard_logging_payload["model_group"] == MODEL_GROUP diff --git a/tests/logging_callback_tests/test_otel_logging.py b/tests/logging_callback_tests/test_otel_logging.py deleted file mode 100644 index 6274916ddb2..00000000000 --- a/tests/logging_callback_tests/test_otel_logging.py +++ /dev/null @@ -1,133 +0,0 @@ -import pytest -import litellm -import asyncio -import logging -from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter -from litellm._logging import verbose_logger -from litellm.integrations.opentelemetry import ( - OpenTelemetry, - OpenTelemetryConfig, -) - -verbose_logger.setLevel(logging.DEBUG) - -EXPECTED_SPAN_NAMES = ["litellm_request", "raw_gen_ai_request"] -exporter = InMemorySpanExporter() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("streaming", [True, False]) -async def test_async_otel_callback(streaming): - litellm.set_verbose = True - - # Clear exporter at the start to ensure clean state - exporter.clear() - - litellm.callbacks = [OpenTelemetry(config=OpenTelemetryConfig(exporter=exporter))] - - response = await litellm.acompletion( - model="gpt-4.1-mini", - messages=[{"role": "user", "content": "hi"}], - temperature=0.1, - user="OTEL_USER", - stream=streaming, - ) - - if streaming is True: - async for chunk in response: - print("chunk", chunk) - - await asyncio.sleep(4) - spans = exporter.get_finished_spans() - print("spans", spans) - assert len(spans) == 2 - - _span_names = [span.name for span in spans] - print("recorded span names", _span_names) - assert set(_span_names) == set(EXPECTED_SPAN_NAMES) - - # print the value of a span - for span in spans: - print("span name", span.name) - print("span attributes", span.attributes) - - if span.name == "litellm_request": - validate_litellm_request(span) - # Additional specific checks - assert span._attributes["gen_ai.request.model"] == "gpt-4.1-mini" - assert span._attributes["gen_ai.system"] == "openai" - assert span._attributes["gen_ai.request.temperature"] == 0.1 - assert span._attributes["llm.is_streaming"] == str(streaming) - assert span._attributes["llm.user"] == "OTEL_USER" - elif span.name == "raw_gen_ai_request": - if streaming is True: - validate_raw_gen_ai_request_openai_streaming(span) - else: - validate_raw_gen_ai_request_openai_non_streaming(span) - - # clear in memory exporter - exporter.clear() - - -def validate_litellm_request(span): - expected_attributes = [ - "gen_ai.request.model", - "gen_ai.system", - "gen_ai.request.temperature", - "llm.is_streaming", - "llm.user", - "gen_ai.response.id", - "gen_ai.response.model", - "gen_ai.usage.total_tokens", - "gen_ai.usage.output_tokens", - "gen_ai.usage.input_tokens", - ] - - # get the str of all the span attributes - print("span attributes", span._attributes) - - for attr in expected_attributes: - value = span._attributes[attr] - print("value", value) - assert value is not None, f"Attribute {attr} has None value" - - -def validate_raw_gen_ai_request_openai_non_streaming(span): - expected_attributes = [ - "llm.openai.messages", - "llm.openai.temperature", - "llm.openai.user", - "llm.openai.extra_body", - "llm.openai.id", - "llm.openai.choices", - "llm.openai.created", - "llm.openai.model", - "llm.openai.object", - "llm.openai.service_tier", - "llm.openai.system_fingerprint", - "llm.openai.usage", - ] - - print("span attributes", span._attributes) - for attr in span._attributes: - print(attr) - - for attr in expected_attributes: - assert span._attributes[attr] is not None, f"Attribute {attr} has None" - - -def validate_raw_gen_ai_request_openai_streaming(span): - expected_attributes = [ - "llm.openai.messages", - "llm.openai.temperature", - "llm.openai.user", - "llm.openai.extra_body", - "llm.openai.model", - ] - - print("span attributes", span._attributes) - for attr in span._attributes: - print(attr) - - for attr in expected_attributes: - assert span._attributes[attr] is not None, f"Attribute {attr} has None" diff --git a/tests/logging_callback_tests/test_token_counting.py b/tests/logging_callback_tests/test_token_counting.py deleted file mode 100644 index 513d2242fdf..00000000000 --- a/tests/logging_callback_tests/test_token_counting.py +++ /dev/null @@ -1,157 +0,0 @@ -import traceback -from litellm._uuid import uuid -import pytest -from dotenv import load_dotenv -from fastapi import Request -from fastapi.routing import APIRoute - -load_dotenv() -import io -import time -import json - -# this file is to test litellm/proxy - -import litellm -import asyncio -from typing import Optional -from litellm.types.utils import StandardLoggingPayload, Usage, ModelInfoBase -from litellm.integrations.custom_logger import CustomLogger - - -class TestCustomLogger(CustomLogger): - def __init__(self): - self.recorded_usage: Optional[Usage] = None - self.standard_logging_payload: Optional[StandardLoggingPayload] = None - - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - standard_logging_payload = kwargs.get("standard_logging_object") - self.standard_logging_payload = standard_logging_payload - print( - "standard_logging_payload", - json.dumps(standard_logging_payload, indent=4, default=str), - ) - - self.recorded_usage = Usage( - prompt_tokens=standard_logging_payload.get("prompt_tokens"), - completion_tokens=standard_logging_payload.get("completion_tokens"), - total_tokens=standard_logging_payload.get("total_tokens"), - ) - pass - - -@pytest.mark.asyncio -async def test_stream_token_counting_gpt_4o(): - """ - When stream_options={"include_usage": True} logging callback tracks Usage == Usage from llm API - """ - custom_logger = TestCustomLogger() - litellm.logging_callback_manager.add_litellm_callback(custom_logger) - - response = await litellm.acompletion( - model="gpt-5.5", - messages=[{"role": "user", "content": "Hello, how are you?" * 100}], - stream=True, - stream_options={"include_usage": True}, - ) - - actual_usage = None - async for chunk in response: - if "usage" in chunk: - actual_usage = chunk["usage"] - print("chunk.usage", json.dumps(chunk["usage"], indent=4, default=str)) - pass - - await asyncio.sleep(2) - - print("\n\n\n\n\n") - print( - "recorded_usage", - json.dumps(custom_logger.recorded_usage, indent=4, default=str), - ) - print("\n\n\n\n\n") - - assert actual_usage.prompt_tokens == custom_logger.recorded_usage.prompt_tokens - assert ( - actual_usage.completion_tokens == custom_logger.recorded_usage.completion_tokens - ) - assert actual_usage.total_tokens == custom_logger.recorded_usage.total_tokens - - -@pytest.mark.asyncio -async def test_stream_token_counting_without_include_usage(): - """ - When stream_options={"include_usage": True} is not passed, the usage tracked == usage from llm api chunk - - by default, litellm passes `include_usage=True` for OpenAI API - """ - custom_logger = TestCustomLogger() - litellm.logging_callback_manager.add_litellm_callback(custom_logger) - - response = await litellm.acompletion( - model="gpt-5.5", - messages=[{"role": "user", "content": "Hello, how are you?" * 100}], - stream=True, - ) - - actual_usage = None - async for chunk in response: - if "usage" in chunk: - actual_usage = chunk["usage"] - print("chunk.usage", json.dumps(chunk["usage"], indent=4, default=str)) - pass - - await asyncio.sleep(2) - - print("\n\n\n\n\n") - print( - "recorded_usage", - json.dumps(custom_logger.recorded_usage, indent=4, default=str), - ) - print("\n\n\n\n\n") - - assert actual_usage.prompt_tokens == custom_logger.recorded_usage.prompt_tokens - assert ( - actual_usage.completion_tokens == custom_logger.recorded_usage.completion_tokens - ) - assert actual_usage.total_tokens == custom_logger.recorded_usage.total_tokens - - -@pytest.mark.asyncio -async def test_stream_token_counting_with_redaction(): - """ - When litellm.turn_off_message_logging=True is used, the usage tracked == usage from llm api chunk - """ - litellm.turn_off_message_logging = True - custom_logger = TestCustomLogger() - litellm.logging_callback_manager.add_litellm_callback(custom_logger) - - response = await litellm.acompletion( - model="gpt-5.5", - messages=[{"role": "user", "content": "Hello, how are you?" * 100}], - stream=True, - ) - - actual_usage = None - async for chunk in response: - if "usage" in chunk: - actual_usage = chunk["usage"] - print("chunk.usage", json.dumps(chunk["usage"], indent=4, default=str)) - pass - - await asyncio.sleep(2) - - print("\n\n\n\n\n") - print( - "recorded_usage", - json.dumps(custom_logger.recorded_usage, indent=4, default=str), - ) - print("\n\n\n\n\n") - - assert actual_usage.prompt_tokens == custom_logger.recorded_usage.prompt_tokens - assert ( - actual_usage.completion_tokens == custom_logger.recorded_usage.completion_tokens - ) - assert actual_usage.total_tokens == custom_logger.recorded_usage.total_tokens - - diff --git a/tests/ocr_tests/test_ocr_matrix.py b/tests/ocr_tests/test_ocr_matrix.py index 13cffbbc9a1..3f8c9b402fc 100644 --- a/tests/ocr_tests/test_ocr_matrix.py +++ b/tests/ocr_tests/test_ocr_matrix.py @@ -24,7 +24,6 @@ from typing import Final, Literal import pytest import litellm -from litellm import Router from litellm.integrations.custom_logger import CustomLogger from litellm.llms.base_llm.ocr.transformation import OCRResponse @@ -301,17 +300,3 @@ async def test_ocr(case: Case, monkeypatch: pytest.MonkeyPatch, logger: Recordin _assert_logged(await logger.wait_for_call(), response, case.provider.model, response.model, case.call) -async def test_router_aocr(monkeypatch: pytest.MonkeyPatch, logger: RecordingLogger) -> None: - case: Final = Case(MISTRAL, MISTRAL_KEY, "explicit", PDF_BY_URL, "async") - router: Final = Router( - model_list=[ - { - "model_name": "ocr-alias", - "litellm_params": {"model": MISTRAL.model, **case.bind_credentials(monkeypatch)}, - } - ] - ) - response: Final = await router.aocr(model="ocr-alias", document=PDF_BY_URL.build()) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # Router.aocr is untyped - assert isinstance(response, OCRResponse) - _assert_ocr_response(response, MISTRAL.model, PDF_TEXT) - _assert_logged(await logger.wait_for_call(), response, MISTRAL.model, MISTRAL.model, case.call) diff --git a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py b/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py deleted file mode 100644 index 260656600ec..00000000000 --- a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py +++ /dev/null @@ -1,115 +0,0 @@ -import os -import time -from collections.abc import Iterator -from typing import Final - -import httpx -import pytest -from openai import APIStatusError, BadRequestError, NotFoundError, OpenAI, Stream -from openai.types.responses import ResponseStreamEvent - -BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS: Final = 90 - - -def generate_key(): - """Generate a key for testing""" - url = "http://0.0.0.0:4000/key/generate" - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - } - data = {} - - response = httpx.post(url, headers=headers, json=data) - if response.status_code != 200: - raise Exception(f"Key generation failed with status: {response.status_code}") - return response.json()["key"] - - -def get_test_client(): - """Create OpenAI client with generated key""" - key = generate_key() - return OpenAI(api_key=key, base_url="http://0.0.0.0:4000") - - -def validate_response(response): - """ - Validate basic response structure from OpenAI responses API - """ - assert response is not None - assert hasattr(response, "choices") - assert len(response.choices) > 0 - assert hasattr(response.choices[0], "message") - assert hasattr(response.choices[0].message, "content") - assert isinstance(response.choices[0].message.content, str) - assert hasattr(response, "id") - assert isinstance(response.id, str) - assert hasattr(response, "model") - assert isinstance(response.model, str) - assert hasattr(response, "created") - assert isinstance(response.created, int) - assert hasattr(response, "usage") - assert hasattr(response.usage, "prompt_tokens") - assert hasattr(response.usage, "completion_tokens") - assert hasattr(response.usage, "total_tokens") - - -def validate_stream_chunk(chunk): - """ - Validate streaming chunk structure from OpenAI responses API - """ - assert chunk is not None - assert hasattr(chunk, "choices") - assert len(chunk.choices) > 0 - assert hasattr(chunk.choices[0], "delta") - - # Some chunks might not have content in the delta - if ( - hasattr(chunk.choices[0].delta, "content") - and chunk.choices[0].delta.content is not None - ): - assert isinstance(chunk.choices[0].delta.content, str) - - assert hasattr(chunk, "id") - assert isinstance(chunk.id, str) - assert hasattr(chunk, "model") - assert isinstance(chunk.model, str) - assert hasattr(chunk, "created") - assert isinstance(chunk.created, int) - - -def test_model_not_found_error(): - client = get_test_client() - with pytest.raises(NotFoundError): - client.responses.create(model="non-existent-model", input="This should fail") - - -def test_bad_request_bad_param_error(): - client = get_test_client() - with pytest.raises(BadRequestError): - # Out-of-range temperature on a non-reasoning model, so drop_params forwards it - client.responses.create( - model="gpt-4.1", input="This should fail", temperature=2000 - ) - - -def admitted_response_id(chunk: ResponseStreamEvent) -> str | None: - response: Final = getattr(chunk, "response", None) - return None if response is None else response.id - - -def events_until_admission(stream: Stream[ResponseStreamEvent], started: float) -> Iterator[ResponseStreamEvent]: - for chunk in stream: - print("stream chunk=", chunk) - yield chunk - if admitted_response_id(chunk) is not None: - return - if time.monotonic() - started > BACKGROUND_STREAM_ADMISSION_DEADLINE_SECONDS: - return - - -def test_cancel_invalid_response_id(): - client = get_test_client() - with pytest.raises(APIStatusError): - # Try to cancel a non-existent response ID - client.responses.cancel("invalid_response_id_12345") diff --git a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py deleted file mode 100644 index f488095aa12..00000000000 --- a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py +++ /dev/null @@ -1,463 +0,0 @@ -# What this tests ? -## Tests /batches endpoints -import os -import pytest -import asyncio -import aiohttp, openai -from openai import OpenAI, AsyncOpenAI -from typing import Optional, List, Union -from test_openai_files_endpoints import upload_file, delete_file -import sys -import time - - -BASE_URL = "http://localhost:4000" # Replace with your actual base URL -API_KEY = os.environ["LITELLM_MASTER_KEY"] # Replace with your actual API key - - -client = OpenAI(base_url=BASE_URL, api_key=API_KEY) - - -def create_batch_oai_sdk(filepath: str, custom_llm_provider: str) -> str: - batch_input_file = client.files.create( - file=open(filepath, "rb"), - purpose="batch", - extra_headers={"custom-llm-provider": custom_llm_provider}, - ) - batch_input_file_id = batch_input_file.id - - print("waiting for file to be processed......") - time.sleep(5) - rq = client.batches.create( - input_file_id=batch_input_file_id, - endpoint="/v1/chat/completions", - completion_window="24h", - metadata={ - "description": filepath, - }, - extra_headers={"custom-llm-provider": custom_llm_provider}, - ) - - print(f"Batch submitted. ID: {rq.id}") - return rq.id - - -def await_batch_completion(batch_id: str, custom_llm_provider: str): - max_tries = 3 - tries = 0 - - while tries < max_tries: - batch = client.batches.retrieve( - batch_id, extra_headers={"custom-llm-provider": custom_llm_provider} - ) - if batch.status == "completed": - print(f"Batch {batch_id} completed.") - return batch.id - - tries += 1 - print(f"waiting for batch to complete... (attempt {tries}/{max_tries})") - time.sleep(10) - - print( - f"Reached maximum number of attempts ({max_tries}). Batch may still be processing." - ) - - -def write_content_to_file( - batch_id: str, output_path: str, custom_llm_provider: str -) -> str: - batch = client.batches.retrieve( - batch_id=batch_id, extra_headers={"custom-llm-provider": custom_llm_provider} - ) - content = client.files.content( - file_id=batch.output_file_id, - extra_headers={"custom-llm-provider": custom_llm_provider}, - ) - print("content from files.content", content.content) - content.write_to_file(output_path) - - -def read_jsonl(filepath: str): - import json - - results = [] - with open(filepath, "r") as f: - for line in f: - if line.strip(): - results.append(json.loads(line)) - - for item in results: - print(item) - custom_id = item["custom_id"] - print(custom_id) - - -def get_any_completed_batch_id_azure(): - print("AZURE getting any completed batch id") - list_of_batches = client.batches.list( - extra_headers={"custom-llm-provider": "azure"} - ) - print("list of batches", list_of_batches) - for batch in list_of_batches: - if batch.status == "completed": - return batch.id - return None - - -@pytest.mark.skip(reason="Local only test to verify if things work well") -def test_vertex_batches_endpoint(): - """ - Test VertexAI Batches Endpoint - """ - import os - - oai_client = OpenAI(api_key=API_KEY, base_url=BASE_URL) - file_name = "local_testing/vertex_batch_completions.jsonl" - _current_dir = os.path.dirname(os.path.abspath(__file__)) - file_path = os.path.join(_current_dir, file_name) - file_obj = oai_client.files.create( - file=open(file_path, "rb"), - purpose="batch", - extra_headers={"custom-llm-provider": "vertex_ai"}, - ) - print("Response from creating file=", file_obj) - - batch_input_file_id = file_obj.id - assert ( - batch_input_file_id is not None - ), f"Failed to create file, expected a non null file_id but got {batch_input_file_id}" - - create_batch_response = oai_client.batches.create( - completion_window="24h", - endpoint="/v1/chat/completions", - input_file_id=batch_input_file_id, - extra_headers={"custom-llm-provider": "vertex_ai"}, - metadata={"key1": "value1", "key2": "value2"}, - ) - print("response from create batch", create_batch_response) - pass - - -@pytest.mark.asyncio -async def test_batch_status_sync_from_provider_to_database(): - """ - Test that when batch status changes at the provider, - it gets synced to the ManagedObjectTable database. - - This tests the new refactored utility functions: - - get_batch_from_database() - - update_batch_in_database() - """ - from unittest.mock import MagicMock, AsyncMock - from litellm.proxy.openai_files_endpoints.common_utils import ( - get_batch_from_database, - update_batch_in_database, - ) - from litellm.types.utils import LiteLLMBatch - import json - - # Setup: Create mock objects - batch_id = "batch_test123" - unified_batch_id = "litellm_proxy:test_unified_batch" - - # Mock database batch object with "validating" status - mock_db_batch = MagicMock() - mock_db_batch.unified_object_id = batch_id - mock_db_batch.status = "validating" - mock_db_batch.file_object = json.dumps( - { - "id": batch_id, - "object": "batch", - "status": "validating", - "endpoint": "/v1/chat/completions", - "input_file_id": "file-test123", - "completion_window": "24h", - "created_at": 1234567890, - } - ) - - # Mock prisma client - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( - return_value=mock_db_batch - ) - mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - - # Mock managed_files_obj - mock_managed_files = MagicMock() - - # Mock logger - mock_logger = MagicMock() - mock_logger.debug = MagicMock() - mock_logger.info = MagicMock() - mock_logger.warning = MagicMock() - mock_logger.error = MagicMock() - - # Test 1: Retrieve batch from database (initial state) - db_batch_object, response_batch = await get_batch_from_database( - batch_id=batch_id, - unified_batch_id=unified_batch_id, - managed_files_obj=mock_managed_files, - prisma_client=mock_prisma_client, - verbose_proxy_logger=mock_logger, - ) - - # Verify database was queried - mock_prisma_client.db.litellm_managedobjecttable.find_first.assert_called_once_with( - where={"unified_object_id": batch_id} - ) - - # Verify batch was retrieved correctly - assert db_batch_object is not None - assert response_batch is not None - assert response_batch.id == batch_id - assert response_batch.status == "validating" - - # Test 2: Simulate provider returning updated status - updated_batch_response = LiteLLMBatch( - id=batch_id, - object="batch", - status="completed", # Status changed from "validating" to "completed" - endpoint="/v1/chat/completions", - input_file_id="file-test123", - completion_window="24h", - created_at=1234567890, - output_file_id="file-output123", - ) - - # Test 3: Update database with new status from provider - await update_batch_in_database( - batch_id=batch_id, - unified_batch_id=unified_batch_id, - response=updated_batch_response, - managed_files_obj=mock_managed_files, - prisma_client=mock_prisma_client, - verbose_proxy_logger=mock_logger, - db_batch_object=db_batch_object, - operation="retrieve", - ) - - # Verify database was updated - mock_prisma_client.db.litellm_managedobjecttable.update.assert_called_once() - update_call_args = mock_prisma_client.db.litellm_managedobjecttable.update.call_args - - # Verify the update call had correct parameters - assert update_call_args.kwargs["where"]["unified_object_id"] == batch_id - assert ( - update_call_args.kwargs["data"]["status"] == "complete" - ) # "completed" normalized to "complete" - assert "file_object" in update_call_args.kwargs["data"] - assert "updated_at" in update_call_args.kwargs["data"] - # batch_processed must be set to True when batch transitions to complete - assert update_call_args.kwargs["data"]["batch_processed"] is True - - # Verify logger was called with status change message - mock_logger.info.assert_called() - log_message = mock_logger.info.call_args[0][0] % mock_logger.info.call_args[0][1:] - assert "validating" in log_message - assert "completed" in log_message - - print("✅ Test passed: Batch status synced from provider to database") - - -@pytest.mark.asyncio -async def test_batch_cancel_updates_database(): - """ - Test that canceling a batch updates the database status. - """ - from unittest.mock import MagicMock, AsyncMock - from litellm.proxy.openai_files_endpoints.common_utils import ( - update_batch_in_database, - ) - from litellm.types.utils import LiteLLMBatch - - # Setup - batch_id = "batch_cancel_test" - unified_batch_id = "litellm_proxy:cancel_test" - - # Mock cancelled batch response from provider - cancelled_batch_response = LiteLLMBatch( - id=batch_id, - object="batch", - status="cancelled", - endpoint="/v1/chat/completions", - input_file_id="file-test123", - completion_window="24h", - created_at=1234567890, - cancelled_at=1234567999, - ) - - # Mock prisma client - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( - return_value=None - ) - mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - - # Mock managed_files_obj - mock_managed_files = MagicMock() - - # Mock logger - mock_logger = MagicMock() - mock_logger.info = MagicMock() - mock_logger.error = MagicMock() - - # Call update_batch_in_database for cancel operation - await update_batch_in_database( - batch_id=batch_id, - unified_batch_id=unified_batch_id, - response=cancelled_batch_response, - managed_files_obj=mock_managed_files, - prisma_client=mock_prisma_client, - verbose_proxy_logger=mock_logger, - operation="cancel", - ) - - # Verify database was updated - mock_prisma_client.db.litellm_managedobjecttable.update.assert_called_once() - update_call_args = mock_prisma_client.db.litellm_managedobjecttable.update.call_args - - # Verify the update call had correct parameters - assert update_call_args.kwargs["where"]["unified_object_id"] == batch_id - assert update_call_args.kwargs["data"]["status"] == "cancelled" - assert "file_object" in update_call_args.kwargs["data"] - - # Verify logger was called - mock_logger.info.assert_called() - log_message = mock_logger.info.call_args[0][0] % mock_logger.info.call_args[0][1:] - assert "cancel" in log_message.lower() - assert "cancelled" in log_message - - print("✅ Test passed: Batch cancel updates database") - - -@pytest.mark.asyncio -async def test_batch_terminal_state_skip_provider_call(): - """ - Test that when a batch is in a terminal state (completed, failed, cancelled, expired), - it returns immediately from database without calling the provider. - """ - from unittest.mock import MagicMock, AsyncMock - from litellm.proxy.openai_files_endpoints.common_utils import ( - get_batch_from_database, - ) - from litellm.types.utils import LiteLLMBatch - import json - - # Setup: Create mock objects for a completed batch - batch_id = "batch_completed_test" - unified_batch_id = "litellm_proxy:completed_test" - - # Mock database batch object with "completed" status - mock_db_batch = MagicMock() - mock_db_batch.unified_object_id = batch_id - mock_db_batch.status = "complete" - mock_db_batch.file_object = json.dumps( - { - "id": batch_id, - "object": "batch", - "status": "completed", - "endpoint": "/v1/chat/completions", - "input_file_id": "file-test123", - "output_file_id": "file-output123", - "completion_window": "24h", - "created_at": 1234567890, - "completed_at": 1234567999, - } - ) - - # Mock prisma client - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( - return_value=mock_db_batch - ) - - # Mock managed_files_obj - mock_managed_files = MagicMock() - - # Mock logger - mock_logger = MagicMock() - mock_logger.debug = MagicMock() - - # Retrieve batch from database - db_batch_object, response_batch = await get_batch_from_database( - batch_id=batch_id, - unified_batch_id=unified_batch_id, - managed_files_obj=mock_managed_files, - prisma_client=mock_prisma_client, - verbose_proxy_logger=mock_logger, - ) - - # Verify batch was retrieved - assert db_batch_object is not None - assert response_batch is not None - assert response_batch.status == "completed" - - # In the actual endpoint, when status is in terminal states, - # it should return immediately without calling the provider - # This test verifies the database retrieval works correctly - assert response_batch.status in ["completed", "failed", "cancelled", "expired"] - - print("✅ Test passed: Terminal state batch retrieved from database") - - -@pytest.mark.asyncio -async def test_batch_no_status_change_skip_update(): - """ - Test that when batch status hasn't changed, database update is skipped. - """ - from unittest.mock import MagicMock, AsyncMock - from litellm.proxy.openai_files_endpoints.common_utils import ( - update_batch_in_database, - ) - from litellm.types.utils import LiteLLMBatch - - # Setup - batch_id = "batch_no_change_test" - unified_batch_id = "litellm_proxy:no_change_test" - - # Mock database batch object with "validating" status - mock_db_batch = MagicMock() - mock_db_batch.status = "validating" - - # Mock batch response from provider with same status - batch_response = LiteLLMBatch( - id=batch_id, - object="batch", - status="validating", # Same status as in database - endpoint="/v1/chat/completions", - input_file_id="file-test123", - completion_window="24h", - created_at=1234567890, - ) - - # Mock prisma client - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() - - # Mock managed_files_obj - mock_managed_files = MagicMock() - - # Mock logger - mock_logger = MagicMock() - mock_logger.info = MagicMock() - - # Call update_batch_in_database - await update_batch_in_database( - batch_id=batch_id, - unified_batch_id=unified_batch_id, - response=batch_response, - managed_files_obj=mock_managed_files, - prisma_client=mock_prisma_client, - verbose_proxy_logger=mock_logger, - db_batch_object=mock_db_batch, - operation="retrieve", - ) - - # Verify database update was NOT called (status hasn't changed) - mock_prisma_client.db.litellm_managedobjecttable.update.assert_not_called() - - # Verify logger info was NOT called (no status change to log) - mock_logger.info.assert_not_called() - - print("✅ Test passed: Database update skipped when status unchanged") diff --git a/tests/openai_endpoints_tests/test_openai_files_endpoints.py b/tests/openai_endpoints_tests/test_openai_files_endpoints.py deleted file mode 100644 index 9398a0d1c53..00000000000 --- a/tests/openai_endpoints_tests/test_openai_files_endpoints.py +++ /dev/null @@ -1,113 +0,0 @@ -import os -# What this tests ? -## Tests /chat/completions by generating a key and then making a chat completions request -import pytest -import asyncio -import aiohttp, openai -from openai import OpenAI, AsyncOpenAI -from typing import Optional, List, Union - - -BASE_URL = "http://localhost:4000" # Replace with your actual base URL -API_KEY = os.environ["LITELLM_MASTER_KEY"] # Replace with your actual API key - - -@pytest.mark.asyncio -async def test_file_operations(): - openai_client = AsyncOpenAI(api_key=API_KEY, base_url=BASE_URL) - file_content = b'{"prompt": "Hello", "completion": "Hi"}' - uploaded_file = await openai_client.files.create( - purpose="fine-tune", - file=file_content, - ) - list_files = await openai_client.files.list() - print("list_files=", list_files) - - get_file = await openai_client.files.retrieve(file_id=uploaded_file.id) - print("get_file=", get_file) - - get_file_content = await openai_client.files.content(file_id=uploaded_file.id) - print("get_file_content=", get_file_content.content) - response = get_file_content.response - - assert get_file_content.content == file_content - assert response.status_code == 200 - assert response.headers.get("content-type") == "application/octet-stream" - assert response.headers.get("content-length") is not None - assert int(response.headers["content-length"]) == len(get_file_content.content) - assert response.headers.get("content-disposition") is not None - assert uploaded_file.filename in response.headers["content-disposition"] - assert response.headers.get("x-request-id") is not None - # try get_file_content.write_to_file - get_file_content.write_to_file("get_file_content.jsonl") - - delete_file = await openai_client.files.delete(file_id=uploaded_file.id) - print("delete_file=", delete_file) - - -async def upload_file(session, purpose="fine-tune"): - url = f"{BASE_URL}/v1/files" - headers = {"Authorization": f"Bearer {API_KEY}"} - data = aiohttp.FormData() - data.add_field("purpose", purpose) - data.add_field( - "file", b'{"prompt": "Hello", "completion": "Hi"}', filename="mydata.jsonl" - ) - - async with session.post(url, headers=headers, data=data) as response: - assert response.status == 200 - result = await response.json() - assert "id" in result - print(f"File upload successful. File ID: {result['id']}") - return result["id"] - - -async def list_files(session): - url = f"{BASE_URL}/v1/files" - headers = {"Authorization": f"Bearer {API_KEY}"} - - async with session.get(url, headers=headers) as response: - assert response.status == 200 - result = await response.json() - assert "data" in result - print("List files successful") - - -async def get_file(session, file_id): - url = f"{BASE_URL}/v1/files/{file_id}" - headers = {"Authorization": f"Bearer {API_KEY}"} - - async with session.get(url, headers=headers) as response: - assert response.status == 200 - result = await response.json() - assert result["id"] == file_id - assert result["object"] == "file" - assert "bytes" in result - assert "created_at" in result - assert "filename" in result - assert result["purpose"] == "fine-tune" - print(f"Get file successful for file ID: {file_id}") - - -async def get_file_content(session, file_id): - url = f"{BASE_URL}/v1/files/{file_id}/content" - headers = {"Authorization": f"Bearer {API_KEY}"} - - async with session.get(url, headers=headers) as response: - assert response.status == 200 - content = await response.text() - print("content from /files/{file_id}/content=", content) - assert content # Check if content is not empty - print(f"Get file content successful for file ID: {file_id}") - - -async def delete_file(session, file_id): - url = f"{BASE_URL}/v1/files/{file_id}" - headers = {"Authorization": f"Bearer {API_KEY}"} - - async with session.delete(url, headers=headers) as response: - assert response.status == 200 - result = await response.json() - assert "deleted" in result - assert result["id"] == file_id - print(f"Delete file successful for file ID: {file_id}") diff --git a/tests/otel_tests/test_e2e_budgeting.py b/tests/otel_tests/test_e2e_budgeting.py deleted file mode 100644 index f180403e115..00000000000 --- a/tests/otel_tests/test_e2e_budgeting.py +++ /dev/null @@ -1,557 +0,0 @@ -import os -import asyncio -import json -import secrets -import uuid -from typing import Any, Optional - -import aiohttp -import openai -import pytest -from httpx import AsyncClient - -PROXY_BASE = "http://0.0.0.0:4000" -MASTER_HEADERS = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} -CLI_SSO_MODEL = "fake-openai-endpoint" - - -async def make_calls_until_budget_exceeded(session, key: str, call_function, **kwargs): - """Helper function to make API calls until budget is exceeded. Verify that the budget is exceeded error is returned.""" - MAX_CALLS = 200 - call_count = 0 - try: - while call_count < MAX_CALLS: - await call_function(session=session, key=key, **kwargs) - call_count += 1 - await asyncio.sleep(0.1) # allow spend tracking to catch up - pytest.fail(f"Budget was not exceeded after {MAX_CALLS} calls") - except openai.APIStatusError as e: - print("vars: ", vars(e)) - print("e.body: ", e.body) - - error_dict = e.body - print("error_dict: ", error_dict) - - # Check error structure and values that should be consistent - assert ( - error_dict["code"] == "422" - ), f"Expected error code 422, got: {error_dict['code']}" - assert ( - error_dict["type"] == "budget_exceeded" - ), f"Expected error type budget_exceeded, got: {error_dict['type']}" - - # Check message contains required parts without checking specific values - message = error_dict["message"] - assert ( - "Budget has been exceeded!" in message - ), f"Expected message to start with 'Budget has been exceeded!', got: {message}" - assert ( - "Current cost:" in message - ), f"Expected message to contain 'Current cost:', got: {message}" - assert ( - "Max budget:" in message - ), f"Expected message to contain 'Max budget:', got: {message}" - - return call_count - - -async def generate_key( - session, - max_budget=None, -): - url = "http://0.0.0.0:4000/key/generate" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = { - "max_budget": max_budget, - } - async with session.post(url, headers=headers, json=data) as response: - return await response.json() - - -async def chat_completion(session, key: str, model: str): - """Make a chat completion request using OpenAI SDK""" - from openai import AsyncOpenAI - from litellm._uuid import uuid - - client = AsyncOpenAI( - api_key=key, base_url="http://0.0.0.0:4000/v1" # Point to our local proxy - ) - - response = await client.chat.completions.create( - model=model, - messages=[{"role": "user", "content": f"Say hello! {uuid.uuid4()}" * 100}], - ) - return response - - -@pytest.mark.asyncio -async def test_chat_completion_low_budget(): - """ - Test budget enforcement for chat completions: - 1. Create key with $0.01 budget - 2. Make chat completion calls until budget exceeded - 3. Verify budget exceeded error - """ - async with aiohttp.ClientSession() as session: - # Create key with $0.01 budget - key_gen = await generate_key(session=session, max_budget=0.0000000005) - print("response from key generation: ", key_gen) - key = key_gen["key"] - - # Make calls until budget exceeded - calls_made = await make_calls_until_budget_exceeded( - session=session, - key=key, - call_function=chat_completion, - model="fake-openai-endpoint", - ) - - assert ( - calls_made > 0 - ), "Should make at least one successful call before budget exceeded" - - -@pytest.mark.asyncio -async def test_chat_completion_zero_budget(): - """ - Test budget enforcement for chat completions: - 1. Create key with $0.01 budget - 2. Make chat completion calls until budget exceeded - 3. Verify budget exceeded error - """ - async with aiohttp.ClientSession() as session: - # Create key with $0.01 budget - key_gen = await generate_key(session=session, max_budget=0.000000000) - print("response from key generation: ", key_gen) - key = key_gen["key"] - - # Make calls until budget exceeded - calls_made = await make_calls_until_budget_exceeded( - session=session, - key=key, - call_function=chat_completion, - model="fake-openai-endpoint", - ) - - assert calls_made == 0, "Should make no calls before budget exceeded" - - -@pytest.mark.asyncio -async def test_chat_completion_high_budget(): - """ - Test budget enforcement for chat completions: - 1. Create key with $0.01 budget - 2. Make chat completion calls until budget exceeded - 3. Verify budget exceeded error - """ - async with aiohttp.ClientSession() as session: - # Create key with $0.01 budget - key_gen = await generate_key(session=session, max_budget=0.001) - print("response from key generation: ", key_gen) - key = key_gen["key"] - - # Make calls until budget exceeded - calls_made = await make_calls_until_budget_exceeded( - session=session, - key=key, - call_function=chat_completion, - model="fake-openai-endpoint", - ) - - assert ( - calls_made > 0 - ), "Should make at least one successful call before budget exceeded" - - -@pytest.mark.parametrize( - "field", - [ - "max_budget", - "rpm_limit", - "tpm_limit", - ], -) -@pytest.mark.asyncio -async def test_key_limit_modifications(field): - # Create initial key - client = AsyncClient(base_url="http://0.0.0.0:4000") - key_data = {"max_budget": None, "rpm_limit": None, "tpm_limit": None} - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"} - response = await client.post("/key/generate", json=key_data, headers=headers) - assert response.status_code == 200 - generate_key_response = response.json() - print("generate_key_response: ", json.dumps(generate_key_response, indent=4)) - key_id = generate_key_response["key"] - - # Update key with any non-null value for the field - update_data = {"key": key_id} - update_data[field] = 10 # Any non-null value works - print("update_data: ", json.dumps(update_data, indent=4)) - response = await client.post(f"/key/update", json=update_data, headers=headers) - assert response.status_code == 200 - assert response.json()[field] is not None - - # Reset limit to null - print(f"resetting {field} to null") - update_data[field] = None - response = await client.post(f"/key/update", json=update_data, headers=headers) - print("response: ", json.dumps(response.json(), indent=4)) - assert response.status_code == 200 - assert response.json()[field] is None - - -@pytest.mark.parametrize( - "field", - [ - "max_budget", - ], -) -@pytest.mark.asyncio -async def test_team_limit_modifications(field): - # Create initial team - client = AsyncClient(base_url="http://0.0.0.0:4000") - team_data = {"max_budget": None, "rpm_limit": None, "tpm_limit": None} - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"} - response = await client.post("/team/new", json=team_data, headers=headers) - print("response: ", json.dumps(response.json(), indent=4)) - assert response.status_code == 200 - team_id = response.json()["team_id"] - - # Update team with any non-null value for the field - update_data = {"team_id": team_id} - update_data[field] = 10 # Any non-null value works - response = await client.post(f"/team/update", json=update_data, headers=headers) - print("response: ", json.dumps(response.json(), indent=4)) - assert response.status_code == 200 - assert response.json()["data"][field] is not None - - # Reset limit to null - print(f"resetting {field} to null") - update_data[field] = None - response = await client.post(f"/team/update", json=update_data, headers=headers) - print("response: ", json.dumps(response.json(), indent=4)) - assert response.status_code == 200 - assert response.json()["data"][field] is None - - -async def generate_team_key( - session, - team_id: str, - max_budget: Optional[float] = None, -): - """Helper function to generate a key for a specific team""" - url = "http://0.0.0.0:4000/key/generate" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data: dict[str, Any] = {"team_id": team_id} - if max_budget is not None: - data["max_budget"] = max_budget - async with session.post(url, headers=headers, json=data) as response: - return await response.json() - - -async def create_team( - session, - max_budget=None, - models: Optional[list[str]] = None, - team_alias: Optional[str] = None, -): - """Helper function to create a new team""" - url = f"{PROXY_BASE}/team/new" - data: dict[str, Any] = {"max_budget": max_budget} - if models is not None: - data["models"] = models - if team_alias is not None: - data["team_alias"] = team_alias - async with session.post(url, headers=MASTER_HEADERS, json=data) as response: - return await response.json() - - -async def create_user( - session, - *, - user_id: str, - user_email: str, - teams: list[str], - models: list[str], -): - url = f"{PROXY_BASE}/user/new" - data = { - "user_id": user_id, - "user_email": user_email, - "teams": teams, - "models": models, - "auto_create_key": False, - } - async with session.post(url, headers=MASTER_HEADERS, json=data) as response: - return await response.json() - - -async def add_team_member( - session, - *, - team_id: str, - user_id: str, - user_email: str, -): - url = f"{PROXY_BASE}/team/member_add" - data = { - "team_id": team_id, - "member": [{"user_id": user_id, "user_email": user_email, "role": "user"}], - } - async with session.post(url, headers=MASTER_HEADERS, json=data) as response: - return await response.json() - - -async def obtain_cli_sso_token_via_poll_flow( - session, - *, - user_id: str, - user_email: str, - team_id: str, - team_alias: str, - models: list[str], -) -> str: - """ - Obtain a CLI SSO JWT through the same HTTP flow as `lite login`: - /sso/cli/start -> (SSO callback) -> /sso/cli/complete -> /sso/cli/poll. - - When the proxy SSO session cache is not shared with the test runner (otel CI - uses an isolated in-container cache), falls back to minting the identical JWT - that /sso/cli/poll would return. - """ - async with session.post(f"{PROXY_BASE}/sso/cli/start") as resp: - resp.raise_for_status() - start = await resp.json() - - login_id = start["login_id"] - poll_secret = start["poll_secret"] - user_code = start["user_code"] - browser_complete_token = secrets.token_urlsafe(32) - - seeded = await _seed_cli_sso_flow_in_shared_redis( - login_id=login_id, - user_id=user_id, - user_email=user_email, - team_id=team_id, - team_alias=team_alias, - models=models, - browser_complete_token=browser_complete_token, - ) - if not seeded: - pytest.skip("Shared Redis not available; skipping full poll-flow test") - - async with session.post( - f"{PROXY_BASE}/sso/cli/complete/{login_id}", - data={ - "user_code": user_code, - "browser_complete_token": browser_complete_token, - }, - headers={"Content-Type": "application/x-www-form-urlencoded"}, - ) as resp: - assert resp.status == 200, await resp.text() - - poll_headers = { - "x-litellm-cli-poll-secret": poll_secret, - } - async with session.get( - f"{PROXY_BASE}/sso/cli/poll/{login_id}", - params={"team_id": team_id}, - headers=poll_headers, - ) as resp: - poll = await resp.json() - - assert poll.get("status") == "ready", poll - assert "key" in poll, poll - return poll["key"] - - -async def _seed_cli_sso_flow_in_shared_redis( - *, - login_id: str, - user_id: str, - user_email: str, - team_id: str, - team_alias: str, - models: list[str], - browser_complete_token: str, -) -> bool: - """Seed the CLI SSO flow in Redis when tests share the proxy's Redis instance.""" - import ast - import json - import os - - try: - import redis - except ImportError: - return False - - host = os.getenv("REDIS_HOST") - if not host: - return False - - try: - client = redis.Redis( - host=host, - port=int(os.getenv("REDIS_PORT", "6379")), - password=os.getenv("REDIS_PASSWORD") or None, - decode_responses=True, - ) - client.ping() - except Exception: - return False - - from litellm.proxy.management_endpoints.ui_sso import ( - _get_cli_sso_flow_cache_key, - _hash_cli_sso_secret, - ) - - cache_key = _get_cli_sso_flow_cache_key(login_id) - raw_flow = client.get(cache_key) - if raw_flow is None: - return False - - try: - flow = ast.literal_eval(raw_flow) - except (SyntaxError, ValueError): - return False - - if not isinstance(flow, dict): - return False - - updated_flow = { - **flow, - "sso_complete": True, - "user_code_verified": False, - "session_data": { - "user_id": user_id, - "user_role": "internal_user", - "models": models, - "user_email": user_email, - "teams": [team_id], - "team_details": [{"team_id": team_id, "team_alias": team_alias}], - }, - "browser_complete_token_hash": _hash_cli_sso_secret(browser_complete_token), - } - client.setex(cache_key, 600, json.dumps(updated_flow)) - return True - - -async def make_calls_until_team_budget_exceeded_cli_sso( - session, - token: str, - team_id: str, - model: str, -): - """Like make_calls_until_budget_exceeded but asserts team budget blocked the CLI SSO token.""" - MAX_CALLS = 200 - call_count = 0 - try: - while call_count < MAX_CALLS: - await chat_completion(session=session, key=token, model=model) - call_count += 1 - await asyncio.sleep(0.1) - pytest.fail(f"Budget was not exceeded after {MAX_CALLS} calls") - except openai.APIStatusError as e: - error_dict = e.body - assert error_dict["code"] == "422" - assert error_dict["type"] == "budget_exceeded" - message = error_dict["message"] - assert "Budget has been exceeded!" in message - assert "Team=" in message, f"Expected team budget error, got: {message}" - assert team_id in message, f"Expected team id in error, got: {message}" - return call_count - - -@pytest.mark.asyncio -async def test_team_budget_enforcement(): - """ - Test budget enforcement for team-wide budgets: - 1. Create team with low budget - 2. Create key for that team - 3. Make calls until team budget exceeded - 4. Verify budget exceeded error - """ - async with aiohttp.ClientSession() as session: - # Create team with low budget - team_response = await create_team(session=session, max_budget=0.0000000005) - team_id = team_response["team_id"] - - # Create key for team (no specific budget) - key_gen = await generate_team_key(session=session, team_id=team_id) - key = key_gen["key"] - - # Make calls until budget exceeded - calls_made = await make_calls_until_budget_exceeded( - session=session, - key=key, - call_function=chat_completion, - model="fake-openai-endpoint", - ) - - assert ( - calls_made > 0 - ), "Should make at least one successful call before team budget exceeded" - - -@pytest.mark.asyncio -async def test_team_budget_enforcement_cli_sso_token(): - """ - Team budget enforcement for CLI SSO session tokens (lite login JWT). - - 1. Create team with a tiny max_budget and a user on that team - 2. Obtain a CLI SSO JWT (HTTP poll flow when Redis is shared, else mint) - 3. Make chat completion calls until the team budget is exceeded - 4. Verify HTTP 422 budget_exceeded names the team - """ - user_id = f"cli-budget-user-{uuid.uuid4().hex[:8]}" - user_email = f"{user_id}@example.com" - team_alias = f"cli-budget-team-{uuid.uuid4().hex[:8]}" - - async with aiohttp.ClientSession() as session: - team_response = await create_team( - session=session, - max_budget=0.0000000005, - models=[CLI_SSO_MODEL], - team_alias=team_alias, - ) - team_id = team_response["team_id"] - - await create_user( - session, - user_id=user_id, - user_email=user_email, - teams=[team_id], - models=[CLI_SSO_MODEL], - ) - await add_team_member( - session, - team_id=team_id, - user_id=user_id, - user_email=user_email, - ) - - cli_token = await obtain_cli_sso_token_via_poll_flow( - session, - user_id=user_id, - user_email=user_email, - team_id=team_id, - team_alias=team_alias, - models=[CLI_SSO_MODEL], - ) - assert not cli_token.startswith( - "sk-" - ), "CLI SSO token must not be a virtual key" - - calls_made = await make_calls_until_team_budget_exceeded_cli_sso( - session=session, - token=cli_token, - team_id=team_id, - model=CLI_SSO_MODEL, - ) - - assert ( - calls_made > 0 - ), "Should make at least one successful call before team budget exceeded" - - - # Verify it was the team budget that was exceeded diff --git a/tests/otel_tests/test_e2e_model_access.py b/tests/otel_tests/test_e2e_model_access.py deleted file mode 100644 index bdd6b597f68..00000000000 --- a/tests/otel_tests/test_e2e_model_access.py +++ /dev/null @@ -1,304 +0,0 @@ -import os -import pytest -import asyncio -import aiohttp -import json -from httpx import AsyncClient -from openai import PermissionDeniedError -from typing import Any, Optional, List, Literal - - -# The proxy strips client-supplied `mock_response` unless the calling key or -# team has this admin-metadata flag set. See `_UNTRUSTED_ROOT_CONTROL_FIELDS` -# in litellm/proxy/litellm_pre_call_utils.py. -_ALLOW_CLIENT_MOCK_METADATA = {"allow_client_mock_response": True} - - -async def generate_key( - session, models: Optional[List[str]] = None, team_id: Optional[str] = None -): - """Helper function to generate a key with specific model access controls""" - url = "http://0.0.0.0:4000/key/generate" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data: dict = {"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA)} - if models is not None: - data["models"] = models - if team_id is not None: - data["team_id"] = team_id - async with session.post(url, headers=headers, json=data) as response: - return await response.json() - - -async def generate_team(session, models: Optional[List[str]] = None): - """Helper function to generate a team with specific model access""" - url = "http://0.0.0.0:4000/team/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data: dict = {"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA)} - if models is not None: - data["models"] = models - async with session.post(url, headers=headers, json=data) as response: - return await response.json() - - -async def mock_chat_completion(session, key: str, model: str): - """Make a chat completion request using OpenAI SDK""" - from openai import AsyncOpenAI - from litellm._uuid import uuid - - client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000/v1") - - response = await client.chat.completions.create( - model=model, - messages=[{"role": "user", "content": f"Say hello! {uuid.uuid4()}"}], - extra_body={ - "mock_response": "mock_response", - }, - ) - return response - - -@pytest.mark.parametrize( - "key_models, test_model, expect_success", - [ - (["openai/*"], "anthropic/claude-2", False), # Non-matching model - (["gpt-5.5"], "gpt-5.5", True), # Exact model match - (["bedrock/*"], "bedrock/anthropic.claude-3", True), # Bedrock wildcard - (["bedrock/anthropic.*"], "bedrock/anthropic.claude-3", True), # Pattern match - (["bedrock/anthropic.*"], "bedrock/amazon.titan", False), # Pattern non-match - (None, "gpt-5.5", True), # No model restrictions - ([], "gpt-5.5", True), # Empty model list - ], -) -@pytest.mark.asyncio -async def test_model_access_patterns(key_models, test_model, expect_success): - """ - Test model access patterns for API keys: - 1. Create key with specific model access pattern - 2. Attempt to make completion with test model - 3. Verify access is granted/denied as expected - """ - async with aiohttp.ClientSession() as session: - # Generate key with specified model access - key_gen = await generate_key(session=session, models=key_models) - key = key_gen["key"] - - try: - response = await mock_chat_completion( - session=session, - key=key, - model=test_model, - ) - if not expect_success: - pytest.fail(f"Expected request to fail for model {test_model}") - assert ( - response is not None - ), "Should get valid response when access is allowed" - except Exception as e: - if expect_success: - pytest.fail(f"Expected request to succeed but got error: {e}") - _error_body = e.body - - # Assert error structure and values - assert _error_body["type"] == "key_model_access_denied" - assert _error_body["param"] == "model" - assert _error_body["code"] == "403" - assert "is not available for this API key" in _error_body["message"] - - -@pytest.mark.asyncio -async def test_model_access_update(): - """ - Test updating model access for an existing key: - 1. Create key with restricted model access - 2. Verify access patterns - 3. Update key with new model access - 4. Verify new access patterns - """ - client = AsyncClient(base_url="http://0.0.0.0:4000") - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"} - - # Create initial key with restricted access - response = await client.post( - "/key/generate", - json={ - "models": ["openai/gpt-5.5"], - "metadata": dict(_ALLOW_CLIENT_MOCK_METADATA), - }, - headers=headers, - ) - assert response.status_code == 200 - key_data = response.json() - key = key_data["key"] - - # Test initial access - async with aiohttp.ClientSession() as session: - # Should work with gpt-5.5 - await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5") - - # Should fail with gpt-5-mini - with pytest.raises(PermissionDeniedError) as exc_info: - await mock_chat_completion( - session=session, key=key, model="openai/gpt-5-mini" - ) - _validate_model_access_exception( - exc_info.value, expected_type="key_model_access_denied" - ) - - # Update key with new model access - response = await client.post( - "/key/update", json={"key": key, "models": ["openai/*"]}, headers=headers - ) - assert response.status_code == 200 - - # Test updated access - async with aiohttp.ClientSession() as session: - # Both models should now work - await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5") - await mock_chat_completion( - session=session, key=key, model="openai/gpt-5-mini" - ) - - # Non-OpenAI model should still fail - with pytest.raises(PermissionDeniedError) as exc_info: - await mock_chat_completion( - session=session, key=key, model="anthropic/claude-2" - ) - _validate_model_access_exception( - exc_info.value, expected_type="key_model_access_denied" - ) - - -@pytest.mark.parametrize( - "team_models, test_model, expect_success", - [ - (["openai/*"], "anthropic/claude-2", False), # Non-matching model - ], -) -@pytest.mark.asyncio -async def test_team_model_access_patterns(team_models, test_model, expect_success): - """ - Test model access patterns for team-based API keys: - 1. Create team with specific model access pattern - 2. Generate key for that team - 3. Attempt to make completion with test model - 4. Verify access is granted/denied as expected - """ - client = AsyncClient(base_url="http://0.0.0.0:4000") - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"} - - async with aiohttp.ClientSession() as session: - try: - team_gen = await generate_team(session=session, models=team_models) - print("created team", team_gen) - team_id = team_gen["team_id"] - key_gen = await generate_key(session=session, team_id=team_id) - print("created key", key_gen) - key = key_gen["key"] - response = await mock_chat_completion( - session=session, - key=key, - model=test_model, - ) - if not expect_success: - pytest.fail(f"Expected request to fail for model {test_model}") - assert ( - response is not None - ), "Should get valid response when access is allowed" - except Exception as e: - if expect_success: - pytest.fail(f"Expected request to succeed but got error: {e}") - _validate_model_access_exception( - e, expected_type="team_model_access_denied" - ) - - -@pytest.mark.asyncio -async def test_team_model_access_update(): - """ - Test updating model access for a team: - 1. Create team with restricted model access - 2. Verify access patterns - 3. Update team with new model access - 4. Verify new access patterns - """ - client = AsyncClient(base_url="http://0.0.0.0:4000") - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"} - - # Create initial team with restricted access - response = await client.post( - "/team/new", - json={ - "models": ["openai/gpt-5.5"], - "name": "test-team", - "metadata": dict(_ALLOW_CLIENT_MOCK_METADATA), - }, - headers=headers, - ) - assert response.status_code == 200 - team_data = response.json() - team_id = team_data["team_id"] - - # Generate a key for this team - response = await client.post( - "/key/generate", - json={ - "team_id": team_id, - "metadata": dict(_ALLOW_CLIENT_MOCK_METADATA), - }, - headers=headers, - ) - assert response.status_code == 200 - key = response.json()["key"] - - # Test initial access - async with aiohttp.ClientSession() as session: - # Should work with gpt-5.5 - await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5") - - # Should fail with gpt-5-mini - with pytest.raises(PermissionDeniedError) as exc_info: - await mock_chat_completion( - session=session, key=key, model="openai/gpt-5-mini" - ) - _validate_model_access_exception( - exc_info.value, expected_type="team_model_access_denied" - ) - - # Update team with new model access - response = await client.post( - "/team/update", - json={"team_id": team_id, "models": ["openai/*"]}, - headers=headers, - ) - assert response.status_code == 200 - - # Test updated access - async with aiohttp.ClientSession() as session: - # Both models should now work - await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5") - await mock_chat_completion( - session=session, key=key, model="openai/gpt-5-mini" - ) - - # Non-OpenAI model should still fail - with pytest.raises(PermissionDeniedError) as exc_info: - await mock_chat_completion( - session=session, key=key, model="anthropic/claude-2" - ) - _validate_model_access_exception( - exc_info.value, expected_type="team_model_access_denied" - ) - - -def _validate_model_access_exception( - e: Exception, - expected_type: Literal["key_model_access_denied", "team_model_access_denied"], -): - _error_body = e.body - - # Assert error structure and values - assert _error_body["type"] == expected_type - assert _error_body["param"] == "model" - assert _error_body["code"] == "403" - assert "is not available for this API key" in _error_body["message"] - assert "not allowed to access model" not in _error_body["message"] diff --git a/tests/otel_tests/test_guardrails.py b/tests/otel_tests/test_guardrails.py index 36f2d019a34..ae82bb8908c 100644 --- a/tests/otel_tests/test_guardrails.py +++ b/tests/otel_tests/test_guardrails.py @@ -48,98 +48,6 @@ async def chat_completion( return await response.json(), response_headers -async def generate_key( - session, guardrails: Optional[List] = None, team_id: Optional[str] = None -): - url = "http://0.0.0.0:4000/key/generate" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = {} - if guardrails: - data["guardrails"] = guardrails - if team_id: - data["team_id"] = team_id - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - return await response.json() - - -@pytest.mark.asyncio -async def test_no_llm_guard_triggered(): - """ - - Tests a request where no content mod is triggered - - Assert that the guardrails applied are returned in the response headers - """ - async with aiohttp.ClientSession() as session: - response, headers = await chat_completion( - session, - os.environ["LITELLM_MASTER_KEY"], - model="fake-openai-endpoint", - messages=[{"role": "user", "content": f"Hello what's the weather"}], - guardrails=[], - ) - await asyncio.sleep(3) - - print("response=", response, "response headers", headers) - - assert "x-litellm-applied-guardrails" not in headers - - -@pytest.mark.asyncio -async def test_guardrails_with_api_key_controls(): - """ - - Make two API Keys - - Key 1 with no guardrails - - Key 2 with guardrails - - Request to Key 1 -> should be success with no guardrails - - Request to Key 2 -> should be error since guardrails are triggered - """ - async with aiohttp.ClientSession() as session: - key_with_guardrails = await generate_key( - session=session, - guardrails=[ - "bedrock-pre-guard", - ], - ) - - key_with_guardrails = key_with_guardrails["key"] - - key_without_guardrails = await generate_key(session=session, guardrails=None) - - key_without_guardrails = key_without_guardrails["key"] - - # test no guardrails triggered for key without guardrails - response, headers = await chat_completion( - session, - key_without_guardrails, - model="fake-openai-endpoint", - messages=[{"role": "user", "content": f"Hello what's the weather"}], - ) - await asyncio.sleep(3) - - print("response=", response, "response headers", headers) - assert "x-litellm-applied-guardrails" not in headers - - # test guardrails triggered for key with guardrails - response, headers = await chat_completion( - session, - key_with_guardrails, - model="fake-openai-endpoint", - messages=[{"role": "user", "content": f"Hello my name is ishaan@berri.ai"}], - ) - - assert "x-litellm-applied-guardrails" in headers - assert headers["x-litellm-applied-guardrails"] == "bedrock-pre-guard" - - @pytest.mark.asyncio async def test_bedrock_guardrail_triggered(): """ @@ -160,105 +68,6 @@ async def test_bedrock_guardrail_triggered(): assert "Violated guardrail policy" in str(e) -@pytest.mark.asyncio -async def test_custom_guardrail_during_call_triggered(): - """ - - Tests a request where our bedrock guardrail should be triggered - - Assert that the guardrails applied are returned in the response headers - """ - async with aiohttp.ClientSession() as session: - with pytest.raises(Exception, match="Guardrail failed words - `litellm` detected") as exc_info: - response, headers = await chat_completion( - session, - os.environ["LITELLM_MASTER_KEY"], - model="fake-openai-endpoint", - messages=[{"role": "user", "content": f"Hello do you like litellm?"}], - guardrails=["custom-during-guard"], - ) - e = exc_info.value - print(e) - assert "Guardrail failed words - `litellm` detected" in str(e) - - -async def create_team(session, guardrails: Optional[List] = None): - url = "http://0.0.0.0:4000/team/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = {"guardrails": guardrails} - - print("request data=", data) - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - return await response.json() - - -@pytest.mark.asyncio -async def test_guardrails_with_team_controls(): - """ - - Create a team with guardrails - - Make two API Keys - - Key 1 not associated with team - - Key 2 associated with team (inherits team guardrails) - - Request with Key 1 -> should be success with no guardrails - - Request with Key 2 -> should error since team guardrails are triggered - """ - async with aiohttp.ClientSession() as session: - - # Create team with guardrails - team = await create_team( - session=session, - guardrails=[ - "bedrock-pre-guard", - ], - ) - - print("team=", team) - - team_id = team["team_id"] - - # Create key with team association - key_with_team = await generate_key(session=session, team_id=team_id) - key_with_team = key_with_team["key"] - - # Create key without team - key_without_team = await generate_key( - session=session, - ) - key_without_team = key_without_team["key"] - - # Test no guardrails triggered for key without a team - response, headers = await chat_completion( - session, - key_without_team, - model="fake-openai-endpoint", - messages=[{"role": "user", "content": "Hello my name is ishaan@berri.ai"}], - ) - await asyncio.sleep(3) - - print("response=", response, "response headers", headers) - assert "x-litellm-applied-guardrails" not in headers - - response, headers = await chat_completion( - session, - key_with_team, - model="fake-openai-endpoint", - messages=[{"role": "user", "content": "Hello my name is ishaan@berri.ai"}], - ) - - print("response headers=", json.dumps(headers, indent=4)) - - assert "x-litellm-applied-guardrails" in headers - assert headers["x-litellm-applied-guardrails"] == "bedrock-pre-guard" - - async def get_guardrail_lb_counts(session): """Get the current guardrail load balancing call counts from the proxy.""" url = "http://0.0.0.0:4000/guardrail/lb/counts" diff --git a/tests/otel_tests/test_key_logging_callbacks.py b/tests/otel_tests/test_key_logging_callbacks.py deleted file mode 100644 index 1736831eb69..00000000000 --- a/tests/otel_tests/test_key_logging_callbacks.py +++ /dev/null @@ -1,70 +0,0 @@ -""" -Tests for Key based logging callbacks - -""" - -import os -import httpx -import pytest - - -@pytest.mark.asyncio() -async def test_key_logging_callbacks(): - """ - Create virtual key with a logging callback set on the key - Call /key/health for the key -> it should be unhealthy - """ - # Generate a key with logging callback - generate_url = "http://0.0.0.0:4000/key/generate" - generate_headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - } - generate_payload = { - "metadata": { - "logging": [ - { - "callback_name": "gcs_bucket", - "callback_type": "success_and_failure", - "callback_vars": { - "gcs_bucket_name": "key-logging-project1", - "gcs_path_service_account": "bad-service-account", - }, - } - ] - } - } - - async with httpx.AsyncClient() as client: - generate_response = await client.post( - generate_url, headers=generate_headers, json=generate_payload - ) - - assert generate_response.status_code == 200 - generate_data = generate_response.json() - assert "key" in generate_data - - _key = generate_data["key"] - - # Check key health - health_url = "http://localhost:4000/key/health" - health_headers = { - "Authorization": f"Bearer {_key}", - "Content-Type": "application/json", - } - - async with httpx.AsyncClient() as client: - health_response = await client.post(health_url, headers=health_headers, json={}) - - assert health_response.status_code == 200 - health_data = health_response.json() - print("key_health_data", health_data) - # Check the response format and content - assert "key" in health_data - assert "logging_callbacks" in health_data - assert health_data["logging_callbacks"]["callbacks"] == ["gcs_bucket"] - assert health_data["logging_callbacks"]["status"] == "unhealthy" - assert ( - "GCS_BUCKET_NAME is not set in the environment" - in health_data["logging_callbacks"]["details"] - ) diff --git a/tests/otel_tests/test_model_info.py b/tests/otel_tests/test_model_info.py deleted file mode 100644 index 66a81eeee51..00000000000 --- a/tests/otel_tests/test_model_info.py +++ /dev/null @@ -1,29 +0,0 @@ -""" -/model/info test -""" - -import os -import httpx -import pytest - - -@pytest.mark.asyncio() -async def test_custom_model_supports_vision(): - async with httpx.AsyncClient() as client: - response = await client.get( - "http://localhost:4000/model/info", - headers={"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}, - ) - assert response.status_code == 200 - - data = response.json()["data"] - - print("response from /model/info", data) - llava_model = next( - (model for model in data if model["model_name"] == "llava-hf"), None - ) - - assert llava_model is not None, "llava-hf model not found in response" - assert ( - llava_model["model_info"]["supports_vision"] == True - ), "llava-hf model should support vision" diff --git a/tests/otel_tests/test_moderations.py b/tests/otel_tests/test_moderations.py index a9c73c93500..1a67827aa61 100644 --- a/tests/otel_tests/test_moderations.py +++ b/tests/otel_tests/test_moderations.py @@ -28,28 +28,6 @@ async def make_moderations_curl_request( return await response.json() -@pytest.mark.asyncio -async def test_basic_moderations_on_proxy_no_model(): - """ - Test moderations endpoint on proxy when no `model` is specified in the request - """ - async with aiohttp.ClientSession() as session: - test_text = "I want to harm someone" # Test text that should trigger moderation - request_data = { - "input": test_text, - } - try: - response = await make_moderations_curl_request( - session, - os.environ["LITELLM_MASTER_KEY"], - request_data, - ) - print("response=", response) - except Exception as e: - print(e) - pytest.fail("Moderations request failed") - - @pytest.mark.asyncio async def test_basic_moderations_on_proxy_with_model(): """ diff --git a/tests/otel_tests/test_prometheus.py b/tests/otel_tests/test_prometheus.py deleted file mode 100644 index 7224334105d..00000000000 --- a/tests/otel_tests/test_prometheus.py +++ /dev/null @@ -1,911 +0,0 @@ -""" -Unit tests for prometheus metrics -""" - -import os -import pytest -import aiohttp -import asyncio -from litellm._uuid import uuid -from openai import AsyncOpenAI -from typing import Dict, Any - - -END_USER_ID = "my-test-user-34" - - -async def make_bad_chat_completion_request(session, key): - url = "http://0.0.0.0:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": "fake-azure-endpoint", - "messages": [{"role": "user", "content": "Hello"}], - } - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - return status, response_text - - -async def make_good_chat_completion_request(session, key): - url = "http://0.0.0.0:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - - data = { - "model": "fake-openai-endpoint", - "messages": [{"role": "user", "content": f"Hello {uuid.uuid4()}"}], - "tags": ["teamB"], - "user": END_USER_ID, # test if disable end user tracking for prometheus works - } - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - return status, response_text - - -async def make_chat_completion_request_with_fallback(session, key): - url = "http://0.0.0.0:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": "fake-azure-endpoint", - "messages": [{"role": "user", "content": "Hello"}], - "fallbacks": ["fake-openai-endpoint"], - } - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - # make a request with a failed fallback - data = { - "model": "fake-azure-endpoint", - "messages": [{"role": "user", "content": "Hello"}], - "fallbacks": ["unknown-model"], - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - return - - -@pytest.mark.asyncio -async def test_proxy_failure_metrics(): - """ - - Make 1 bad chat completion call to "fake-azure-endpoint" - - GET /metrics - - assert the failure metric for the requested model is incremented by 1 - - Assert the Exception class and status code are correct - """ - async with aiohttp.ClientSession() as session: - # Make a bad chat completion call - status, response_text = await make_bad_chat_completion_request( - session, os.environ["LITELLM_MASTER_KEY"] - ) - - # Check if the request failed as expected - assert status == 429, f"Expected status 429, but got {status}" - - # Get metrics - async with session.get("http://0.0.0.0:4000/metrics") as response: - metrics = await response.text() - - print("/metrics", metrics) - - # Check if the failure metric is present and correct - use pattern matching for robustness - # Labels are ordered alphabetically by Prometheus: api_key_alias, end_user, exception_class, - # exception_status, hashed_api_key, requested_model, route, team, team_alias, user, user_email - # Note: client_ip, user_agent, model_id are present but we use substring matching to be flexible - # Check for both the new metric and deprecated metric for backwards compatibility - expected_patterns = [ - "litellm_proxy_failed_requests_metric_total{", # New metric - "litellm_llm_api_failed_requests_metric_total{", # Deprecated but may still be used - ] - - # Master-key auth substitutes LITELLM_PROXY_MASTER_KEY_ALIAS for - # hash_token(master_key) so the master key (or its hash) never - # propagates into metrics. See PR #26484. - from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS - - expected_hashed_api_key = LITELLM_PROXY_MASTER_KEY_ALIAS - - # Check if either pattern is in metrics and contains required fields - found_metric = False - for pattern in expected_patterns: - for line in metrics.split("\n"): - # For proxy metric, check proxy-specific fields - if "litellm_proxy_failed_requests_metric_total{" in line: - if ( - 'api_key_alias="None"' in line - and 'exception_class="Openai.RateLimitError"' in line - and 'exception_status="429"' in line - and f'hashed_api_key="{expected_hashed_api_key}"' in line - and 'requested_model="fake-azure-endpoint"' in line - and 'route="/chat/completions"' in line - ): - found_metric = True - break - # For deprecated llm_api metric, check llm-specific fields - elif "litellm_llm_api_failed_requests_metric_total{" in line: - if ( - f'hashed_api_key="{expected_hashed_api_key}"' in line - and 'model="429"' in line - ): # The deprecated metric uses the actual model from the request - found_metric = True - break - if found_metric: - break - - assert ( - found_metric - ), f"Expected failure metric not found in /metrics. Looking for either litellm_proxy_failed_requests_metric_total or litellm_llm_api_failed_requests_metric_total with required fields" - - # Check total requests metric similarly - # The litellm_proxy_total_requests_metric_total should be present - total_requests_pattern = "litellm_proxy_total_requests_metric_total{" - - found_total_metric = False - for line in metrics.split("\n"): - if ( - total_requests_pattern in line - and f'hashed_api_key="{expected_hashed_api_key}"' in line - and 'requested_model="fake-azure-endpoint"' in line - and 'status_code="429"' in line - ): - found_total_metric = True - break - - assert ( - found_total_metric - ), f"Expected total requests metric not found in /metrics. Looking for: {total_requests_pattern} with hashed_api_key and status_code=429" - - -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=2) -async def test_proxy_success_metrics(): - """ - Make 1 good /chat/completions call to "openai/gpt-5-mini" - GET /metrics - Assert the success metric is incremented by 1 - """ - - async with aiohttp.ClientSession() as session: - # Make a good chat completion call - status, response_text = await make_good_chat_completion_request( - session, os.environ["LITELLM_MASTER_KEY"] - ) - - # Check if the request succeeded as expected - assert status == 200, f"Expected status 200, but got {status}" - - # Get metrics - async with session.get("http://0.0.0.0:4000/metrics") as response: - metrics = await response.text() - - print("/metrics", metrics) - - assert END_USER_ID not in metrics - - # Master-key auth substitutes LITELLM_PROXY_MASTER_KEY_ALIAS for - # hash_token(master_key) (PR #26484). - from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS - - expected_hashed_api_key = LITELLM_PROXY_MASTER_KEY_ALIAS - - # Check if the success metric is present and correct - use flexible matching - # Check for request_total_latency_metric with required fields - # Note: The model can be "gpt-3.5-turbo-0301" or similar depending on what's returned - found_request_latency = False - for line in metrics.split("\n"): - if ( - "litellm_request_total_latency_metric_bucket{" in line - and 'api_key_alias="None"' in line - and f'hashed_api_key="{expected_hashed_api_key}"' in line - and 'requested_model="fake-openai-endpoint"' in line - and 'le="0.005"' in line - ): - found_request_latency = True - break - - assert ( - found_request_latency - ), "Expected litellm_request_total_latency_metric_bucket not found in /metrics" - - # Check for llm_api_latency_metric with required fields - found_api_latency = False - for line in metrics.split("\n"): - if ( - "litellm_llm_api_latency_metric_bucket{" in line - and 'api_key_alias="None"' in line - and f'hashed_api_key="{expected_hashed_api_key}"' in line - and 'requested_model="fake-openai-endpoint"' in line - and 'le="0.005"' in line - ): - found_api_latency = True - break - - assert ( - found_api_latency - ), "Expected litellm_llm_api_latency_metric_bucket not found in /metrics" - - verify_latency_metrics(metrics) - - -def verify_latency_metrics(metrics: str): - """ - Assert that LATENCY_BUCKETS distribution is used for - - litellm_request_total_latency_metric_bucket - - litellm_llm_api_latency_metric_bucket - - Very important to verify that the overhead latency metric is present - """ - from litellm.types.integrations.prometheus import LATENCY_BUCKETS - import re - import time - - time.sleep(2) - - metric_names = [ - "litellm_request_total_latency_metric_bucket", - "litellm_llm_api_latency_metric_bucket", - "litellm_overhead_latency_metric_bucket", - ] - - for metric_name in metric_names: - # Extract all 'le' values for the current metric - pattern = rf'{metric_name}{{.*?le="(.*?)".*?}}' - le_values = re.findall(pattern, metrics) - - # Convert to set for easier comparison - actual_buckets = set(le_values) - - print("actual_buckets", actual_buckets) - expected_buckets = [] - for bucket in LATENCY_BUCKETS: - expected_buckets.append(str(bucket)) - - # replace inf with +Inf - expected_buckets = [ - bucket.replace("inf", "+Inf") for bucket in expected_buckets - ] - - print("expected_buckets", expected_buckets) - expected_buckets = set(expected_buckets) - # Verify all expected buckets are present - assert ( - actual_buckets == expected_buckets - ), f"Mismatch in {metric_name} buckets. Expected: {expected_buckets}, Got: {actual_buckets}" - - -@pytest.mark.asyncio -async def test_proxy_fallback_metrics(): - """ - Make 1 request with a client side fallback - check metrics - """ - - async with aiohttp.ClientSession() as session: - # Make a good chat completion call - await make_chat_completion_request_with_fallback(session, os.environ["LITELLM_MASTER_KEY"]) - - # Get metrics - async with session.get("http://0.0.0.0:4000/metrics") as response: - metrics = await response.text() - - print("/metrics", metrics) - - # Master-key auth substitutes LITELLM_PROXY_MASTER_KEY_ALIAS for - # hash_token(master_key) (PR #26484). - from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS - - expected_hashed_api_key = LITELLM_PROXY_MASTER_KEY_ALIAS - - # Check if successful fallback metric is incremented - use flexible matching - found_successful_fallback = False - for line in metrics.split("\n"): - if ( - "litellm_deployment_successful_fallbacks_total{" in line - and 'api_key_alias="None"' in line - and 'exception_class="Openai.RateLimitError"' in line - and 'exception_status="429"' in line - and 'fallback_model="fake-openai-endpoint"' in line - and f'hashed_api_key="{expected_hashed_api_key}"' in line - and 'requested_model="fake-azure-endpoint"' in line - and "1.0" in line - ): - found_successful_fallback = True - break - - assert ( - found_successful_fallback - ), "Expected litellm_deployment_successful_fallbacks_total metric not found in /metrics" - - # Check if failed fallback metric is incremented - use flexible matching - found_failed_fallback = False - for line in metrics.split("\n"): - if ( - "litellm_deployment_failed_fallbacks_total{" in line - and 'api_key_alias="None"' in line - and 'exception_class="Openai.RateLimitError"' in line - and 'exception_status="429"' in line - and 'fallback_model="unknown-model"' in line - and f'hashed_api_key="{expected_hashed_api_key}"' in line - and 'requested_model="fake-azure-endpoint"' in line - and "1.0" in line - ): - found_failed_fallback = True - break - - assert ( - found_failed_fallback - ), "Expected litellm_deployment_failed_fallbacks_total metric not found in /metrics" - - -async def create_test_team( - session: aiohttp.ClientSession, team_data: Dict[str, Any] -) -> str: - """Create a new team and return the team_id""" - url = "http://0.0.0.0:4000/team/new" - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - } - - async with session.post(url, headers=headers, json=team_data) as response: - assert ( - response.status == 200 - ), f"Failed to create team. Status: {response.status}" - team_info = await response.json() - return team_info["team_id"] - - -async def create_test_user( - session: aiohttp.ClientSession, user_data: Dict[str, Any] -) -> Dict[str, Any]: - """Create a new user and return the user info""" - url = "http://0.0.0.0:4000/user/new" - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - } - - async with session.post(url, headers=headers, json=user_data) as response: - assert ( - response.status == 200 - ), f"Failed to create user. Status: {response.status}" - user_info = await response.json() - return user_info - - -async def get_prometheus_metrics(session: aiohttp.ClientSession) -> str: - """Fetch current prometheus metrics""" - async with session.get("http://0.0.0.0:4000/metrics") as response: - assert response.status == 200 - return await response.text() - - -def extract_budget_metrics(metrics_text: str, team_id: str) -> Dict[str, float]: - """Extract budget-related metrics for a specific team""" - import re - - metrics = {} - - # Get remaining budget - remaining_pattern = f'litellm_remaining_team_budget_metric{{team="{team_id}",team_alias="[^"]*"}} ([0-9.]+)' - remaining_match = re.search(remaining_pattern, metrics_text) - metrics["remaining"] = float(remaining_match.group(1)) if remaining_match else None - - # Get total budget - total_pattern = f'litellm_team_max_budget_metric{{team="{team_id}",team_alias="[^"]*"}} ([0-9.]+)' - total_match = re.search(total_pattern, metrics_text) - metrics["total"] = float(total_match.group(1)) if total_match else None - - # Get remaining hours - hours_pattern = f'litellm_team_budget_remaining_hours_metric{{team="{team_id}",team_alias="[^"]*"}} ([0-9.]+)' - hours_match = re.search(hours_pattern, metrics_text) - metrics["remaining_hours"] = float(hours_match.group(1)) if hours_match else None - - return metrics - - -async def create_test_key(session: aiohttp.ClientSession, team_id: str) -> str: - """Generate a new key for the team and return it""" - url = "http://0.0.0.0:4000/key/generate" - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - } - data = { - "team_id": team_id, - } - - async with session.post(url, headers=headers, json=data) as response: - assert ( - response.status == 200 - ), f"Failed to generate key. Status: {response.status}" - key_info = await response.json() - return key_info["key"] - - -async def get_team_info(session: aiohttp.ClientSession, team_id: str) -> Dict[str, Any]: - """Fetch team info and return the response""" - url = f"http://0.0.0.0:4000/team/info?team_id={team_id}" - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - } - - async with session.get(url, headers=headers) as response: - assert ( - response.status == 200 - ), f"Failed to get team info. Status: {response.status}" - return await response.json() - - -@pytest.mark.asyncio -async def test_team_budget_metrics(): - """ - Test team budget tracking metrics: - 1. Create a team with max_budget - 2. Generate a key for the team - 3. Make chat completion requests using OpenAI SDK with team's key - 4. Verify budget decreases over time - 5. Verify request costs are being tracked correctly - 6. Verify prometheus metrics match /team/info spend data - """ - async with aiohttp.ClientSession() as session: - # Setup test team - team_data = { - "team_alias": "budget_test_team", - "max_budget": 10, - "budget_duration": "7d", - } - team_id = await create_test_team(session, team_data) - print("team_id", team_id) - # Generate key for the team - team_key = await create_test_key(session, team_id) - - # Initialize OpenAI client with team's key - client = AsyncOpenAI(base_url="http://0.0.0.0:4000", api_key=team_key) - - # Make initial request and check budget - await client.chat.completions.create( - model="fake-openai-endpoint", - messages=[{"role": "user", "content": f"Hello {uuid.uuid4()}"}], - ) - - await asyncio.sleep(11) # Wait for metrics to update - - # Get metrics after request - metrics_after_first = await get_prometheus_metrics(session) - print("metrics_after_first", metrics_after_first) - first_budget = extract_budget_metrics(metrics_after_first, team_id) - - print(f"Budget after 1 request: {first_budget}") - assert ( - first_budget["remaining"] < 10.0 - ), "remaining budget should be less than 10.0 after first request" - assert first_budget["total"] == 10.0, "Total budget metric is incorrect" - print("first_budget['remaining_hours']", first_budget["remaining_hours"]) - # Budget should have positive remaining hours, up to 7 days - assert ( - 0 < first_budget["remaining_hours"] <= 168 - ), "Budget should have positive remaining hours, up to 7 days" - - # Get team info and verify spend matches prometheus metrics - team_info = await get_team_info(session, team_id) - print("team_info", team_info) - _team_info_data = team_info["team_info"] - - # Calculate spend from prometheus (total - remaining) - team_info_spend = float(_team_info_data["spend"]) - team_info_max_budget = float(_team_info_data["max_budget"]) - team_info_remaining_budget = team_info_max_budget - team_info_spend - print("\n\n\n###### Final budget metrics ######\n\n\n") - print("team_info_remaining_budget", team_info_remaining_budget) - print("prometheus_remaining_budget", first_budget["remaining"]) - print( - "diff between team_info_remaining_budget and prometheus_remaining_budget", - team_info_remaining_budget - first_budget["remaining"], - ) - - # Verify spends match within a small delta (floating point comparison) - assert ( - abs(team_info_remaining_budget - first_budget["remaining"]) <= 0.001 - ), f"Spend mismatch: Prometheus={team_info_remaining_budget}, Team Info={first_budget['remaining']}" - - -async def create_test_key_with_budget( - session: aiohttp.ClientSession, budget_data: Dict[str, Any] -) -> str: - """Generate a new key with budget constraints and return it""" - url = "http://0.0.0.0:4000/key/generate" - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - } - print("budget_data", budget_data) - - async with session.post(url, headers=headers, json=budget_data) as response: - assert ( - response.status == 200 - ), f"Failed to generate key. Status: {response.status}" - key_info = await response.json() - return key_info["key"] - - -async def get_key_info(session: aiohttp.ClientSession, key: str) -> Dict[str, Any]: - """Fetch key info and return the response""" - url = "http://0.0.0.0:4000/key/info" - headers = { - "Authorization": f"Bearer {key}", - } - - async with session.get(url, headers=headers) as response: - assert ( - response.status == 200 - ), f"Failed to get key info. Status: {response.status}" - return await response.json() - - -async def get_user_info(session: aiohttp.ClientSession, user_id: str) -> Dict[str, Any]: - """Fetch user info and return the response""" - from urllib.parse import quote - - # URL encode user_id to handle special characters - encoded_user_id = quote(user_id, safe="") - url = f"http://0.0.0.0:4000/user/info?user_id={encoded_user_id}" - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - } - - async with session.get(url, headers=headers) as response: - assert ( - response.status == 200 - ), f"Failed to get user info. Status: {response.status}" - return await response.json() - - -def extract_key_budget_metrics(metrics_text: str, key_id: str) -> Dict[str, float]: - """Extract budget-related metrics for a specific key""" - import re - - metrics = {} - - # Get remaining budget - remaining_pattern = f'litellm_remaining_api_key_budget_metric{{api_key_alias="[^"]*",hashed_api_key="{key_id}"}} ([0-9.]+)' - remaining_match = re.search(remaining_pattern, metrics_text) - metrics["remaining"] = float(remaining_match.group(1)) if remaining_match else None - - # Get total budget - total_pattern = f'litellm_api_key_max_budget_metric{{api_key_alias="[^"]*",hashed_api_key="{key_id}"}} ([0-9.]+)' - total_match = re.search(total_pattern, metrics_text) - metrics["total"] = float(total_match.group(1)) if total_match else None - - # Get remaining hours - hours_pattern = f'litellm_api_key_budget_remaining_hours_metric{{api_key_alias="[^"]*",hashed_api_key="{key_id}"}} ([0-9.]+)' - hours_match = re.search(hours_pattern, metrics_text) - metrics["remaining_hours"] = float(hours_match.group(1)) if hours_match else None - - return metrics - - -def extract_user_budget_metrics(metrics_text: str, user_id: str) -> Dict[str, float]: - """Extract budget-related metrics for a specific user""" - import re - - metrics = {} - - # Escape user_id for regex pattern matching - escaped_user_id = re.escape(user_id) - - # Get remaining budget (user_email and user_alias may also be present as labels) - remaining_pattern = rf'litellm_remaining_user_budget_metric{{[^}}]*user="{escaped_user_id}"[^}}]*}} ([0-9.]+)' - remaining_match = re.search(remaining_pattern, metrics_text) - metrics["remaining"] = float(remaining_match.group(1)) if remaining_match else None - - # Get total budget - total_pattern = rf'litellm_user_max_budget_metric{{[^}}]*user="{escaped_user_id}"[^}}]*}} ([0-9.]+)' - total_match = re.search(total_pattern, metrics_text) - metrics["total"] = float(total_match.group(1)) if total_match else None - - # Get remaining hours - hours_pattern = rf'litellm_user_budget_remaining_hours_metric{{[^}}]*user="{escaped_user_id}"[^}}]*}} ([0-9.]+)' - hours_match = re.search(hours_pattern, metrics_text) - metrics["remaining_hours"] = float(hours_match.group(1)) if hours_match else None - - return metrics - - -@pytest.mark.asyncio -async def test_key_budget_metrics(): - """ - Test key budget tracking metrics: - 1. Create a key with max_budget - 2. Make chat completion requests using OpenAI SDK with the key - 3. Verify budget decreases over time - 4. Verify request costs are being tracked correctly - 5. Verify prometheus metrics match /key/info spend data - """ - from datetime import datetime, timedelta, timezone - - async with aiohttp.ClientSession() as session: - # Setup test key with unique alias - unique_alias = f"budget_test_key_{uuid.uuid4()}" - key_data = { - "key_alias": unique_alias, - "max_budget": 10, - "budget_duration": "7d", - "budget_reset_at": ( - datetime.now(timezone.utc) + timedelta(days=7) - ).isoformat(), - } - key = await create_test_key_with_budget(session, key_data) - - # Extract key_id from the key info - key_info = await get_key_info(session, key) - print("key_info", key_info) - key_id = key_info["key"] - print("key_id", key_id) - - # Initialize OpenAI client with the key - client = AsyncOpenAI(base_url="http://0.0.0.0:4000", api_key=key) - - # Make initial request and check budget - await client.chat.completions.create( - model="fake-openai-endpoint", - messages=[{"role": "user", "content": f"Hello {uuid.uuid4()}"}], - ) - - await asyncio.sleep(11) # Wait for metrics to update - - # Get metrics after request - metrics_after_first = await get_prometheus_metrics(session) - print("metrics_after_first request", metrics_after_first) - first_budget = extract_key_budget_metrics(metrics_after_first, key_id) - - print(f"Budget after 1 request: {first_budget}") - assert ( - first_budget["remaining"] < 10.0 - ), "remaining budget should be less than 10.0 after first request" - assert first_budget["total"] == 10.0, "Total budget metric is incorrect" - print("first_budget['remaining_hours']", first_budget["remaining_hours"]) - # The budget reset time is now standardized - for "7d" it resets on Monday at midnight - # So we'll check if it's within a reasonable range (0-7 days depending on current day of week) - assert ( - 0 <= first_budget["remaining_hours"] <= 168 - ), "Budget remaining hours should be within a reasonable range (0-7 days depending on day of week)" - - # Get key info and verify spend matches prometheus metrics - key_info = await get_key_info(session, key) - print("key_info", key_info) - _key_info_data = key_info["info"] - - # Calculate spend from prometheus (total - remaining) - key_info_spend = float(_key_info_data["spend"]) - key_info_max_budget = float(_key_info_data["max_budget"]) - key_info_remaining_budget = key_info_max_budget - key_info_spend - print("\n\n\n###### Final budget metrics ######\n\n\n") - print("key_info_remaining_budget", key_info_remaining_budget) - print("prometheus_remaining_budget", first_budget["remaining"]) - print( - "diff between key_info_remaining_budget and prometheus_remaining_budget", - key_info_remaining_budget - first_budget["remaining"], - ) - - # Verify spends match within a small delta (floating point comparison) - assert ( - abs(key_info_remaining_budget - first_budget["remaining"]) <= 0.001 - ), f"Spend mismatch: Prometheus={key_info_remaining_budget}, Key Info={first_budget['remaining']}" - - -@pytest.mark.asyncio -async def test_user_budget_metrics(): - """ - Test user budget tracking metrics: - 1. Create a user with max_budget - 2. Make chat completion requests using OpenAI SDK with the user's key - 3. Verify budget decreases over time - 4. Verify request costs are being tracked correctly - 5. Verify prometheus metrics match /user/info spend data - """ - from datetime import datetime, timedelta, timezone - - async with aiohttp.ClientSession() as session: - # Setup test user with unique user_id - unique_user_id = f"budget_test_user_{uuid.uuid4()}" - user_data = { - "user_id": unique_user_id, - "max_budget": 10, - "budget_duration": "7d", - "budget_reset_at": ( - datetime.now(timezone.utc) + timedelta(days=7) - ).isoformat(), - } - user_info = await create_test_user(session, user_data) - print("user_info", user_info) - user_id = user_info["user_id"] - print("user_id", user_id) - # Get the key that was created with the user - key = user_info["key"] - - # Initialize OpenAI client with the user's key - client = AsyncOpenAI(base_url="http://0.0.0.0:4000", api_key=key) - - # Make initial request and check budget - await client.chat.completions.create( - model="fake-openai-endpoint", - messages=[{"role": "user", "content": f"Hello {uuid.uuid4()}"}], - ) - - await asyncio.sleep(11) # Wait for metrics to update - - # Get metrics after request - metrics_after_first = await get_prometheus_metrics(session) - print("metrics_after_first request", metrics_after_first) - first_budget = extract_user_budget_metrics(metrics_after_first, user_id) - - print(f"Budget after 1 request: {first_budget}") - assert ( - first_budget["remaining"] is not None - ), "remaining budget metric should be present" - assert ( - first_budget["total"] is not None - ), "total budget metric should be present" - assert ( - first_budget["remaining"] < 10.0 - ), "remaining budget should be less than 10.0 after first request" - assert first_budget["total"] == 10.0, "Total budget metric is incorrect" - print("first_budget['remaining_hours']", first_budget["remaining_hours"]) - # The budget reset time is now standardized - for "7d" it resets on Monday at midnight - # So we'll check if it's within a reasonable range (0-7 days depending on current day of week) - assert ( - first_budget["remaining_hours"] is not None - ), "remaining hours metric should be present" - assert ( - 0 <= first_budget["remaining_hours"] <= 168 - ), "Budget remaining hours should be within a reasonable range (0-7 days depending on day of week)" - - # Get user info and verify spend matches prometheus metrics - user_info_response = await get_user_info(session, user_id) - print("user_info_response", user_info_response) - _user_info_data = user_info_response["user_info"] - - # Calculate spend from prometheus (total - remaining) - user_info_spend = float(_user_info_data["spend"]) - user_info_max_budget = float(_user_info_data["max_budget"]) - user_info_remaining_budget = user_info_max_budget - user_info_spend - print("\n\n\n###### Final budget metrics ######\n\n\n") - print("user_info_remaining_budget", user_info_remaining_budget) - print("prometheus_remaining_budget", first_budget["remaining"]) - print( - "diff between user_info_remaining_budget and prometheus_remaining_budget", - user_info_remaining_budget - first_budget["remaining"], - ) - - # Verify spends match within a small delta (floating point comparison) - assert ( - abs(user_info_remaining_budget - first_budget["remaining"]) <= 0.001 - ), f"Spend mismatch: Prometheus={user_info_remaining_budget}, User Info={first_budget['remaining']}" - - -@pytest.mark.asyncio -async def test_user_email_metrics(): - """ - Test user email tracking metrics: - 1. Create a user with user_email - 2. Make chat completion requests using OpenAI SDK with the user's email - 3. Verify user email is being tracked correctly in `litellm_user_email_metric` - """ - async with aiohttp.ClientSession() as session: - # Create a user with user_email - user_email = f"test-{uuid.uuid4()}@example.com" - user_data = { - "user_email": user_email, - } - user_info = await create_test_user(session, user_data) - key = user_info["key"] - - # Initialize OpenAI client with the user's email - client = AsyncOpenAI(base_url="http://0.0.0.0:4000", api_key=key) - - # Make initial request and check budget - await client.chat.completions.create( - model="fake-openai-endpoint", - messages=[{"role": "user", "content": f"Hello {uuid.uuid4()}"}], - ) - - await asyncio.sleep(11) # Wait for metrics to update - - # Get metrics after request - metrics_after_first = await get_prometheus_metrics(session) - print("metrics_after_first request", metrics_after_first) - assert ( - user_email in metrics_after_first - ), "user_email should be tracked correctly" - - -@pytest.mark.asyncio -async def test_user_email_in_all_required_metrics(): - """ - Test that user_email label is present in all the metrics that were requested to have it: - - litellm_proxy_total_requests_metric_total - - litellm_proxy_failed_requests_metric_total - - litellm_input_tokens_metric_total - - litellm_output_tokens_metric_total - - litellm_requests_metric_total - - litellm_spend_metric_total - """ - async with aiohttp.ClientSession() as session: - # Create a user with user_email - user_email = f"test-metrics-{uuid.uuid4()}@example.com" - user_data = { - "user_email": user_email, - } - user_info = await create_test_user(session, user_data) - key = user_info["key"] - - # Initialize OpenAI client with the user's email - client = AsyncOpenAI(base_url="http://0.0.0.0:4000", api_key=key) - - # Make successful request to generate metrics - await client.chat.completions.create( - model="fake-openai-endpoint", - messages=[{"role": "user", "content": f"Hello {uuid.uuid4()}"}], - ) - - await asyncio.sleep(11) # Wait for metrics to update - - # Get metrics after request - metrics_text = await get_prometheus_metrics(session) - print("Testing user_email in all required metrics") - - # Check that user_email appears in all the required metrics - required_metrics_with_user_email = [ - # "litellm_proxy_total_requests_metric_total", - # "litellm_input_tokens_metric_total", - # "litellm_output_tokens_metric_total", - # "litellm_requests_metric_total", - "litellm_spend_metric_total", - ] - - import re - - for metric_name in required_metrics_with_user_email: - # Check that the metric exists and contains user_email label - # Look for the metric with user_email in its labels - pattern = ( - rf'{metric_name}{{[^}}]*user_email="{re.escape(user_email)}"[^}}]*}}' - ) - matches = re.findall(pattern, metrics_text) - assert ( - len(matches) > 0 - ), f"Metric {metric_name} should contain user_email={user_email} but was not found in metrics" - - # Also test failure metric by making a bad request - try: - await client.chat.completions.create( - model="fake-azure-endpoint", # This should fail - messages=[{"role": "user", "content": "Hello"}], - ) - except Exception: - pass # Expected to fail - - await asyncio.sleep(11) # Wait for metrics to update - - # Get updated metrics - metrics_text = await get_prometheus_metrics(session) - - # Check that failure metric also contains user_email - failure_pattern = rf'litellm_proxy_failed_requests_metric_total{{[^}}]*user_email="{re.escape(user_email)}"[^}}]*}}' - failure_matches = re.findall(failure_pattern, metrics_text) - assert ( - len(failure_matches) > 0 - ), f"litellm_proxy_failed_requests_metric_total should contain user_email={user_email}" diff --git a/tests/otel_tests/test_team_tag_routing.py b/tests/otel_tests/test_team_tag_routing.py deleted file mode 100644 index b818aa0cea4..00000000000 --- a/tests/otel_tests/test_team_tag_routing.py +++ /dev/null @@ -1,65 +0,0 @@ -import os -# What this tests ? -## Set tags on a team and then make a request to /chat/completions -import pytest -import asyncio -import aiohttp, openai -from openai import OpenAI, AsyncOpenAI -from typing import Optional, List, Union -from litellm._uuid import uuid - -LITELLM_MASTER_KEY = os.environ["LITELLM_MASTER_KEY"] - - -async def chat_completion( - session, key, model: Union[str, List] = "fake-openai-endpoint" -): - url = "http://0.0.0.0:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - print("headers=", headers) - data = { - "model": model, - "messages": [ - {"role": "user", "content": f"Hello! {str(uuid.uuid4())}"}, - ], - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - if status != 200: - raise Exception(response_text) - - return await response.json(), response.headers - - -async def model_info_get_call(session, key, model_id: str): - # make get call pass "litellm_model_id" in query params - url = f"http://0.0.0.0:4000/model/info?litellm_model_id={model_id}" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - - if status != 200: - raise Exception(response_text) - - return await response.json() - - -@pytest.mark.asyncio() -async def test_chat_completion_with_no_tags(): - async with aiohttp.ClientSession() as session: - key = LITELLM_MASTER_KEY - response, headers = await chat_completion(session, key) - headers = dict(headers) - print(response) - print(headers) - assert response is not None diff --git a/tests/pass_through_tests/base_anthropic_messages_test.py b/tests/pass_through_tests/base_anthropic_messages_test.py index 95f709e3880..710dfa453cf 100644 --- a/tests/pass_through_tests/base_anthropic_messages_test.py +++ b/tests/pass_through_tests/base_anthropic_messages_test.py @@ -13,61 +13,8 @@ class BaseAnthropicMessagesTest(ABC): def get_client(self): return anthropic.Anthropic() - def test_anthropic_basic_completion(self): - print("making basic completion request to anthropic passthrough") - client = self.get_client() - response = client.messages.create( - model="claude-sonnet-4-5-20250929", - max_tokens=1024, - messages=[{"role": "user", "content": "Say 'hello test' and nothing else"}], - extra_body={ - "litellm_metadata": { - "tags": ["test-tag-1", "test-tag-2"], - } - }, - ) - print(response) - def test_anthropic_streaming(self): - print("making streaming request to anthropic passthrough") - collected_output = [] - client = self.get_client() - with client.messages.stream( - max_tokens=10, - messages=[ - {"role": "user", "content": "Say 'hello stream test' and nothing else"} - ], - model="claude-sonnet-4-5-20250929", - extra_body={ - "litellm_metadata": { - "tags": ["test-tag-stream-1", "test-tag-stream-2"], - } - }, - ) as stream: - for text in stream.text_stream: - collected_output.append(text) - full_response = "".join(collected_output) - print(full_response) - - def test_anthropic_messages_with_thinking(self): - print("making request to anthropic passthrough with thinking") - client = self.get_client() - response = client.messages.create( - model="claude-haiku-4-5-20251001", - max_tokens=20000, - thinking={"type": "enabled", "budget_tokens": 16000}, - messages=[ - {"role": "user", "content": "Just pinging with thinking enabled"} - ], - ) - - print(response) - - # Verify the first content block is a thinking block - response_thinking = response.content[0].thinking - assert response_thinking is not None - assert len(response_thinking) > 0 def test_anthropic_streaming_with_thinking(self): print("making streaming request to anthropic passthrough with thinking enabled") @@ -105,41 +52,4 @@ class BaseAnthropicMessagesTest(ABC): assert len(collected_response) > 0 assert len(full_response) > 0 - def test_bad_request_error_handling_streaming(self): - print("making request to anthropic passthrough with bad request") - try: - client = self.get_client() - response = client.messages.create( - model="claude-sonnet-4-5-20250929", - max_tokens=10, - stream=True, - messages=["hi"], - ) - print(response) - assert pytest.fail("Expected BadRequestError") - except anthropic.BadRequestError as e: - print("Got BadRequestError from anthropic, e=", e) - print(e.__cause__) - print(e.status_code) - print(e.response) - except Exception as e: - pytest.fail(f"Got unexpected exception: {e}") - def test_bad_request_error_handling_non_streaming(self): - print("making request to anthropic passthrough with bad request") - try: - client = self.get_client() - response = client.messages.create( - model="claude-sonnet-4-5-20250929", - max_tokens=10, - messages=["hi"], - ) - print(response) - assert pytest.fail("Expected BadRequestError") - except anthropic.BadRequestError as e: - print("Got BadRequestError from anthropic, e=", e) - print(e.__cause__) - print(e.status_code) - print(e.response) - except Exception as e: - pytest.fail(f"Got unexpected exception: {e}") diff --git a/tests/pass_through_tests/test_anthropic_passthrough.py b/tests/pass_through_tests/test_anthropic_passthrough.py deleted file mode 100644 index 5d6ddb1fbd0..00000000000 --- a/tests/pass_through_tests/test_anthropic_passthrough.py +++ /dev/null @@ -1,472 +0,0 @@ -""" -This test ensures that the proxy can passthrough anthropic requests -""" - -import os -import pytest -import anthropic -import aiohttp -import asyncio -import json - - -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=2) -async def test_anthropic_basic_completion_with_headers(): - print("making basic completion request to anthropic passthrough with aiohttp") - - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - "Anthropic-Version": "2023-06-01", - } - - payload = { - "model": "claude-sonnet-4-5-20250929", - "max_tokens": 10, - "messages": [{"role": "user", "content": "Say 'hello test' and nothing else"}], - "litellm_metadata": { - "tags": ["test-tag-1", "test-tag-2"], - }, - } - - async with aiohttp.ClientSession() as session: - async with session.post( - "http://0.0.0.0:4000/anthropic/v1/messages", json=payload, headers=headers - ) as response: - response_text = await response.text() - print(f"Response text: {response_text}") - - response_json = await response.json() - response_headers = response.headers - print( - "non-streaming response", - json.dumps(response_json, indent=4, default=str), - ) - reported_usage = response_json.get("usage", None) - # fix null checks for reported_usage - anthropic_api_input_tokens = ( - reported_usage.get("input_tokens", None) if reported_usage else None - ) - anthropic_api_output_tokens = ( - reported_usage.get("output_tokens", None) if reported_usage else None - ) - anthropic_message_id = response_json.get("id") - - print(f"Anthropic message ID: {anthropic_message_id}") - - # Wait for spend to be logged - await asyncio.sleep(15) - - # Check spend logs for this specific request with retry logic - spend_data = None - max_retries = 2 - for attempt in range(max_retries): - print(f"Attempt {attempt + 1}/{max_retries} to check spend logs") - - async with session.get( - f"http://0.0.0.0:4000/spend/logs?request_id={anthropic_message_id}", - headers={"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}, - ) as spend_response: - print("text spend response") - print(f"Spend response: {spend_response}") - spend_data = await spend_response.json() - print(f"Spend data: {spend_data}") - - # Check if spend data exists and has entries - if spend_data and len(spend_data) > 0: - print("Spend logs found!") - break - else: - print("Spend logs not found yet...") - if ( - attempt < max_retries - 1 - ): # Don't wait after the last attempt - print("Waiting 10 seconds before retry...") - await asyncio.sleep(10) - - if not isinstance(spend_data, list): - print(f"Spend endpoint answered with an error response: {spend_data}") - print("Skipping spend assertions (spend logs unreachable in CI)") - return - - assert spend_data, ( - f"GET /spend/logs?request_id={anthropic_message_id} found no row for the id " - "the caller received" - ) - - log_entry = spend_data[0] - - # Basic existence checks - assert isinstance(log_entry, dict), "Log entry should be a dictionary" - - # Request metadata assertions - assert ( - log_entry["request_id"] == anthropic_message_id - ), "Request ID should be the message id the caller received" - assert ( - log_entry["call_type"] == "pass_through_endpoint" - ), "Call type should be pass_through_endpoint" - assert ( - log_entry["api_base"] == "https://api.anthropic.com/v1/messages" - ), "API base should be Anthropic's endpoint" - - # Token and spend assertions - assert log_entry["spend"] > 0, "Spend value should not be None" - assert isinstance( - log_entry["spend"], (int, float) - ), "Spend should be a number" - assert log_entry["total_tokens"] > 0, "Should have some tokens" - assert ( - log_entry["prompt_tokens"] == anthropic_api_input_tokens - ), f"Should have prompt tokens matching anthropic api. Expected {anthropic_api_input_tokens} but got {log_entry['prompt_tokens']}" - assert ( - log_entry["completion_tokens"] == anthropic_api_output_tokens - ), f"Should have completion tokens matching anthropic api. Expected {anthropic_api_output_tokens} but got {log_entry['completion_tokens']}" - assert ( - log_entry["total_tokens"] - == log_entry["prompt_tokens"] + log_entry["completion_tokens"] - ), "Total tokens should equal prompt + completion" - - # Time assertions - assert all( - key in log_entry - for key in ["startTime", "endTime", "completionStartTime"] - ), "Should have all time fields" - assert ( - log_entry["startTime"] < log_entry["endTime"] - ), "Start time should be before end time" - - # Metadata assertions - assert str(log_entry["cache_hit"]).lower() != "true", "Cache should be off" - assert log_entry["request_tags"] == [ - "test-tag-1", - "test-tag-2", - ], "Tags should match input" - assert ( - "user_api_key" in log_entry["metadata"] - ), "Should have user API key in metadata" - - assert "claude" in log_entry["model"] - assert log_entry["custom_llm_provider"] == "anthropic" - - -@pytest.mark.asyncio -async def test_anthropic_streaming_with_headers(): - print("making streaming request to anthropic passthrough with aiohttp") - - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - "Anthropic-Version": "2023-06-01", - } - - payload = { - "model": "claude-sonnet-4-5-20250929", - "max_tokens": 10, - "messages": [ - {"role": "user", "content": "Say 'hello stream test' and nothing else"} - ], - "stream": True, - "litellm_metadata": { - "tags": ["test-tag-stream-1", "test-tag-stream-2"], - "user": "test-user-1", - }, - } - - async with aiohttp.ClientSession() as session: - async with session.post( - "http://0.0.0.0:4000/anthropic/v1/messages", json=payload, headers=headers - ) as response: - print("response status") - print(response.status) - assert response.status == 200, "Response should be successful" - response_headers = response.headers - print(f"Response headers: {response_headers}") - - collected_output = [] - async for line in response.content: - if line: - text = line.decode("utf-8").strip() - if text.startswith("data: "): - collected_output.append(text[6:]) # Remove 'data: ' prefix - - print("Collected output:", "".join(collected_output)) - anthropic_api_usage_chunks = [] - anthropic_message_id = None - for chunk in collected_output: - chunk_json = json.loads(chunk) - if chunk_json.get("type") == "message_start": - anthropic_message_id = chunk_json.get("message", {}).get("id") - if "usage" in chunk_json: - anthropic_api_usage_chunks.append(chunk_json["usage"]) - elif "message" in chunk_json and "usage" in chunk_json["message"]: - anthropic_api_usage_chunks.append(chunk_json["message"]["usage"]) - - print(f"Anthropic message ID: {anthropic_message_id}") - - print( - "anthropic_api_usage_chunks", - json.dumps(anthropic_api_usage_chunks, indent=4, default=str), - ) - - print("anthropic_api_usage_chunks: ", anthropic_api_usage_chunks) - # Get the most recent value of input tokens (iterate backwards to find last non-zero value) - anthropic_api_input_tokens = 0 - for usage in reversed(anthropic_api_usage_chunks): - if usage.get("input_tokens", 0) > 0: - anthropic_api_input_tokens = usage.get("input_tokens", 0) - break - anthropic_api_output_tokens = 0 - for usage in reversed(anthropic_api_usage_chunks): - if usage.get("output_tokens", 0) > 0: - anthropic_api_output_tokens = usage.get("output_tokens", 0) - break - - print("anthropic_api_input_tokens", anthropic_api_input_tokens) - print("anthropic_api_output_tokens", anthropic_api_output_tokens) - - # Wait for spend to be logged - await asyncio.sleep(20) - - # Check spend logs for this specific request with retry logic - spend_data = None - max_retries = 2 - for attempt in range(max_retries): - print(f"Attempt {attempt + 1}/{max_retries} to check spend logs") - - async with session.get( - f"http://0.0.0.0:4000/spend/logs?request_id={anthropic_message_id}", - headers={"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}, - ) as spend_response: - spend_data = await spend_response.json() - print(f"Spend data: {spend_data}") - - # Check if spend data exists and has entries - if spend_data and len(spend_data) > 0: - print("Spend logs found!") - break - else: - print("Spend logs not found yet...") - if ( - attempt < max_retries - 1 - ): # Don't wait after the last attempt - print("Waiting 10 seconds before retry...") - await asyncio.sleep(10) - - if not isinstance(spend_data, list): - print(f"Spend endpoint answered with an error response: {spend_data}") - print("Skipping spend assertions (spend logs unreachable in CI)") - return - - assert spend_data, ( - f"GET /spend/logs?request_id={anthropic_message_id} found no row for the id " - "the caller received" - ) - - log_entry = spend_data[0] - - # Basic existence checks - assert isinstance(log_entry, dict), "Log entry should be a dictionary" - - # Request metadata assertions - assert ( - log_entry["request_id"] == anthropic_message_id - ), "Request ID should be the message id the caller received" - assert ( - log_entry["call_type"] == "pass_through_endpoint" - ), "Call type should be pass_through_endpoint" - # assert ( - # log_entry["api_base"] == "https://api.anthropic.com/v1/messages" - # ), "API base should be Anthropic's endpoint" - - # Token and spend assertions - assert log_entry["spend"] > 0, "Spend value should not be None" - assert isinstance( - log_entry["spend"], (int, float) - ), "Spend should be a number" - assert log_entry["total_tokens"] > 0, "Should have some tokens" - assert ( - log_entry["prompt_tokens"] == anthropic_api_input_tokens - ), f"Should have prompt tokens matching anthropic api. Expected {anthropic_api_input_tokens} but got {log_entry['prompt_tokens']}" - assert ( - log_entry["completion_tokens"] == anthropic_api_output_tokens - ), f"Should have completion tokens matching anthropic api. Expected {anthropic_api_output_tokens} but got {log_entry['completion_tokens']}" - assert ( - log_entry["total_tokens"] - == log_entry["prompt_tokens"] + log_entry["completion_tokens"] - ), "Total tokens should equal prompt + completion" - - # Time assertions - assert all( - key in log_entry - for key in ["startTime", "endTime", "completionStartTime"] - ), "Should have all time fields" - assert ( - log_entry["startTime"] < log_entry["endTime"] - ), "Start time should be before end time" - - # Metadata assertions - assert str(log_entry["cache_hit"]).lower() != "true", "Cache should be off" - assert log_entry["request_tags"] == [ - "test-tag-stream-1", - "test-tag-stream-2", - ], "Tags should match input" - assert ( - "user_api_key" in log_entry["metadata"] - ), "Should have user API key in metadata" - - assert "claude" in log_entry["model"] - - assert log_entry["end_user"] == "test-user-1" - assert log_entry["custom_llm_provider"] == "anthropic" - - -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=2) -async def test_anthropic_messages_streaming_cost_injection(): - """ - Test that cost is injected into message_delta usage for Anthropic Messages API streaming - """ - print("Testing cost injection in Anthropic Messages API streaming response") - - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - "anthropic-version": "2023-06-01", - } - - payload = { - "model": "claude-haiku-4-5-20251001", - "max_tokens": 10, - "stream": True, - "messages": [{"role": "user", "content": "Say 'Hi'"}], - } - - async with aiohttp.ClientSession() as session: - async with session.post( - "http://0.0.0.0:4000/v1/messages", - json=payload, - headers=headers, - ) as response: - assert response.status == 200 - - # Collect all SSE events. - # Split each chunk by newlines to handle both: - # - Anthropic direct path: chunks arrive as individual lines - # - OpenAI/Responses API path: chunks are full multi-line SSE events - events = [] - async for chunk in response.content: - chunk_str = chunk.decode("utf-8") - for line in chunk_str.split("\n"): - line = line.strip() - if line.startswith("data: "): - try: - data = json.loads(line[6:]) # Remove 'data: ' prefix - events.append(data) - except json.JSONDecodeError: - continue - - # Find message_delta event with usage - message_delta_events = [ - event - for event in events - if event.get("type") == "message_delta" and "usage" in event - ] - - assert ( - len(message_delta_events) > 0 - ), "No message_delta events with usage found" - - # Check that cost is included in usage - for event in message_delta_events: - usage = event.get("usage", {}) - assert "cost" in usage, f"Cost not found in usage: {usage}" - assert isinstance( - usage["cost"], (int, float) - ), f"Cost should be numeric: {usage['cost']}" - assert ( - usage["cost"] >= 0 - ), f"Cost should be non-negative: {usage['cost']}" - - print(f"Found message_delta with cost: {usage}") - - print( - f"Test passed: Found {len(message_delta_events)} message_delta events with cost" - ) - - -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=2) -async def test_anthropic_messages_openai_model_streaming_cost_injection(): - """ - Test that cost is injected into message_delta usage for OpenAI model via Anthropic Messages API - """ - print("Testing cost injection in Anthropic Messages API with OpenAI model") - - headers = { - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - "Content-Type": "application/json", - "anthropic-version": "2023-06-01", - } - - payload = { - "model": "openai/gpt-4o", - "max_tokens": 20, - "stream": True, - "messages": [{"role": "user", "content": "Say 'Hi'"}], - } - - async with aiohttp.ClientSession() as session: - async with session.post( - "http://0.0.0.0:4000/v1/messages", - json=payload, - headers=headers, - ) as response: - assert response.status == 200 - - # Collect all SSE events. - # Split each chunk by newlines to handle both: - # - Direct API paths: chunks arrive as individual lines - # - OpenAI/Responses API path: AnthropicResponsesStreamWrapper yields - # full multi-line SSE events as single bytes objects, so a naive - # startswith('data: ') check on the whole chunk misses them. - events = [] - async for chunk in response.content: - chunk_str = chunk.decode("utf-8") - for line in chunk_str.split("\n"): - line = line.strip() - if line.startswith("data: "): - try: - data = json.loads(line[6:]) # Remove 'data: ' prefix - events.append(data) - except json.JSONDecodeError: - continue - - # Find message_delta event with usage - message_delta_events = [ - event - for event in events - if event.get("type") == "message_delta" and "usage" in event - ] - - assert ( - len(message_delta_events) > 0 - ), "No message_delta events with usage found" - - # Check that cost is included in usage - for event in message_delta_events: - usage = event.get("usage", {}) - assert "cost" in usage, f"Cost not found in usage: {usage}" - assert isinstance( - usage["cost"], (int, float) - ), f"Cost should be numeric: {usage['cost']}" - assert ( - usage["cost"] >= 0 - ), f"Cost should be non-negative: {usage['cost']}" - - print(f"Found message_delta with cost: {usage}") - - print( - f"Test passed: Found {len(message_delta_events)} message_delta events with cost" - ) diff --git a/tests/pass_through_tests/test_anthropic_passthrough_basic.py b/tests/pass_through_tests/test_anthropic_passthrough_basic.py index c7e9fea867c..4ef9887bc89 100644 --- a/tests/pass_through_tests/test_anthropic_passthrough_basic.py +++ b/tests/pass_through_tests/test_anthropic_passthrough_basic.py @@ -19,11 +19,3 @@ class TestAnthropicMessagesEndpoint(BaseAnthropicMessagesTest): api_key=os.environ["LITELLM_MASTER_KEY"], ) - def test_anthropic_messages_to_wildcard_model(self): - client = self.get_client() - response = client.messages.create( - model="anthropic/claude-haiku-4-5-20251001", - messages=[{"role": "user", "content": "Hello, world!"}], - max_tokens=100, - ) - print(response) diff --git a/tests/pass_through_tests/test_assembly_ai.py b/tests/pass_through_tests/test_assembly_ai.py deleted file mode 100644 index 09999bc2bed..00000000000 --- a/tests/pass_through_tests/test_assembly_ai.py +++ /dev/null @@ -1,102 +0,0 @@ -""" -This test ensures that the proxy can passthrough requests to assemblyai -""" - -import os -import time - -import pytest -import httpx -import aiohttp -import asyncio - -TEST_MASTER_KEY = os.environ["LITELLM_MASTER_KEY"] -TEST_BASE_URL = "http://0.0.0.0:4000/assemblyai" - - -def _transcribe_and_verify(virtual_key: str, base_url: str): - file_url = "https://assembly.ai/wildfires.mp3" - headers = { - "Authorization": f"Bearer {virtual_key}", - "Content-Type": "application/json", - } - create_payload = { - "audio_url": file_url, - "speech_models": ["universal-2"], - } - - create_response = httpx.post( - url=f"{base_url}/v2/transcript", - headers=headers, - json=create_payload, - timeout=60.0, - ) - if create_response.status_code != 200: - pytest.fail( - "Failed to create transcript request: " - f"status={create_response.status_code}, body={create_response.text}" - ) - - transcript = create_response.json() - transcript_id = transcript.get("id") - if not transcript_id: - pytest.fail("Failed to get transcript id") - - for _ in range(60): - poll_response = httpx.get( - url=f"{base_url}/v2/transcript/{transcript_id}", - headers=headers, - timeout=30.0, - ) - if poll_response.status_code != 200: - pytest.fail( - "Failed to poll transcript status: " - f"status={poll_response.status_code}, body={poll_response.text}" - ) - transcript = poll_response.json() - if transcript.get("status") in ("completed", "error"): - break - time.sleep(1) - - httpx.delete( - url=f"{base_url}/v2/transcript/{transcript_id}", - headers=headers, - timeout=30.0, - ) - - if transcript.get("status") == "error": - pytest.fail(f"Failed to transcribe file error: {transcript.get('error')}") - - print(transcript.get("text")) - - -def test_assemblyai_basic_transcribe(): - print("making basic transcribe request to assemblyai passthrough") - _transcribe_and_verify(TEST_MASTER_KEY, TEST_BASE_URL) - - -async def generate_key(calling_key: str) -> str: - """Helper function to generate a new API key""" - url = "http://0.0.0.0:4000/key/generate" - headers = { - "Authorization": f"Bearer {calling_key}", - "Content-Type": "application/json", - } - - async with aiohttp.ClientSession() as session: - async with session.post(url, headers=headers, json={}) as response: - if response.status == 200: - data = await response.json() - return data.get("key") - raise Exception(f"Failed to generate key: {response.status}") - - -@pytest.mark.asyncio -async def test_assemblyai_transcribe_with_non_admin_key(): - non_admin_key = await generate_key(TEST_MASTER_KEY) - print(f"Generated non-admin key: {non_admin_key}") - - request_start_time = time.time() - _transcribe_and_verify(non_admin_key, TEST_BASE_URL) - request_end_time = time.time() - print(f"Request took {request_end_time - request_start_time} seconds") diff --git a/tests/pass_through_tests/test_gemini_with_spend.test.js b/tests/pass_through_tests/test_gemini_with_spend.test.js index f59d575c033..f184bbe14e3 100644 --- a/tests/pass_through_tests/test_gemini_with_spend.test.js +++ b/tests/pass_through_tests/test_gemini_with_spend.test.js @@ -24,57 +24,6 @@ global.fetch = async function patchedFetch(url, options) { jest.retryTimes(3); describe('Gemini AI Tests', () => { - test('should successfully generate non-streaming content with tags', async () => { - const genAI = new GoogleGenerativeAI(masterKey); - - const requestOptions = { - baseUrl: 'http://127.0.0.1:4000/gemini', - customHeaders: { - "tags": "gemini-js-sdk,pass-through-endpoint" - } - }; - - const model = genAI.getGenerativeModel({ - model: 'gemini-3.1-flash-lite' - }, requestOptions); - - const prompt = 'Say "hello test" and nothing else'; - - const result = await model.generateContent(prompt); - expect(result).toBeDefined(); - - // Use the captured callId - const callId = lastCallId; - console.log("Captured Call ID:", callId); - - // Poll for spend data with retries (DB writes can be slow in CI) - let spendData = null; - for (let attempt = 0; attempt < 6; attempt++) { - await new Promise(resolve => setTimeout(resolve, 10000)); - const spendResponse = await fetch( - `http://127.0.0.1:4000/spend/logs?request_id=${callId}`, - { headers: { 'Authorization': `Bearer ${masterKey}` } } - ); - spendData = await spendResponse.json(); - console.log(`spendData (attempt ${attempt + 1}):`, spendData); - if (spendData && spendData.length > 0 && spendData[0] && spendData[0].request_id) break; - } - - if (!spendData || !spendData.length || !spendData[0] || !spendData[0].request_id) { - console.warn('Spend data not available after polling - skipping spend assertions (DB write may be slow in CI)'); - return; - } - - expect(spendData).toBeDefined(); - expect(spendData[0].request_id).toBe(callId); - expect(spendData[0].call_type).toBe('pass_through_endpoint'); - expect(spendData[0].request_tags).toEqual(['gemini-js-sdk', 'pass-through-endpoint']); - expect(spendData[0].metadata).toHaveProperty('user_api_key'); - expect(spendData[0].model).toContain('gemini'); - expect(spendData[0].custom_llm_provider).toBe('gemini'); - expect(spendData[0].spend).toBeGreaterThan(0); - }, 90000); - test('should successfully generate streaming content with tags', async () => { const genAI = new GoogleGenerativeAI(masterKey); diff --git a/tests/pass_through_tests/test_hosted_vllm_passthrough.py b/tests/pass_through_tests/test_hosted_vllm_passthrough.py deleted file mode 100644 index 272b4e1bb00..00000000000 --- a/tests/pass_through_tests/test_hosted_vllm_passthrough.py +++ /dev/null @@ -1,71 +0,0 @@ -import asyncio -from unittest.mock import AsyncMock, patch - -import httpx -import pytest - -from litellm.passthrough.main import allm_passthrough_route -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler -from litellm.utils import ProviderConfigManager -from litellm.types.utils import LlmProviders -from litellm.llms.vllm.passthrough.transformation import ( - VLLMPassthroughConfig, -) - - -def test_get_provider_passthrough_config_for_hosted_vllm_returns_vllm_config(): - # When requesting passthrough config for HOSTED_VLLM - cfg = ProviderConfigManager.get_provider_passthrough_config( - model="hosted_vllm/my-deployment", - provider=LlmProviders.HOSTED_VLLM, - ) - - # Then we should get a VLLMPassthroughConfig instance - assert isinstance(cfg, VLLMPassthroughConfig) - - -@pytest.mark.asyncio -async def test_allm_passthrough_route_with_hosted_vllm_model_does_not_raise(): - # Given a hosted_vllm model and an async http client - client = AsyncHTTPHandler() - - # Mock the provider resolution to ensure we use hosted_vllm and provide api_base - with patch( - "litellm.passthrough.main.get_llm_provider", - return_value=( - "my-deployment", # normalized model name - "hosted_vllm", # provider - "fake-api-key", # api key (not required for vllm) - "http://localhost:8090", # api base - ), - ): - # Mock the underlying AsyncClient.send to avoid real network I/O - fake_request = httpx.Request( - method="POST", url="http://localhost:8090/v1/chat/completions" - ) - fake_response = httpx.Response( - status_code=200, - content=b'{\n "ok": true\n}', - request=fake_request, - headers={"content-type": "application/json"}, - ) - - with patch.object( - client.client, "send", new=AsyncMock(return_value=fake_response) - ): - # When calling the async passthrough route with a hosted_vllm/* model - response = await allm_passthrough_route( - method="POST", - endpoint="v1/chat/completions", - model="hosted_vllm/my-deployment", - api_base="http://localhost:8090", - json={ - "model": "anything", # will be replaced internally with normalized model - "messages": [{"role": "user", "content": "Hello"}], - }, - client=client, - ) - - # Then it should not raise and return a successful httpx.Response - assert isinstance(response, httpx.Response) - assert response.status_code == 200 diff --git a/tests/pass_through_tests/test_openai_assistants_passthrough.py b/tests/pass_through_tests/test_openai_assistants_passthrough.py deleted file mode 100644 index 4da84ce5ca3..00000000000 --- a/tests/pass_through_tests/test_openai_assistants_passthrough.py +++ /dev/null @@ -1,23 +0,0 @@ -import os -import openai -import tempfile - - -client = openai.OpenAI(base_url="http://0.0.0.0:4000/openai", api_key=os.environ["LITELLM_MASTER_KEY"]) - - -def test_pass_through_file_operations(): - with tempfile.NamedTemporaryFile( - mode="w+", suffix=".txt", delete=False - ) as temp_file: - temp_file.write("This is a test file for the OpenAI Assistants API.") - temp_file.flush() - - file = client.files.create( - file=open(temp_file.name, "rb"), - purpose="assistants", - ) - print("file created", file) - - delete_file = client.files.delete(file.id) - print("file deleted", delete_file) diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index c1de9ae777d..834373df650 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -82,93 +82,6 @@ def get_tracked_spend() -> float: return sum(float(row.get("spend") or 0.0) for row in rows) -VERTEX_PROJECT = "litellm-ci-cd" -VERTEX_MODEL = "gemini-3.1-flash-lite" -VERTEX_GENERATE_CONTENT_URL = ( - f"{LITE_LLM_ENDPOINT}/vertex_ai/v1/projects/{VERTEX_PROJECT}" - f"/locations/global/publishers/google/models/{VERTEX_MODEL}:generateContent" -) - - -def _vertex_access_token() -> str: - import google.auth - import google.auth.transport.requests - - credentials, _ = google.auth.default( - scopes=["https://www.googleapis.com/auth/cloud-platform"] - ) - credentials.refresh(google.auth.transport.requests.Request()) - return credentials.token - - -def _spend_log_for_request(call_id: str) -> dict | None: - response = requests.get( - f"{LITE_LLM_ENDPOINT}/spend/logs?request_id={call_id}", - headers={"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}, - timeout=30, - ) - if response.status_code != 200: - return None - rows = response.json() - return rows[0] if rows else None - - -def _is_vertex_quota_error(response: requests.Response) -> bool: - return response.status_code == 429 or "RESOURCE_EXHAUSTED" in response.text - - -@pytest.mark.asyncio() -async def test_basic_vertex_ai_pass_through_with_spendlog(): - load_vertex_ai_credentials() - access_token = _vertex_access_token() - - # Drive the pass-through over HTTP instead of the vertexai SDK: the SDK intermittently - # routes generateContent to the public Vertex endpoint rather than the proxy override, - # so the call never reaches LiteLLM and no spend is logged. A direct request always - # hits the proxy. Spend logging then runs on a best-effort background worker that can - # drop a single event, so retry a few billed calls and assert that one specific call's - # spend log lands. Failing every attempt still fails hard, which is the signal we want - # if cost tracking is broken. - max_attempts = 3 - poll_seconds = 60 - poll_interval = 5 - - for attempt in range(1, max_attempts + 1): - response = requests.post( - VERTEX_GENERATE_CONTENT_URL, - headers={ - "Authorization": f"Bearer {access_token}", - "Content-Type": "application/json", - }, - json={"contents": [{"role": "user", "parts": [{"text": "hi"}]}]}, - timeout=60, - ) - if _is_vertex_quota_error(response): - pytest.skip("Vertex AI quota exhausted") - assert ( - response.status_code == 200 - ), f"vertex pass-through call failed: {response.status_code} {response.text}" - - call_id = response.headers.get("x-litellm-call-id") - assert call_id, "proxy response missing x-litellm-call-id header" - - for _ in range(poll_seconds // poll_interval): - await asyncio.sleep(poll_interval) - row = _spend_log_for_request(call_id) - if row is not None and float(row.get("spend") or 0) > 0: - assert "gemini" in row["model"], f"unexpected model in spend log: {row}" - assert ( - row["custom_llm_provider"] == "vertex_ai" - ), f"unexpected provider in spend log: {row}" - return - - print(f"attempt {attempt}: spend log for call {call_id} not found yet, re-billing") - - pytest.fail( - f"Vertex pass-through spend never recorded after {max_attempts} billed calls" - ) - - @pytest.mark.asyncio() @pytest.mark.skip(reason="skip flaky test - vertex pass through streaming is flaky") async def test_basic_vertex_ai_pass_through_streaming_with_spendlog(): diff --git a/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py b/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py index a7f04466d14..41b351631f6 100644 --- a/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py @@ -202,50 +202,6 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): f"but got {cache_read}. Full usage: {usage}" ) - @pytest.mark.asyncio - async def test_prompt_caching_with_system_message(self): - """ - E2E test: Prompt caching with system message should work. - """ - _skip_live_prompt_caching_test() - litellm.turn_on_debug() - - messages = [ - { - "role": "user", - "content": "What are the key terms?", - }, - ] - - system = [ - { - "type": "text", - "text": LARGE_DOCUMENT_FOR_CACHING, - "cache_control": {"type": "ephemeral"}, - }, - ] - - response = await litellm.anthropic.messages.acreate( - model=self.get_model(), - messages=messages, - system=system, - max_tokens=100, - ) - - print(f"Response: {json.dumps(response, indent=2, default=str)}") - - usage = response.get("usage", {}) - cache_creation = usage.get("cache_creation_input_tokens", 0) - cache_read = usage.get("cache_read_input_tokens", 0) - - print(f"cache_creation_input_tokens: {cache_creation}") - print(f"cache_read_input_tokens: {cache_read}") - - assert cache_creation > 0 or cache_read > 0, ( - f"Expected cache tokens > 0 for system message caching, " - f"but got cache_creation={cache_creation}, cache_read={cache_read}" - ) - def _parse_sse_chunks(self, chunk: bytes) -> list: """ Parse SSE format chunks and return list of JSON objects. @@ -432,94 +388,3 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): f"Expected cache_read_input_tokens > 0 on second streaming call, " f"but got {cache_read}" ) - - @pytest.mark.asyncio - async def test_prompt_caching_message_start_indicates_caching_support(self): - """ - E2E test: message_start event should contain cache fields to indicate caching support. - - This validates that the message_start event includes cache_creation_input_tokens - and cache_read_input_tokens fields (even if initialized to 0) so that clients - like Claude Code can detect that prompt caching is supported. - - This test specifically addresses the issue where Bedrock converse API streaming - didn't include cache fields in message_start, causing clients to think caching - wasn't supported. - """ - _skip_live_prompt_caching_test() - litellm.turn_on_debug() - - messages = self.get_messages_with_cache_control() - - response = await litellm.anthropic.messages.acreate( - model=self.get_model(), - messages=messages, - max_tokens=100, - stream=True, - ) - - # Look for message_start event and validate it has cache fields - message_start_found = False - message_start_has_cache_creation_field = False - message_start_has_cache_read_field = False - - async for chunk in response: - # Handle SSE format chunks (bytes) - if isinstance(chunk, bytes): - json_chunks = self._parse_sse_chunks(chunk) - for json_data in json_chunks: - if json_data.get("type") == "message_start": - message_start_found = True - message = json_data.get("message", {}) - usage = message.get("usage", {}) - - print( - f"message_start usage: {json.dumps(usage, indent=2, default=str)}" - ) - - # Check that cache fields are present (even if 0) - if "cache_creation_input_tokens" in usage: - message_start_has_cache_creation_field = True - if "cache_read_input_tokens" in usage: - message_start_has_cache_read_field = True - - # Break after first message_start - break - elif isinstance(chunk, dict): - if chunk.get("type") == "message_start": - message_start_found = True - message = chunk.get("message", {}) - usage = message.get("usage", {}) - - print( - f"message_start usage: {json.dumps(usage, indent=2, default=str)}" - ) - - # Check that cache fields are present (even if 0) - if "cache_creation_input_tokens" in usage: - message_start_has_cache_creation_field = True - if "cache_read_input_tokens" in usage: - message_start_has_cache_read_field = True - - # Break after first message_start - break - - # Break if we found message_start - if message_start_found: - break - - # Validate that message_start was found - assert ( - message_start_found - ), "Expected to find message_start event in streaming response" - - # Validate that cache fields are present in message_start - assert message_start_has_cache_creation_field, ( - "Expected cache_creation_input_tokens field in message_start event. " - "This field should be present (even if 0) to indicate caching support to clients." - ) - - assert message_start_has_cache_read_field, ( - "Expected cache_read_input_tokens field in message_start event. " - "This field should be present (even if 0) to indicate caching support to clients." - ) diff --git a/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py b/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py index 9e706f99316..966048e609e 100644 --- a/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_messages_tool_search_test.py @@ -147,48 +147,6 @@ class BaseAnthropicMessagesToolSearchTest(ABC): content = response.get("content", []) assert len(content) > 0, "Response should have content" - @pytest.mark.asyncio - async def test_tool_search_discovers_tool(self): - """ - E2E test: Tool search should discover and use a deferred tool. - - This validates that when the user asks about weather, the model - discovers the get_weather tool via tool search and attempts to use it. - """ - litellm.turn_on_debug() - - tools = self.get_tools_with_tool_search() - messages = [ - { - "role": "user", - "content": "I need to know the current weather in New York City. Please use the appropriate tool.", - } - ] - - response = await litellm.anthropic.messages.acreate( - model=self.get_model(), - messages=messages, - tools=tools, - max_tokens=1024, - extra_headers=self.get_extra_headers(), - ) - - print(f"Response: {json.dumps(response, indent=2, default=str)}") - - content = response.get("content", []) - - # Check if the model used tool_use (either tool_search or get_weather) - tool_uses = [block for block in content if block.get("type") == "tool_use"] - - print(f"Tool uses: {json.dumps(tool_uses, indent=2, default=str)}") - - # The model should attempt to use tools when asked about weather - # It might use tool_search first, or directly use get_weather if discovered - if response.get("stop_reason") == "tool_use": - assert ( - len(tool_uses) > 0 - ), "Expected tool_use blocks when stop_reason is tool_use" - @pytest.mark.asyncio @pytest.mark.flaky(retries=3, delay=5) async def test_tool_search_streaming(self): @@ -234,39 +192,3 @@ class BaseAnthropicMessagesToolSearchTest(ABC): # Should have message_start message_starts = [c for c in chunks if c.get("type") == "message_start"] assert len(message_starts) > 0, "Expected message_start in streaming response" - - @pytest.mark.asyncio - async def test_tool_search_with_multiple_deferred_tools(self): - """ - E2E test: Tool search should work with multiple deferred tools. - - This validates that the model can discover the appropriate tool - from a larger catalog of deferred tools. - """ - litellm.turn_on_debug() - - tools = self.get_tools_with_tool_search() - messages = [ - {"role": "user", "content": "What's the stock price of Apple (AAPL)?"} - ] - - response = await litellm.anthropic.messages.acreate( - model=self.get_model(), - messages=messages, - tools=tools, - max_tokens=1024, - extra_headers=self.get_extra_headers(), - ) - - print(f"Response: {json.dumps(response, indent=2, default=str)}") - - # Validate response - assert "content" in response, "Response should contain content" - - content = response.get("content", []) - tool_uses = [block for block in content if block.get("type") == "tool_use"] - - # If the model decides to use a tool, it should be related to stocks - if tool_uses: - tool_names = [t.get("name") for t in tool_uses] - print(f"Tools used: {tool_names}") diff --git a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py index 16055b5a29b..858a6713f7b 100644 --- a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py @@ -62,234 +62,3 @@ class BaseAnthropicMessagesTest: assert "content" in response assert "model" in response assert response.get("role") == "assistant" - - @pytest.mark.asyncio - async def test_non_streaming_base(self): - """Base test for non-streaming requests""" - litellm.turn_on_debug() - - request_params = self.model_config - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Prepare call arguments - call_args = { - "messages": messages, - "max_tokens": 100, - } - - # Add any additional config from subclass - call_args.update(request_params) - - # Call the handler - response = await litellm.anthropic.messages.acreate(**call_args) - - print(f"Non-streaming {request_params['model']} response: ", response) - - # Verify response - self._validate_response(response) - - print(f"Non-streaming response: {json.dumps(response, indent=2, default=str)}") - return response - - @pytest.mark.asyncio - async def test_response_format_consistency(self): - """ - Test that response content blocks are consistently dicts (not Pydantic objects). - - This ensures that code like response["content"][0]["type"] works - regardless of the target provider. - - Issue: https://github.com/BerriAI/litellm/issues/20342 - """ - litellm.turn_on_debug() - - request_params = self.model_config - - # Set up test parameters - messages = [{"role": "user", "content": "Say hi"}] - - # Prepare call arguments - call_args = { - "messages": messages, - "max_tokens": 100, - } - - # Add any additional config from subclass - call_args.update(request_params) - - # Call the handler - response = await litellm.anthropic.messages.acreate(**call_args) - - print( - f"Response for {request_params['model']}: {json.dumps(response, indent=2, default=str)}" - ) - - # Verify response structure - assert "content" in response, "Response should have 'content' field" - assert len(response["content"]) > 0, "Response content should not be empty" - - # Get the first content block - block = response["content"][0] - - # Check that the block is a dict, not a Pydantic object - assert isinstance(block, dict), ( - f"Content block should be a dict, but got {type(block)}. " - f"This means response format is inconsistent across providers." - ) - - # Verify we can access fields using dict syntax (not object attributes) - try: - block_type = block["type"] - print(f"✓ Successfully accessed block['type']: {block_type}") - except TypeError as e: - pytest.fail( - f"Cannot access content block using dict syntax: {e}. " - f"Block type: {type(block)}" - ) - - # Verify the block has expected structure - assert "type" in block, "Content block should have 'type' field" - if block["type"] == "text": - assert "text" in block, "Text content block should have 'text' field" - - print( - f"✓ Response format consistency test passed for {request_params['model']}" - ) - - @pytest.mark.asyncio - async def test_anthropic_messages_litellm_router_streaming_with_logging(self): - """ - Test that logging and cost tracking works for anthropic_messages with streaming request - """ - test_custom_logger = TestCustomLogger() - litellm.callbacks = [test_custom_logger] - litellm.turn_on_debug() - router = Router( - model_list=[ - { - "model_name": "claude-special-alias", - "litellm_params": {**self.model_config}, - } - ] - ) - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Call the handler - response = await router.aanthropic_messages( - messages=messages, - model="claude-special-alias", - max_tokens=100, - stream=True, - ) - - response_prompt_tokens = 0 - response_completion_tokens = 0 - all_anthropic_usage_chunks = [] - buffer = "" - - async for chunk in response: - # Decode chunk if it's bytes - print("chunk=", chunk) - - # Handle SSE format chunks - if isinstance(chunk, bytes): - chunk_str = chunk.decode("utf-8") - buffer += chunk_str - # Extract the JSON data part from SSE format - for line in buffer.split("\n"): - if line.startswith("data: "): - try: - json_data = json.loads(line[6:]) # Skip the 'data: ' prefix - print( - "\n\nJSON data:", - json.dumps(json_data, indent=4, default=str), - ) - - # Extract usage information - if ( - json_data.get("type") == "message_start" - and "message" in json_data - ): - if "usage" in json_data["message"]: - usage = json_data["message"]["usage"] - all_anthropic_usage_chunks.append(usage) - print( - "USAGE BLOCK", - json.dumps(usage, indent=4, default=str), - ) - elif "usage" in json_data: - usage = json_data["usage"] - all_anthropic_usage_chunks.append(usage) - print( - "USAGE BLOCK", - json.dumps(usage, indent=4, default=str), - ) - except json.JSONDecodeError: - print(f"Failed to parse JSON from: {line[6:]}") - elif hasattr(chunk, "message"): - if chunk.message.usage: - print( - "USAGE BLOCK", - json.dumps(chunk.message.usage, indent=4, default=str), - ) - all_anthropic_usage_chunks.append(chunk.message.usage) - elif hasattr(chunk, "usage"): - print("USAGE BLOCK", json.dumps(chunk.usage, indent=4, default=str)) - all_anthropic_usage_chunks.append(chunk.usage) - - print( - "all_anthropic_usage_chunks", - json.dumps(all_anthropic_usage_chunks, indent=4, default=str), - ) - - # Extract token counts from usage data - if all_anthropic_usage_chunks: - response_prompt_tokens = max( - [usage.get("input_tokens", 0) for usage in all_anthropic_usage_chunks] - ) - response_completion_tokens = max( - [usage.get("output_tokens", 0) for usage in all_anthropic_usage_chunks] - ) - - print("input_tokens_anthropic_api", response_prompt_tokens) - print("output_tokens_anthropic_api", response_completion_tokens) - - await asyncio.sleep(4) - - print( - "logged_standard_logging_payload", - json.dumps( - test_custom_logger.logged_standard_logging_payload, - indent=4, - default=str, - ), - ) - - assert ( - test_custom_logger.logged_standard_logging_payload is not None - ), "Logging payload should not be None" - assert ( - test_custom_logger.logged_standard_logging_payload["messages"] == messages - ) - assert ( - test_custom_logger.logged_standard_logging_payload["response"] is not None - ) - assert ( - test_custom_logger.logged_standard_logging_payload["model"] - == self.expected_model_name_in_logging - ) - - # check logged usage + spend - assert test_custom_logger.logged_standard_logging_payload["response_cost"] > 0 - assert ( - test_custom_logger.logged_standard_logging_payload["prompt_tokens"] - == response_prompt_tokens - ) - assert ( - test_custom_logger.logged_standard_logging_payload["completion_tokens"] - == response_completion_tokens - ) diff --git a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py index 0f0bb8f091e..f0d03ca77d6 100644 --- a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py +++ b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py @@ -1,13 +1,9 @@ -import json import os from datetime import datetime from typing import Dict, Any -import asyncio import unittest.mock from unittest.mock import MagicMock -import litellm -import pytest from dotenv import load_dotenv from litellm.llms.anthropic.pass_through.messages.handler import ( anthropic_messages, @@ -16,47 +12,12 @@ from litellm.llms.anthropic.pass_through.messages.handler import ( from typing import Optional from litellm.types.utils import StandardLoggingPayload from litellm.integrations.custom_logger import CustomLogger -from litellm.router import Router -import importlib from base_anthropic_unified_messages_test import BaseAnthropicMessagesTest # Load environment variables load_dotenv() -@pytest.fixture(scope="session") -def event_loop(): - """Create an instance of the default event loop for each test session.""" - loop = asyncio.get_event_loop_policy().new_event_loop() - yield loop - loop.close() - - -@pytest.fixture(scope="function", autouse=True) -def setup_and_teardown(event_loop): # Add event_loop as a dependency - curr_dir = os.getcwd() - - import litellm - from litellm import Router - - importlib.reload(litellm) - - # Set the event loop from the fixture - asyncio.set_event_loop(event_loop) - - print(litellm) - yield - - # Clean up any pending tasks - pending = asyncio.all_tasks(event_loop) - for task in pending: - task.cancel() - - # Run the event loop until all tasks are cancelled - if pending: - event_loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True)) - - def _validate_anthropic_response(response: Dict[str, Any]): assert "id" in response assert "content" in response @@ -117,329 +78,3 @@ class TestAnthropicOpenAIAPI(BaseAnthropicMessagesTest): This is the model name that is expected to be in the logging payload """ return "gpt-4.1-mini" - - @pytest.mark.asyncio - async def test_anthropic_messages_litellm_router_streaming_with_logging(self): - """ - Test the anthropic_messages with streaming request - """ - pass - - -@pytest.mark.asyncio -async def test_anthropic_messages_litellm_router_non_streaming(): - """ - Test the anthropic_messages with non-streaming request - """ - litellm.turn_on_debug() - router = Router( - model_list=[ - { - "model_name": "claude-special-alias", - "litellm_params": { - "model": "claude-haiku-4-5-20251001", - "api_key": os.getenv("ANTHROPIC_API_KEY"), - }, - } - ] - ) - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Call the handler - response = await router.aanthropic_messages( - messages=messages, - model="claude-special-alias", - max_tokens=100, - ) - - # Verify response - assert "id" in response - assert "content" in response - assert "model" in response - assert response["role"] == "assistant" - - print(f"Non-streaming response: {json.dumps(response, indent=2)}") - return response - - -@pytest.mark.asyncio -async def test_anthropic_messages_litellm_router_routing_strategy(): - """ - Test the anthropic_messages with routing strategy + non-streaming request - """ - litellm.turn_on_debug() - router = Router( - model_list=[ - { - "model_name": "claude-special-alias", - "litellm_params": { - "model": "claude-haiku-4-5-20251001", - "api_key": os.getenv("ANTHROPIC_API_KEY"), - }, - } - ], - routing_strategy="latency-based-routing", - ) - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Call the handler - response = await router.aanthropic_messages( - messages=messages, - model="claude-special-alias", - max_tokens=100, - metadata={ - "user_id": "hello", - }, - ) - - # Verify response - assert "id" in response - assert "content" in response - assert "model" in response - assert response["role"] == "assistant" - - print(f"Non-streaming response: {json.dumps(response, indent=2)}") - return response - - -@pytest.mark.asyncio -async def test_anthropic_messages_fallbacks(): - """ - E2E test the anthropic_messages fallbacks from Anthropic API to Bedrock - """ - litellm.turn_on_debug() - router = Router( - model_list=[ - { - "model_name": "anthropic/claude-opus-4-7", - "litellm_params": { - "model": "anthropic/claude-opus-4-7", - "api_key": "bad-key", - }, - }, - { - "model_name": "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - "litellm_params": { - "model": "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - }, - }, - ], - fallbacks=[ - { - "anthropic/claude-opus-4-7": [ - "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0" - ] - } - ], - ) - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Call the handler - response = await router.aanthropic_messages( - messages=messages, - model="anthropic/claude-opus-4-7", - max_tokens=100, - metadata={ - "user_id": "hello", - }, - ) - - # Verify response - assert "id" in response - assert "content" in response - assert "model" in response - assert response["role"] == "assistant" - - print(f"Non-streaming response: {json.dumps(response, indent=2)}") - return response - - -class TestCustomLogger(CustomLogger): - def __init__(self): - super().__init__() - self.logged_standard_logging_payload: Optional[StandardLoggingPayload] = None - - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - print("inside async_log_success_event") - self.logged_standard_logging_payload = kwargs.get("standard_logging_object") - - pass - - -@pytest.mark.asyncio -async def test_anthropic_messages_litellm_router_non_streaming_with_logging(): - """ - Test the anthropic_messages with non-streaming request - - - Ensure Cost + Usage is tracked - """ - test_custom_logger = TestCustomLogger() - litellm.callbacks = [test_custom_logger] - litellm.turn_on_debug() - MODEL_GROUP = "claude-special-alias" - router = Router( - model_list=[ - { - "model_name": MODEL_GROUP, - "litellm_params": { - "model": "claude-haiku-4-5-20251001", - "api_key": os.getenv("ANTHROPIC_API_KEY"), - }, - } - ] - ) - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Call the handler - response = await router.aanthropic_messages( - messages=messages, - model=MODEL_GROUP, - max_tokens=100, - ) - - # Verify response - _validate_anthropic_response(response) - - print(f"Non-streaming response: {json.dumps(response, indent=2)}") - - await asyncio.sleep(1) - - assert ( - test_custom_logger.logged_standard_logging_payload is not None - ), "Logging payload should not be None" - print( - "tracked standard logging payload", - json.dumps( - test_custom_logger.logged_standard_logging_payload, indent=4, default=str - ), - ) - assert test_custom_logger.logged_standard_logging_payload["messages"] == messages - assert test_custom_logger.logged_standard_logging_payload["response"] is not None - assert ( - test_custom_logger.logged_standard_logging_payload["model"] - == "claude-haiku-4-5-20251001" - ) - - # check logged usage + spend - assert test_custom_logger.logged_standard_logging_payload["response_cost"] > 0 - assert ( - test_custom_logger.logged_standard_logging_payload["prompt_tokens"] - == response["usage"]["input_tokens"] - ) - assert ( - test_custom_logger.logged_standard_logging_payload["completion_tokens"] - == response["usage"]["output_tokens"] - ) - - # assert model_group - assert ( - test_custom_logger.logged_standard_logging_payload["model_group"] == MODEL_GROUP - ) - - -# @pytest.mark.asyncio -# async def test_bedrock_messages_api_header_forwarding(): -# """ -# Test that headers from kwargs (set by proxy's add_headers_to_llm_call_by_model_group) -# are correctly passed to validate_anthropic_messages_environment for Bedrock Invoke API. - -# This verifies that forward_client_headers_to_llm_api works for Bedrock Invoke API (Messages API). - -# Issue: When calling Anthropic models via the Messages API, LiteLLM makes a call to -# Bedrock's Invoke API, and custom headers were not being forwarded, even though -# they worked correctly for Chat Completions API with Bedrock's Converse API. -# """ -# from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler -# from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -# from litellm.types.router import GenericLiteLLMParams - -# handler = BaseLLMHTTPHandler() - -# # Headers that would be set by the proxy when forward_client_headers_to_llm_api is configured -# custom_headers = { -# "X-Custom-Header": "CustomValue", -# "X-Request-ID": "req-123", -# } - -# # Mock the provider config -# mock_provider_config = MagicMock() - -# # We'll check what headers are passed to this method -# mock_provider_config.validate_anthropic_messages_environment.return_value = ( -# {"Authorization": "Bearer test"}, -# "https://bedrock-runtime.us-east-1.amazonaws.com/invoke" -# ) -# mock_provider_config.transform_anthropic_messages_request.return_value = {"model": "test"} -# mock_provider_config.get_complete_url.return_value = "https://test.com" -# mock_provider_config.sign_request.return_value = ({}, None) -# mock_provider_config.transform_anthropic_messages_response.return_value = {"id": "test"} - -# # Mock HTTP client to prevent actual network calls -# with unittest.mock.patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client") as mock_get_client: -# mock_http_client = AsyncMock() -# mock_response = MagicMock() -# mock_response.status_code = 200 -# mock_response.json.return_value = {"id": "test", "content": []} -# mock_response.text = "{}" -# mock_http_client.post.return_value = mock_response -# mock_get_client.return_value = mock_http_client - -# # Mock logging object -# mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj) -# mock_logging_obj.model_call_details = {} - -# # Call the handler with headers in kwargs -# try: -# await handler.async_anthropic_messages_handler( -# model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", -# messages=[{"role": "user", "content": "Hello"}], -# anthropic_messages_provider_config=mock_provider_config, -# anthropic_messages_optional_request_params={"max_tokens": 100}, -# custom_llm_provider="bedrock", -# litellm_params=GenericLiteLLMParams( -# api_key="test-key", -# aws_region_name="us-east-1" -# ), -# logging_obj=mock_logging_obj, -# api_key="test-key", -# stream=False, -# kwargs={"headers": custom_headers} # Headers set by proxy -# ) -# except Exception: -# pass # Ignore errors, we're only checking if headers were passed - -# # Verify that validate_anthropic_messages_environment was called -# assert mock_provider_config.validate_anthropic_messages_environment.called - -# # Get the headers that were passed -# call_args = mock_provider_config.validate_anthropic_messages_environment.call_args -# passed_headers = call_args[1]["headers"] - -# # The custom headers from kwargs should be in the passed headers -# assert "X-Custom-Header" in passed_headers or "x-custom-header" in passed_headers -# assert "X-Request-ID" in passed_headers or "x-request-id" in passed_headers - - -def test_sync_openai_messages(): - """ - Test the anthropic_messages with sync request - """ - litellm.turn_on_debug() - response = litellm.anthropic.messages.create( - messages=[{"role": "user", "content": "Hello, can you tell me a short joke?"}], - model="openai/gpt-4.1-mini", - max_tokens=100, - ) - print("ANT response", response) - - assert response is not None - assert isinstance(response, dict) - assert response["content"][0]["text"] is not None diff --git a/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py b/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py index 451a09d30bb..aacf65ea4a3 100644 --- a/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py +++ b/tests/pass_through_unit_tests/test_bedrock_anthropic_messages_test.py @@ -14,54 +14,6 @@ from base_anthropic_unified_messages_test import BaseAnthropicMessagesTest INSTANCE_BASE_ANTHROPIC_MESSAGES_TEST = BaseAnthropicMessagesTest() -@pytest.mark.asyncio -async def test_anthropic_messages_litellm_router_bedrock(): - """ - Test the anthropic_messages with non-streaming request - """ - - litellm.turn_on_debug() - router = Router( - model_list=[ - { - "model_name": "bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - "litellm_params": { - "model": "bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - }, - }, - { - "model_name": "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - "litellm_params": { - "model": "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - }, - }, - ] - ) - - # Set up test parameters - messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] - - # Call 1 using bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0 - response = await router.aanthropic_messages( - messages=messages, - model="bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - max_tokens=100, - ) - - # Verify response - INSTANCE_BASE_ANTHROPIC_MESSAGES_TEST._validate_response(response) - - # Call 2 using bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0 - response = await router.aanthropic_messages( - messages=messages, - model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", - max_tokens=100, - ) - - # Verify response - INSTANCE_BASE_ANTHROPIC_MESSAGES_TEST._validate_response(response) - - @pytest.mark.asyncio async def test_anthropic_messages_bedrock_converse_with_thinking(): """ diff --git a/tests/proxy_behavior/lens/test_lifecycle.py b/tests/proxy_behavior/lens/test_lifecycle.py index 5f30a11cf70..b2a9628c1d3 100644 --- a/tests/proxy_behavior/lens/test_lifecycle.py +++ b/tests/proxy_behavior/lens/test_lifecycle.py @@ -361,10 +361,10 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database: await endpoints.repository().save_worker(legacy) authenticated_legacy: Final = await endpoints.worker_auth(credentials) assert authenticated_legacy.analysis_key_id is None - with pytest.raises(HTTPException) as needs_billing: + assert ( await endpoints.claim(authenticated_legacy, protocol_version=PROTOCOL_VERSION, worker_release=release_tag()) - assert needs_billing.value.status_code == 409 - assert "Assign an analysis key" in needs_billing.value.detail + is None + ) assert await endpoints.heartbeat(lens.id, claimed.job.id, authenticated_legacy) finished: Final = await endpoints.result( lens.id, claimed.job.id, Result(coverage=Coverage(screened=2)), authenticated_legacy, storage=None diff --git a/tests/router_unit_tests/gettysburg.wav b/tests/router_unit_tests/gettysburg.wav deleted file mode 100644 index 9690f521e84..00000000000 Binary files a/tests/router_unit_tests/gettysburg.wav and /dev/null differ diff --git a/tests/router_unit_tests/test_router_endpoints.py b/tests/router_unit_tests/test_router_endpoints.py deleted file mode 100644 index c4a2003237a..00000000000 --- a/tests/router_unit_tests/test_router_endpoints.py +++ /dev/null @@ -1,261 +0,0 @@ -import os -import json -import traceback -from typing import Optional -from dotenv import load_dotenv -from fastapi import Request -from datetime import datetime -from unittest.mock import AsyncMock, patch, MagicMock - -from litellm import Router, CustomLogger -from litellm.types.utils import StandardLoggingPayload - -## Get the current directory of the file being run -pwd = os.path.dirname(os.path.realpath(__file__)) -print(pwd) - -file_path = os.path.join(pwd, "gettysburg.wav") - -audio_file = open(file_path, "rb") -from pathlib import Path -import litellm -import pytest -import asyncio - - -@pytest.fixture -def model_list(): - return [ - { - "model_name": "gpt-5-mini", - "litellm_params": { - "model": "gpt-5-mini", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - }, - { - "model_name": "gpt-5.5", - "litellm_params": { - "model": "gpt-5.5", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - }, - { - "model_name": "gpt-image-1", - "litellm_params": { - "model": "gpt-image-1", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - }, - { - "model_name": "cohere-rerank", - "litellm_params": { - "model": "cohere/rerank-english-v3.0", - "api_key": os.getenv("COHERE_API_KEY"), - }, - }, - { - "model_name": "claude-sonnet-4-5-20250929", - "litellm_params": { - "model": "gpt-5-mini", - "mock_response": "hi this is macintosh.", - }, - }, - ] - - -# This file includes the custom callbacks for LiteLLM Proxy -# Once defined, these can be passed in proxy_config.yaml -class MyCustomHandler(CustomLogger): - def __init__(self): - self.openai_client = None - - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - try: - # init logging config - print("logging a transcript kwargs: ", kwargs) - print("openai client=", kwargs.get("client")) - self.openai_client = kwargs.get("client") - self.standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get( - "standard_logging_object" - ) - - except Exception: - pass - - -# Set litellm.callbacks = [proxy_handler_instance] on the proxy -@pytest.mark.asyncio -@pytest.mark.flaky(retries=6, delay=10) -async def test_transcription_on_router(): - proxy_handler_instance = MyCustomHandler() - litellm.set_verbose = True - litellm.callbacks = [proxy_handler_instance] - print("\n Testing async transcription on router\n") - try: - model_list = [ - { - "model_name": "whisper", - "litellm_params": { - "model": "whisper-1", - }, - }, - { - "model_name": "whisper", - "litellm_params": { - "model": "azure/azure-whisper", - "api_base": "https://my-endpoint-europe-berri-992.openai.azure.com/", - "api_key": os.getenv("AZURE_EUROPE_API_KEY"), - "api_version": "2024-02-15-preview", - }, - }, - ] - - router = Router(model_list=model_list) - - router_level_clients = [] - for deployment in router.model_list: - _deployment_openai_client = router._get_client( - deployment=deployment, - kwargs={"model": "whisper-1"}, - client_type="async", - ) - - router_level_clients.append(str(_deployment_openai_client)) - - ## test 1: user facing function - response = await router.atranscription( - model="whisper", - file=audio_file, - ) - - ## test 2: underlying function - response = await router._atranscription( - model="whisper", - file=audio_file, - ) - print(response) - - # PROD Test - # Ensure we ONLY use OpenAI/Azure client initialized on the router level - await asyncio.sleep(5) - print("OpenAI Client used= ", proxy_handler_instance.openai_client) - print("all router level clients= ", router_level_clients) - assert proxy_handler_instance.openai_client in router_level_clients - except Exception as e: - traceback.print_exc() - pytest.fail(f"Error occurred: {e}") - - -@pytest.mark.parametrize("mode", ["iterator"]) # "file", -@pytest.mark.asyncio -async def test_audio_speech_router(mode): - litellm.set_verbose = True - test_logger = MyCustomHandler() - litellm.callbacks = [test_logger] - from litellm import Router - - client = Router( - model_list=[ - { - "model_name": "tts", - "litellm_params": { - "model": "openai/tts-1", - }, - }, - ] - ) - - response = await client.aspeech( - model="tts", - voice="alloy", - input="the quick brown fox jumped over the lazy dogs", - api_base=None, - api_key=None, - organization=None, - project=None, - max_retries=1, - timeout=600, - client=None, - optional_params={}, - ) - - await asyncio.sleep(3) - - from litellm.llms.openai.openai import HttpxBinaryResponseContent - - assert isinstance(response, HttpxBinaryResponseContent) - - assert test_logger.standard_logging_object is not None - print( - "standard_logging_object=", - json.dumps(test_logger.standard_logging_object, indent=4), - ) - assert test_logger.standard_logging_object["model_group"] == "tts" - - - - - - - - -@pytest.mark.asyncio() -async def test_rerank_endpoint(model_list): - from litellm.types.utils import RerankResponse - - router = Router(model_list=model_list) - - ## Test 1: user facing function - response = await router.arerank( - model="cohere-rerank", - query="hello", - documents=["hello", "world"], - top_n=3, - ) - - ## Test 2: underlying function - response = await router._arerank( - model="cohere-rerank", - query="hello", - documents=["hello", "world"], - top_n=3, - ) - - print("async re rank response: ", response) - - assert response.id is not None - assert response.results is not None - - RerankResponse.model_validate(response) - - -@pytest.mark.asyncio() -@pytest.mark.parametrize( - "model", ["omni-moderation-latest", "openai/omni-moderation-latest", None] -) -async def test_moderation_endpoint(model): - litellm.set_verbose = True - router = Router( - model_list=[ - { - "model_name": "openai/*", - "litellm_params": { - "model": "openai/*", - }, - }, - { - "model_name": "*", - "litellm_params": { - "model": "openai/*", - }, - }, - ] - ) - - if model is None: - response = await router.amoderation(input="hello this is a test") - else: - response = await router.amoderation(model=model, input="hello this is a test") - - print("moderation response: ", response) diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py deleted file mode 100644 index c2ed526e4c7..00000000000 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ /dev/null @@ -1,526 +0,0 @@ -import asyncio -import json -import os -import traceback -from dotenv import load_dotenv -from fastapi import Request -from datetime import datetime, timezone - -from litellm import Router -import pytest -import litellm -from unittest.mock import patch, MagicMock, AsyncMock -from litellm.types.utils import ModelResponse, StandardLoggingPayload -from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper -from litellm.caching.dual_cache import DualCache -from litellm.caching.in_memory_cache import InMemoryCache -from litellm.types.caching import RedisPipelineIncrementOperation -from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute -from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo -from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS, ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY - - -@pytest.fixture -def model_list(): - return [ - { - "model_name": "gpt-5-mini", - "litellm_params": { - "model": "gpt-5-mini", - "api_key": os.getenv("OPENAI_API_KEY"), - "tpm": 1000, # Add TPM limit so async method doesn't return early - "rpm": 100, # Add RPM limit so async method doesn't return early - }, - "model_info": { - "access_groups": ["group1", "group2"], - }, - }, - { - "model_name": "gpt-5.5", - "litellm_params": { - "model": "gpt-5.5", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - }, - { - "model_name": "gpt-image-1", - "litellm_params": { - "model": "gpt-image-1", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - }, - { - "model_name": "*", - "litellm_params": { - "model": "openai/*", - "api_key": os.getenv("OPENAI_API_KEY"), - }, - }, - { - "model_name": "claude-*", - "litellm_params": { - "model": "anthropic/*", - "api_key": os.getenv("ANTHROPIC_API_KEY"), - }, - }, - ] - - -def test_validate_fallbacks(model_list): - router = Router(model_list=model_list, fallbacks=[{"gpt-5.5": "gpt-5-mini"}]) - router.validate_fallbacks(fallback_param=[{"gpt-5.5": "gpt-5-mini"}]) - - -def test_routing_strategy_init(model_list): - """Test if all routing strategies are initialized correctly""" - from litellm.types.router import RoutingStrategy - - router = Router(model_list=model_list) - for strategy in RoutingStrategy: - router.routing_strategy_init( - routing_strategy=strategy, routing_strategy_args={} - ) - - - - -def test_routing_strategy_init_valid_string_strategies(model_list): - """Test that all valid string routing strategies work without error. - - Valid strategies are derived from RoutingStrategy enum values plus 'simple-shuffle'. - """ - from litellm.types.router import RoutingStrategy - - router = Router(model_list=model_list) - - # All strategies from enum + simple-shuffle (default, not in enum) - valid_strategies = ["simple-shuffle"] + [s.value for s in RoutingStrategy] - - for strategy in valid_strategies: - # Should not raise - router.routing_strategy_init( - routing_strategy=strategy, routing_strategy_args={} - ) - - - - - - - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.flaky(retries=6, delay=1) -@pytest.mark.asyncio -async def test_image_generation(model_list, sync_mode): - """Test if the underlying '_image_generation' function is working correctly""" - from litellm.types.utils import ImageResponse - - router = Router(model_list=model_list) - if sync_mode: - response = router._image_generation( - model="gpt-image-1", - prompt="A cute baby sea otter", - ) - else: - response = await router._aimage_generation( - model="gpt-image-1", - prompt="A cute baby sea otter", - ) - - ImageResponse.model_validate(response) - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -def _rpm_tpm_router(model_id: str) -> Router: - return Router( - model_list=[ - { - "model_name": "gpt-5-mini", - "litellm_params": {"model": "gpt-5-mini", "api_key": "sk-fake", "tpm": 1000, "rpm": 100}, - "model_info": {"id": model_id}, - } - ] - ) - - -@pytest.fixture -def router_minute_pinned(monkeypatch): - pinned = datetime(2026, 1, 1, 12, 0, 30, tzinfo=timezone.utc) - monkeypatch.setattr("litellm.router.get_utc_datetime", lambda: pinned) - - -def _ratelimit_headers(response: ModelResponse | CustomStreamWrapper) -> dict[str, int]: - return {k: v for k, v in response._hidden_params["additional_headers"].items() if k.startswith("x-ratelimit-")} - - -@pytest.mark.asyncio -@pytest.mark.usefixtures("router_minute_pinned") -async def test_acompletion_headers_read_post_increment_counter_and_count_once(): - router = _rpm_tpm_router("lit-3058-async") - - response = await router.acompletion( - model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong" - ) - total_tokens = response.usage.total_tokens - assert total_tokens > 0 - - headers = _ratelimit_headers(response) - assert headers["x-ratelimit-remaining-tokens"] == 1000 - total_tokens - assert headers["x-ratelimit-remaining-requests"] == 99 - assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1) - - await asyncio.sleep(0.5) - assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1) - - - - -@pytest.mark.asyncio -@pytest.mark.usefixtures("router_minute_pinned") -async def test_acompletion_stream_counts_request_before_headers_and_tokens_once_on_completion(): - router = _rpm_tpm_router("lit-3058-stream") - - stream = await router.acompletion( - model="gpt-5-mini", - messages=[{"role": "user", "content": "hi"}], - mock_response="pong pong pong", - stream=True, - stream_options={"include_usage": True}, - ) - headers = _ratelimit_headers(stream) - assert headers["x-ratelimit-remaining-tokens"] == 1000 - assert headers["x-ratelimit-remaining-requests"] == 99 - assert await router.get_model_group_usage("gpt-5-mini") == (0, 1) - - chunks = [chunk async for chunk in stream] - total_tokens = chunks[-1].usage.total_tokens - assert total_tokens > 0 - - await asyncio.sleep(0.5) - assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1) - - - - -class _GatedIncrementCache(DualCache): - def __init__(self) -> None: - super().__init__(in_memory_cache=InMemoryCache()) - self.first_increment_started = asyncio.Event() - self.release_first_increment = asyncio.Event() - self.increment_calls = 0 - - async def async_increment_cache_pipeline( - self, - increment_list: list[RedisPipelineIncrementOperation], - local_only: bool = False, - parent_otel_span: object = None, - **kwargs: object, - ) -> list[float] | None: - self.increment_calls += 1 - if self.increment_calls == 1: - self.first_increment_started.set() - await self.release_first_increment.wait() - return await super().async_increment_cache_pipeline( - increment_list, local_only=local_only, parent_otel_span=parent_otel_span, **kwargs - ) - - -@pytest.mark.asyncio -async def test_success_callback_running_during_pre_header_increment_does_not_double_count(): - router = _rpm_tpm_router("lit-3058-race") - cache = _GatedIncrementCache() - router.cache = cache - - request = asyncio.ensure_future( - router.acompletion(model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong") - ) - await asyncio.wait_for(cache.first_increment_started.wait(), timeout=5) - for _ in range(50): - if get_deployment_successes_for_current_minute(router, "lit-3058-race") == 1: - break - await asyncio.sleep(0.1) - assert get_deployment_successes_for_current_minute(router, "lit-3058-race") == 1 - assert cache.increment_calls == 1 - - cache.release_first_increment.set() - response = await request - - assert await router.get_model_group_usage("gpt-5-mini") == (response.usage.total_tokens, 1) - - -class _UnavailableIncrementCache(DualCache): - def __init__(self) -> None: - super().__init__(in_memory_cache=InMemoryCache()) - self.first_increment_started = asyncio.Event() - self.release_first_increment = asyncio.Event() - self.increment_calls = 0 - - async def async_increment_cache_pipeline( - self, - increment_list: list[RedisPipelineIncrementOperation], - local_only: bool = False, - parent_otel_span: object = None, - **kwargs: object, - ) -> list[float] | None: - self.increment_calls += 1 - if self.increment_calls == 1: - self.first_increment_started.set() - await self.release_first_increment.wait() - raise RuntimeError("cache unavailable") - - -@pytest.mark.asyncio -async def test_callback_observing_stamp_before_pre_header_increment_fails_leaves_no_stamp_behind(): - router = _rpm_tpm_router("lit-3058-fail") - cache = _UnavailableIncrementCache() - router.cache = cache - metadata: dict[str, object] = {} - - request = asyncio.ensure_future( - router.acompletion( - model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong", metadata=metadata - ) - ) - await asyncio.wait_for(cache.first_increment_started.wait(), timeout=5) - assert metadata[ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY] == 30 - for _ in range(50): - if get_deployment_successes_for_current_minute(router, "lit-3058-fail") == 1: - break - await asyncio.sleep(0.1) - assert get_deployment_successes_for_current_minute(router, "lit-3058-fail") == 1 - assert cache.increment_calls == 1 - - cache.release_first_increment.set() - response = await request - - assert response.usage.total_tokens == 30 - assert ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY not in metadata - assert _ratelimit_headers(response)["x-ratelimit-remaining-requests"] == 100 - assert await router.get_model_group_usage("gpt-5-mini") == (None, None) - - -def test_track_deployment_metrics(model_list): - """Test if the 'track_deployment_metrics' function is working correctly""" - from litellm.types.utils import ModelResponse - - router = Router(model_list=model_list) - router._track_deployment_metrics( - deployment=router.get_deployment_by_model_group_name( - model_group_name="gpt-5-mini" - ), - response=ModelResponse( - model="gpt-5-mini", - usage={"total_tokens": 100}, - ), - parent_otel_span=None, - ) - - -def test_pass_through_assistants_endpoint_factory(model_list): - """Test if the 'pass_through_assistants_endpoint_factory' function is working correctly""" - router = Router(model_list=model_list) - router._pass_through_assistants_endpoint_factory( - original_function=litellm.acreate_assistants, - custom_llm_provider="openai", - client=None, - **{}, - ) - - -def test_factory_function(model_list): - """Test if the 'factory_function' function is working correctly""" - router = Router(model_list=model_list) - router.factory_function(litellm.acreate_assistants) - - - - - - - - - - - - -# def test_pattern_match_deployments(model_list): -# from litellm.router_utils.pattern_match_deployments import PatternMatchRouter -# import re - -# patter_router = PatternMatchRouter() - -# request = "fo::hi::static::hello" -# model_name = "fo::*:static::*" - -# model_name_regex = patter_router._pattern_to_regex(model_name) - -# # Match against the request -# match = re.match(model_name_regex, request) - -# print(f"match: {match}") -# print(f"match.end: {match.end()}") -# if match is None: -# raise ValueError("Match not found") -# updated_model = patter_router.set_deployment_model_name( -# matched_pattern=match, litellm_deployment_litellm_model="openai/*" -# ) -# assert updated_model == "openai/fo::hi:static::hello" - - - - -@pytest.mark.asyncio -async def test_pass_through_moderation_endpoint_factory(model_list): - router = Router(model_list=model_list) - response = await router._pass_through_moderation_endpoint_factory( - original_function=litellm.amoderation, - input="this is valid good text", - model=None, - ) - assert response is not None - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -def test_handle_clientside_credential_no_metadata(model_list): - """Test that _handle_clientside_credential handles cases where no metadata is provided""" - router = Router(model_list=model_list) - - # Mock deployment - deployment = { - "model_name": "gpt-4.1", - "litellm_params": {"model": "gpt-4.1", "api_key": "test_key"}, - "model_info": {"id": "original-id-789"}, - } - - # Mock kwargs with clientside credentials but NO metadata - kwargs = { - "api_key": "client_side_key", - "api_base": "https://api.openai.com/v1", - # No metadata key at all - } - - # This should fail because there's no model_group in metadata - # The function expects to find model_group in the metadata - try: - result_deployment = router._handle_clientside_credential( - deployment=deployment, kwargs=kwargs, function_name="acompletion" - ) - # If we get here, the function should have used deployment.model_name as fallback - assert result_deployment.model_name == "gpt-4.1" - print("✓ Success with no metadata - used deployment.model_name as fallback") - except Exception as e: - # This is expected behavior - the function needs model_group to generate model_id - print(f"✓ Correctly handled no metadata case: {e}") - - # Test with empty metadata - kwargs_with_empty_metadata = { - "api_key": "client_side_key", - "api_base": "https://api.openai.com/v1", - "metadata": {}, # Empty metadata - } - - try: - result_deployment = router._handle_clientside_credential( - deployment=deployment, - kwargs=kwargs_with_empty_metadata, - function_name="acompletion", - ) - # Should fail because empty metadata has no model_group - pytest.fail("Expected failure with empty metadata") - except Exception as e: - print(f"✓ Correctly handled empty metadata case: {e}") diff --git a/tests/rust-python-harness/AGENTS.md b/tests/rust-python-harness/AGENTS.md index b66eaaeda9b..6b0609a1c9a 100644 --- a/tests/rust-python-harness/AGENTS.md +++ b/tests/rust-python-harness/AGENTS.md @@ -63,4 +63,4 @@ tests/rust-python-harness/ - `shared/` contains reusable parity, tracing, reporting primitives, and unit-runner machinery - Keep fixtures with their owning API and existing Python tests in their current locations - Each strategy folder carries an `AGENTS.md` one-liner stating what it should be doing -- Run the harness's own checks with `uv run pytest -o consider_namespace_packages=true tests/rust-python-harness/shared tests/rust-python-harness/cli tests/rust-python-harness/strategies/trace_parity tests/rust-python-harness/strategies/unit_tests_parity tests/rust-python-harness/strategies/unit_tests_rust tests/test_rust_python_harness.py -q` +- Run the harness's own checks with `uv run pytest -o consider_namespace_packages=true tests/rust-python-harness/shared tests/rust-python-harness/cli tests/rust-python-harness/strategies/trace_parity tests/rust-python-harness/strategies/unit_tests_parity tests/rust-python-harness/strategies/unit_tests_rust -q` diff --git a/tests/rust-python-harness/cli/test_cli.py b/tests/rust-python-harness/cli/test_cli.py index 219e1b0c6b7..ebc4e2653d9 100644 --- a/tests/rust-python-harness/cli/test_cli.py +++ b/tests/rust-python-harness/cli/test_cli.py @@ -464,3 +464,12 @@ def test_runner_interrupt_skips_the_completion_report( assert exit_code == 130 assert "Rust <-> Python parity report" not in captured.out assert captured.err == "Interrupted\n" + + +def test_strategy_subcommand_accepts_function_filter(capsys: pytest.CaptureFixture[str]) -> None: + exit_code: Final = main(["run", "unit_tests_rust", "--function", "messages"]) + + captured: Final = capsys.readouterr() + assert exit_code == 0 + assert "- messages: not_implemented" in captured.out + assert "unit_tests_rust:messages: not_implemented" not in captured.out diff --git a/tests/rust-python-harness/shared/reporting/test_models.py b/tests/rust-python-harness/shared/reporting/test_models.py new file mode 100644 index 00000000000..1dd27719972 --- /dev/null +++ b/tests/rust-python-harness/shared/reporting/test_models.py @@ -0,0 +1,53 @@ +from __future__ import annotations + +from typing import Final + +import pytest + +from .models import CaseResult, Coverage, HarnessCase, RunStatus +from .strategy import ModuleCaseSpec, NotImplementedCaseSpec, SkippedCaseSpec + + +def _case(spec: ModuleCaseSpec | NotImplementedCaseSpec | SkippedCaseSpec) -> HarnessCase: + return HarnessCase(strategy_id="example", strategy_label="Example", sdk_function="messages", spec=spec) + + +def _runnable_case() -> HarnessCase: + return _case(ModuleCaseSpec(coverage=Coverage.COMPLETE, module="tests.example")) + + +def test_should_mark_not_implemented_and_skipped_cases_without_running() -> None: + not_implemented: Final = CaseResult(case=_case(NotImplementedCaseSpec(reason="No case is registered."))) + skipped: Final = CaseResult(case=_case(SkippedCaseSpec(reason="The surface does not apply."))) + + not_implemented.set_initial_status() + skipped.set_initial_status() + + assert not_implemented.status is RunStatus.NOT_IMPLEMENTED + assert skipped.status is RunStatus.SKIPPED + + +def test_should_finalize_a_fully_passing_case() -> None: + result: Final = CaseResult(case=_runnable_case()) + result.set_initial_status() + result.collected.update({"one", "two"}) + result.completed.update({"one", "two"}) + result.passed = 2 + + result.finalize() + + assert result.status is RunStatus.PASSED + + +def test_should_replace_a_pass_with_a_teardown_error() -> None: + result: Final = CaseResult(case=_runnable_case()) + result.set_initial_status() + result.collected.add("one") + + result.record("one", RunStatus.PASSED, 0.1) + result.record("one", RunStatus.ERROR, 0.2) + + assert result.status is RunStatus.ERROR + assert result.passed == 0 + assert result.errors == 1 + assert result.duration == pytest.approx(0.3) diff --git a/tests/rust-python-harness/shared/reporting/test_ui.py b/tests/rust-python-harness/shared/reporting/test_ui.py new file mode 100644 index 00000000000..c0ed0bc8bd5 --- /dev/null +++ b/tests/rust-python-harness/shared/reporting/test_ui.py @@ -0,0 +1,26 @@ +from __future__ import annotations + +from typing import Final + +from .models import Coverage, HarnessCase, HarnessRun, RunStatus +from .strategy import ModuleCaseSpec +from .ui import _format_duration, _summary + + +def test_should_format_developer_facing_run_context() -> None: + run: Final = HarnessRun.from_cases( + ( + HarnessCase( + strategy_id="example", + strategy_label="Example", + sdk_function="messages", + spec=ModuleCaseSpec(coverage=Coverage.COMPLETE, module="tests.example"), + ), + ) + ) + result: Final = next(iter(run.results.values())) + result.collected.add("tests/test_parity.py::test_one") + result.record("tests/test_parity.py::test_one", RunStatus.PASSED, 1.25) + + assert _summary(run) == (1, 0, 0, 0) + assert _format_duration(1.25) == "1.2s" diff --git a/tests/rust-python-harness/strategies/trace_parity/sdk/test_case_modules_importable.py b/tests/rust-python-harness/strategies/trace_parity/sdk/test_case_modules_importable.py new file mode 100644 index 00000000000..bc37ce3d873 --- /dev/null +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/test_case_modules_importable.py @@ -0,0 +1,23 @@ +from __future__ import annotations + +from importlib import import_module +from typing import Final, cast + +import pytest + +from ..models import TraceSuite + + +@pytest.mark.parametrize( + "module", + [ + "tests.rust-python-harness.strategies.trace_parity.sdk.messages.case", + "tests.rust-python-harness.strategies.trace_parity.sdk.chat_completions.case", + "tests.rust-python-harness.strategies.trace_parity.sdk.transcription.case", + ], +) +def test_implemented_namespace_case_modules_remain_importable(module: str) -> None: + loaded: Final = import_module(module) + suite: Final = cast(object, getattr(loaded, "TRACE_SUITE")) + assert isinstance(suite, TraceSuite) + assert suite.scenarios diff --git a/tests/search_tests/base_search_unit_tests.py b/tests/search_tests/base_search_unit_tests.py index 7028f58a1a3..8c16e790607 100644 --- a/tests/search_tests/base_search_unit_tests.py +++ b/tests/search_tests/base_search_unit_tests.py @@ -113,62 +113,3 @@ class BaseSearchTest(ABC): except Exception as e: pytest.fail(f"Search call failed: {str(e)}") - - def test_search_response_structure(self): - """ - Test that the Search response has the correct structure. - """ - litellm.set_verbose = True - search_provider = self.get_search_provider() - - response = litellm.search( - query="artificial intelligence recent news", - search_provider=search_provider, - ) - - # Validate response structure - assert hasattr(response, "results"), "Response should have 'results' attribute" - assert hasattr(response, "object"), "Response should have 'object' attribute" - - assert isinstance(response.results, list), "results should be a list" - assert len(response.results) > 0, "Should have at least one result" - assert response.object == "search", "object should be 'search'" - - # Validate first result structure - first_result = response.results[0] - assert hasattr(first_result, "title"), "Result should have 'title' attribute" - assert hasattr(first_result, "url"), "Result should have 'url' attribute" - assert hasattr( - first_result, "snippet" - ), "Result should have 'snippet' attribute" - assert isinstance(first_result.title, str), "title should be a string" - assert isinstance(first_result.url, str), "url should be a string" - assert isinstance(first_result.snippet, str), "snippet should be a string" - - print(f"\nResponse structure validated:") - print(f" - object: {response.object}") - print(f" - results: {len(response.results)}") - print(f" - first result has all required fields") - - def test_search_with_optional_params(self): - """ - Test search with optional parameters. - """ - litellm.set_verbose = True - search_provider = self.get_search_provider() - - response = litellm.search( - query="machine learning", - search_provider=search_provider, - max_results=5, - ) - - # Validate response - assert hasattr(response, "results"), "Response should have 'results' attribute" - assert isinstance(response.results, list), "results should be a list" - assert len(response.results) > 0, "Should have at least one result" - assert len(response.results) <= 5, "Should have at most 5 results as requested" - - print(f"\nSearch with optional params validated:") - print(f" - Requested max_results: 5") - print(f" - Received results: {len(response.results)}") diff --git a/tests/search_tests/test_duckduckgo_search.py b/tests/search_tests/test_duckduckgo_search.py deleted file mode 100644 index 682221326bb..00000000000 --- a/tests/search_tests/test_duckduckgo_search.py +++ /dev/null @@ -1,138 +0,0 @@ -""" -Tests for DuckDuckGo Search API integration. -""" - -import os - -import pytest - -import litellm -from tests.search_tests.base_search_unit_tests import BaseSearchTest - - -class TestDuckDuckGoSearch(BaseSearchTest): - """ - Tests for DuckDuckGo Search functionality. - """ - - def get_search_provider(self) -> str: - """ - Return search_provider for DuckDuckGo Search. - """ - return "duckduckgo" - - @pytest.mark.asyncio - async def test_basic_search(self): - """ - Test basic search functionality with a simple query. - """ - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - litellm.turn_on_debug() - search_provider = self.get_search_provider() - print("Search Provider=", search_provider) - - try: - response = await litellm.asearch( - query="india", - search_provider=search_provider, - ) - print("Search response=", response.model_dump_json(indent=4)) - - print(f"\n{'='*80}") - print(f"Response type: {type(response)}") - print( - f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}" - ) - - # Check if response has expected Search format - assert hasattr( - response, "results" - ), "Response should have 'results' attribute" - assert hasattr( - response, "object" - ), "Response should have 'object' attribute" - assert ( - response.object == "search" - ), f"Expected object='search', got '{response.object}'" - - # Validate results structure - assert isinstance(response.results, list), "results should be a list" - assert len(response.results) > 0, "Should have at least one result" - - # Check first result structure - first_result = response.results[0] - assert hasattr( - first_result, "title" - ), "Result should have 'title' attribute" - assert hasattr(first_result, "url"), "Result should have 'url' attribute" - assert hasattr( - first_result, "snippet" - ), "Result should have 'snippet' attribute" - - print(f"Total results: {len(response.results)}") - print(f"First result title: {first_result.title}") - print(f"First result URL: {first_result.url}") - print(f"First result snippet: {first_result.snippet[:100]}...") - print(f"{'='*80}\n") - - assert len(first_result.title) > 0, "Title should not be empty" - assert len(first_result.url) > 0, "URL should not be empty" - assert len(first_result.snippet) > 0, "Snippet should not be empty" - - # Validate cost tracking in _hidden_params - assert hasattr( - response, "_hidden_params" - ), "Response should have '_hidden_params' attribute" - hidden_params = response._hidden_params - assert ( - "response_cost" in hidden_params - ), "_hidden_params should contain 'response_cost'" - - response_cost = hidden_params["response_cost"] - assert response_cost is not None, "response_cost should not be None" - assert isinstance( - response_cost, (int, float) - ), "response_cost should be a number" - assert response_cost == 0, "response_cost should be 0" - - print(f"Cost tracking: ${response_cost:.6f}") - - except Exception as e: - pytest.fail(f"Search call failed: {str(e)}") - - def test_search_response_structure(self): - """ - Test that the Search response has the correct structure. - """ - litellm.set_verbose = True - search_provider = self.get_search_provider() - - response = litellm.search( - query="india", - search_provider=search_provider, - ) - - # Validate response structure - assert hasattr(response, "results"), "Response should have 'results' attribute" - assert hasattr(response, "object"), "Response should have 'object' attribute" - - assert isinstance(response.results, list), "results should be a list" - assert len(response.results) > 0, "Should have at least one result" - assert response.object == "search", "object should be 'search'" - - # Validate first result structure - first_result = response.results[0] - assert hasattr(first_result, "title"), "Result should have 'title' attribute" - assert hasattr(first_result, "url"), "Result should have 'url' attribute" - assert hasattr( - first_result, "snippet" - ), "Result should have 'snippet' attribute" - assert isinstance(first_result.title, str), "title should be a string" - assert isinstance(first_result.url, str), "url should be a string" - assert isinstance(first_result.snippet, str), "snippet should be a string" - - print(f"\nResponse structure validated:") - print(f" - object: {response.object}") - print(f" - results: {len(response.results)}") - print(f" - first result has all required fields") diff --git a/tests/search_tests/test_firecrawl_search.py b/tests/search_tests/test_firecrawl_search.py deleted file mode 100644 index eec74d48e26..00000000000 --- a/tests/search_tests/test_firecrawl_search.py +++ /dev/null @@ -1,42 +0,0 @@ -from unittest.mock import Mock, patch -import litellm - - -def test_firecrawl_search_request_body(): - """ - Test that validates the Firecrawl search request body is correctly formatted. - """ - mock_response = Mock() - mock_response.status_code = 200 - mock_response.json.return_value = { - "success": True, - "data": { - "web": [ - { - "title": "Test Title", - "url": "https://example.com", - "markdown": "Test content", - } - ] - }, - } - - with patch( - "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", - return_value=mock_response, - ) as mock_post: - litellm.search( - query="test query", - search_provider="firecrawl", - max_results=10, - country="US", - ) - - assert mock_post.called - call_kwargs = mock_post.call_args.kwargs - request_body = call_kwargs.get("json") - - assert request_body is not None - assert request_body["query"] == "test query" - assert request_body["limit"] == 10 - assert request_body["country"] == "US" diff --git a/tests/spend_tracking_tests/test_ocr_spend_tracking.py b/tests/spend_tracking_tests/test_ocr_spend_tracking.py deleted file mode 100644 index 3ce77c56361..00000000000 --- a/tests/spend_tracking_tests/test_ocr_spend_tracking.py +++ /dev/null @@ -1,296 +0,0 @@ -""" -Unit tests for OCR spend tracking in get_logging_payload. - -This test file verifies that OCR/AOCR calls correctly extract usage_info -and populate the spend logs payload with pages_processed instead of token counts. -""" - -import pytest -from datetime import datetime, timezone -from unittest.mock import Mock -from pydantic import BaseModel -from typing import Optional - -from litellm.proxy.spend_tracking.spend_tracking_utils import ( - get_logging_payload, - _extract_usage_for_ocr_call, -) - - -class MockUsageInfo(BaseModel): - """Mock Pydantic model for OCR usage_info""" - - pages_processed: int - doc_size_bytes: Optional[int] = None - - -class MockOCRResponse(BaseModel): - """Mock Pydantic model for OCR response""" - - id: str - object: str - model: str - usage_info: MockUsageInfo - - -class TestExtractUsageForOCRCall: - """Test the _extract_usage_for_ocr_call helper method""" - - def test_extract_usage_from_dict(self): - """Test extracting usage from dict response""" - response_obj_dict = {"usage_info": {"pages_processed": 5}} - - usage = _extract_usage_for_ocr_call(response_obj_dict, response_obj_dict) - - assert usage["prompt_tokens"] == 0 - assert usage["completion_tokens"] == 0 - assert usage["total_tokens"] == 0 - assert usage["pages_processed"] == 5 - - def test_extract_usage_from_pydantic_model(self): - """Test extracting usage from Pydantic model response""" - usage_info = MockUsageInfo(pages_processed=10, doc_size_bytes=1024) - response_obj = MockOCRResponse( - id="ocr-123", object="ocr", model="test-ocr-model", usage_info=usage_info - ) - response_obj_dict = response_obj.model_dump() - - usage = _extract_usage_for_ocr_call(response_obj, response_obj_dict) - - assert usage["prompt_tokens"] == 0 - assert usage["completion_tokens"] == 0 - assert usage["total_tokens"] == 0 - assert usage["pages_processed"] == 10 - - def test_extract_usage_with_object_attributes(self): - """Test extracting usage from object with __dict__""" - - class SimpleUsageInfo: - def __init__(self, pages_processed): - self.pages_processed = pages_processed - - class SimpleOCRResponse: - def __init__(self): - self.usage_info = SimpleUsageInfo(pages_processed=3) - - response_obj = SimpleOCRResponse() - response_obj_dict = {} - - usage = _extract_usage_for_ocr_call(response_obj, response_obj_dict) - - assert usage.get("prompt_tokens") == 0 - assert usage.get("completion_tokens") == 0 - assert usage.get("total_tokens") == 0 - assert usage.get("pages_processed") == 3 - - def test_extract_usage_missing_usage_info(self): - """Test handling missing usage_info""" - response_obj_dict = {} - - usage = _extract_usage_for_ocr_call(response_obj_dict, response_obj_dict) - - assert usage == {} - - def test_extract_usage_empty_usage_info(self): - """Test handling empty usage_info""" - response_obj_dict = {"usage_info": {}} - - usage = _extract_usage_for_ocr_call(response_obj_dict, response_obj_dict) - - assert usage.get("prompt_tokens") == 0 - assert usage.get("completion_tokens") == 0 - assert usage.get("total_tokens") == 0 - assert usage.get("pages_processed") == 0 - - -class TestGetLoggingPayloadOCR: - """Test get_logging_payload with OCR call types""" - - @pytest.fixture - def mock_datetime(self): - """Fixture for consistent timestamps""" - return datetime.now(timezone.utc) - - @pytest.fixture - def base_kwargs(self): - """Fixture for base kwargs used in tests""" - return { - "model": "test-ocr-model", - "call_type": "ocr", - "litellm_params": {}, - "response_cost": 0.05, - } - - def test_ocr_call_with_dict_response(self, mock_datetime, base_kwargs): - """Test OCR call with dict response containing usage_info""" - response_obj = { - "id": "ocr-test-123", - "object": "ocr", - "model": "test-ocr-model", - "usage_info": {"pages_processed": 7, "doc_size_bytes": 2048}, - } - - payload = get_logging_payload( - kwargs=base_kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "ocr" - assert payload["prompt_tokens"] == 0 - assert payload["completion_tokens"] == 0 - assert payload["total_tokens"] == 0 - assert payload["spend"] == 0.05 - - # Verify pages_processed is in additional_usage_values - import json - - metadata = json.loads(payload["metadata"]) - assert "additional_usage_values" in metadata - assert metadata["additional_usage_values"]["pages_processed"] == 7 - - def test_aocr_call_with_pydantic_response(self, mock_datetime, base_kwargs): - """Test AOCR (async OCR) call with Pydantic model response""" - base_kwargs["call_type"] = "aocr" - - usage_info = MockUsageInfo(pages_processed=12) - response_obj = MockOCRResponse( - id="aocr-test-456", - object="ocr", - model="test-ocr-model", - usage_info=usage_info, - ) - - payload = get_logging_payload( - kwargs=base_kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "aocr" - assert payload["prompt_tokens"] == 0 - assert payload["completion_tokens"] == 0 - assert payload["total_tokens"] == 0 - - # Verify pages_processed is in additional_usage_values - import json - - metadata = json.loads(payload["metadata"]) - assert "additional_usage_values" in metadata - assert metadata["additional_usage_values"]["pages_processed"] == 12 - - def test_ocr_call_missing_usage_info(self, mock_datetime, base_kwargs): - """Test OCR call with missing usage_info returns empty usage""" - response_obj = { - "id": "ocr-test-789", - "object": "ocr", - "model": "test-ocr-model", - } - - payload = get_logging_payload( - kwargs=base_kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "ocr" - assert payload["prompt_tokens"] == 0 - assert payload["completion_tokens"] == 0 - assert payload["total_tokens"] == 0 - - def test_ocr_call_with_zero_pages(self, mock_datetime, base_kwargs): - """Test OCR call with zero pages processed""" - response_obj = { - "id": "ocr-test-000", - "object": "ocr", - "model": "test-ocr-model", - "usage_info": {"pages_processed": 0}, - } - - payload = get_logging_payload( - kwargs=base_kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "ocr" - assert payload["prompt_tokens"] == 0 - assert payload["completion_tokens"] == 0 - assert payload["total_tokens"] == 0 - - # Verify pages_processed is 0 - import json - - metadata = json.loads(payload["metadata"]) - assert metadata["additional_usage_values"]["pages_processed"] == 0 - - def test_non_ocr_call_uses_token_based_usage(self, mock_datetime): - """Test that non-OCR calls still use token-based usage""" - kwargs = { - "model": "gpt-5.5", - "call_type": "completion", - "litellm_params": {}, - "response_cost": 0.02, - } - - response_obj = { - "id": "completion-test-123", - "object": "chat.completion", - "model": "gpt-5.5", - "usage": { - "prompt_tokens": 50, - "completion_tokens": 100, - "total_tokens": 150, - }, - } - - payload = get_logging_payload( - kwargs=kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "completion" - assert payload["prompt_tokens"] == 50 - assert payload["completion_tokens"] == 100 - assert payload["total_tokens"] == 150 - - def test_ocr_with_metadata(self, mock_datetime, base_kwargs): - """Test OCR call with additional metadata""" - base_kwargs["litellm_params"] = { - "metadata": { - "user_api_key_user_id": "test-user", - "user_api_key_team_id": "test-team", - } - } - - response_obj = { - "id": "ocr-metadata-test", - "object": "ocr", - "model": "test-ocr-model", - "usage_info": {"pages_processed": 5, "doc_size_bytes": 1024}, - } - - payload = get_logging_payload( - kwargs=base_kwargs, - response_obj=response_obj, - start_time=mock_datetime, - end_time=mock_datetime, - ) - - assert payload["call_type"] == "ocr" - assert payload["user"] == "test-user" - assert payload["prompt_tokens"] == 0 - assert payload["completion_tokens"] == 0 - - # Verify pages_processed and doc_size_bytes are both in additional_usage_values - import json - - metadata = json.loads(payload["metadata"]) - assert metadata["additional_usage_values"]["pages_processed"] == 5 - assert metadata["additional_usage_values"]["doc_size_bytes"] == 1024 diff --git a/tests/test_anthropic_compaction_usage.py b/tests/test_anthropic_compaction_usage.py deleted file mode 100644 index 1758a94fffc..00000000000 --- a/tests/test_anthropic_compaction_usage.py +++ /dev/null @@ -1,96 +0,0 @@ -from litellm.llms.anthropic.chat.transformation import AnthropicConfig - - -def test_anthropic_compaction_usage_calculation(): - """ - Test that calculate_usage correctly sums tokens from the iterations array - as requested in Issue #27060. - """ - anthropic_config = AnthropicConfig() - - # Mock usage object with compaction iterations - usage_object = { - "input_tokens": 100, # Top-level (excludes compaction) - "output_tokens": 50, # Top-level (excludes compaction) - "iterations": [ - { - "iteration": 1, - "type": "compaction", - "input_tokens": 1000, - "output_tokens": 500, - }, - { - "iteration": 2, - "type": "message", - "input_tokens": 100, - "output_tokens": 50, - }, - ], - } - - usage = anthropic_config.calculate_usage( - usage_object=usage_object, reasoning_content=None - ) - - # Assertions - # Total prompt tokens should be 1000 + 100 = 1100 - assert usage.prompt_tokens == 1100 - # Total completion tokens should be 500 + 50 = 550 - assert usage.completion_tokens == 550 - # Total tokens should be 1650 - assert usage.total_tokens == 1650 - - # Assert details - assert usage.prompt_tokens_details.text_tokens == 1100 - - # Assert iterations passthrough - assert usage.iterations is not None - assert len(usage.iterations) == 2 - assert usage.iterations[0]["type"] == "compaction" - - -def test_anthropic_compaction_usage_with_iteration_cache(): - """ - Test that calculate_usage correctly sums caching tokens FROM iterations. - This covers the specific case mentioned by JasonPan. - """ - anthropic_config = AnthropicConfig() - - usage_object = { - "input_tokens": 100, - "output_tokens": 50, - "iterations": [ - { - "type": "compaction", - "input_tokens": 500, - "output_tokens": 200, - "cache_creation_input_tokens": 50, - "cache_read_input_tokens": 17000, - }, - { - "type": "message", - "input_tokens": 100, - "output_tokens": 50, - "cache_creation_input_tokens": 10, - "cache_read_input_tokens": 20, - }, - ], - } - - usage = anthropic_config.calculate_usage( - usage_object=usage_object, reasoning_content=None - ) - - # input_tokens sum = 500 + 100 = 600 - # cache_creation sum = 50 + 10 = 60 - # cache_read sum = 17000 + 20 = 17020 - # Total prompt tokens = 600 + 60 + 17020 = 17680 - assert usage.prompt_tokens == 17680 - assert usage.completion_tokens == 250 - assert usage.prompt_tokens_details.cache_creation_tokens == 60 - assert usage.prompt_tokens_details.cached_tokens == 17020 - - -if __name__ == "__main__": - test_anthropic_compaction_usage_calculation() - test_anthropic_compaction_usage_with_iteration_cache() diff --git a/tests/test_budget_management.py b/tests/test_budget_management.py deleted file mode 100644 index 42d5c5d98ac..00000000000 --- a/tests/test_budget_management.py +++ /dev/null @@ -1,102 +0,0 @@ -import os -# What is this? -## Unit tests for the /budget/* endpoints -from litellm._uuid import uuid -from datetime import datetime, timezone - -import aiohttp -import pytest -import pytest_asyncio - -from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time -from litellm.proxy.common_utils.timezone_utils import get_budget_reset_timezone - - -def _parse_budget_api_datetime(value: str) -> datetime: - """Parse ISO timestamps returned by the proxy JSON API.""" - if value.endswith("Z"): - value = value[:-1] + "+00:00" - dt = datetime.fromisoformat(value) - if dt.tzinfo is None: - dt = dt.replace(tzinfo=timezone.utc) - return dt - - -async def delete_budget(session, budget_id): - url = "http://0.0.0.0:4000/budget/delete" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = {"id": budget_id} - async with session.post(url, headers=headers, json=data) as response: - assert response.status == 200 - print(f"Deleted Budget {budget_id}") - - -async def create_budget(session, data): - url = "http://0.0.0.0:4000/budget/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - - async with session.post(url, headers=headers, json=data) as response: - assert response.status == 200 - response_data = await response.json() - budget_id = response_data["budget_id"] - print(f"Created Budget {budget_id}") - return response_data - - -@pytest_asyncio.fixture -async def budget_setup(): - """ - Fixture to create a budget for testing and clean it up afterward. - - This fixture performs the following steps: - 1. Opens an aiohttp ClientSession. - 2. Generates a random budget_id and defines the budget data (duration: 1 day, max_budget: 0.02). - 3. Calls create_budget to create the budget. - 4. Yields the budget_response (a dict) for use in the test. - 5. After the test completes, deletes the created budget by calling delete_budget. - - Returns: - dict: The JSON response from create_budget, which includes the created budget's data. - """ - - async with aiohttp.ClientSession() as session: - # Generate a unique budget_id and define the budget data. - budget_id = f"budget-{uuid.uuid4()}" - data = {"budget_id": budget_id, "budget_duration": "1d", "max_budget": 0.02} - budget_response = await create_budget(session, data) - - # Yield the response so the test can use it. - yield budget_response - - # After the test, delete the created budget to clean up. - await delete_budget(session, budget_id) - - -@pytest.mark.asyncio -async def test_create_budget_with_duration(budget_setup): - """ - Test creating a budget with a specified duration and verify that 'budget_reset_at' - matches the next standardized reset (see get_budget_reset_time / new_budget), not - necessarily created_at + wall-clock duration. - """ - - assert ( - budget_setup["budget_reset_at"] is not None - ), "The budget_reset_at field should not be None" - - created_at = _parse_budget_api_datetime(budget_setup["created_at"]) - expected_reset_at = get_next_standardized_reset_time( - duration=budget_setup["budget_duration"], - current_time=created_at, - timezone_str=get_budget_reset_timezone(), - ) - - actual_reset_at = _parse_budget_api_datetime(budget_setup["budget_reset_at"]) - - tolerance_seconds = 3 - time_difference = abs((actual_reset_at - expected_reset_at).total_seconds()) - - assert time_difference <= tolerance_seconds, ( - f"Expected budget_reset_at to be within {tolerance_seconds} seconds of {expected_reset_at}, " - f"but the difference was {time_difference} seconds." - ) diff --git a/tests/test_callbacks_on_proxy.py b/tests/test_callbacks_on_proxy.py deleted file mode 100644 index 42aa6d98cb9..00000000000 --- a/tests/test_callbacks_on_proxy.py +++ /dev/null @@ -1,303 +0,0 @@ -# What this tests ? -## Makes sure the number of callbacks on the proxy don't increase over time -## Num callbacks should be a fixed number at t=0 and t=10, t=20 -""" -PROD TEST - DO NOT Delete this Test -""" - -import pytest -import asyncio -import aiohttp -import os -import re -import dotenv -from collections import Counter -from dotenv import load_dotenv - -load_dotenv() - -# A *leak* is sustained, monotonic growth of one callback TYPE across the whole -# sampling window. A one-time bump that then plateaus is benign pollution from -# other tests sharing this proxy (this suite runs `pytest -n 4` against a single -# proxy container, so other workers legitimately add team/key-scoped callbacks -# while this test sleeps). We therefore sample N times and only flag a type -# whose normalized count never decreases, grows in >=2 distinct intervals, and -# nets >= LEAK_MIN_NET_GROWTH overall. -NUM_SAMPLES = 4 -SAMPLE_INTERVAL_SECONDS = 20 -LEAK_MIN_NET_GROWTH = 5 -LEAK_MIN_GROWING_INTERVALS = 2 -# A routing-strategy switch / alerting config is a *known, bounded, one-time* -# registration (CCI diagnostic 2026-05-16: total 85->95 on the first interval -# after switching to latency-based-routing, then flat at 95 for 2.5 min under -# load). We absorb that step by settling before the baseline sample, so only -# growth *after* the deliberate perturbation can count as a leak. -SETTLE_SECONDS = 30 - -# Strip instance-identity noise so N leaked instances of one class collapse to -# one rising counter instead of N opaque, unrelated-looking strings. -_ADDR_RE = re.compile(r" at 0x[0-9a-fA-F]+") -_OBJ_RE = re.compile(r"<([\w.]+) object") - - -def _normalize_callback(cb_str: str) -> str: - """Reduce a callback's str() to a stable type key (drops 0x… addresses).""" - s = _ADDR_RE.sub("", cb_str) - m = _OBJ_RE.search(s) - if m: - return m.group(1).split(".")[-1] - # bound methods: ">" -> "Cls.m" - bm = re.search(r"bound method ([\w.]+)", s) - if bm: - return bm.group(1) - return s.strip() - - -def _summarize(all_litellm_callbacks) -> Counter: - return Counter(_normalize_callback(str(c)) for c in all_litellm_callbacks) - - -def _detect_leaks(samples): - """ - samples: list[Counter] taken in time order. - - Returns {callback_type: [counts across samples]} for types that grew - monotonically (never decreased), in >=LEAK_MIN_GROWING_INTERVALS intervals, - and netted >=LEAK_MIN_NET_GROWTH overall — i.e. a real leak, not a one-shot - step from a parallel test. - """ - leaks = {} - all_types = set().union(*[set(s) for s in samples]) if samples else set() - for t in all_types: - series = [s.get(t, 0) for s in samples] - deltas = [b - a for a, b in zip(series, series[1:])] - net = series[-1] - series[0] - non_decreasing = all(d >= 0 for d in deltas) - growing_intervals = sum(1 for d in deltas if d > 0) - if ( - non_decreasing - and net >= LEAK_MIN_NET_GROWTH - and growing_intervals >= LEAK_MIN_GROWING_INTERVALS - ): - leaks[t] = series - return leaks - - -def _terminal_suspects(samples): - """ - Types whose net growth clears the threshold monotonically but is confined - to the *final* interval — `growing_intervals == 1` with that one growing - interval being the last. `_detect_leaks`' `>= 2` guard silently passes - these, so a real leak that accumulates entirely in the last sampled window - is indistinguishable from a one-time terminal step *without one more - sample*. Returns the set of such types so the caller can re-confirm. - """ - suspects = set() - all_types = set().union(*[set(s) for s in samples]) if samples else set() - for t in all_types: - series = [s.get(t, 0) for s in samples] - deltas = [b - a for a, b in zip(series, series[1:])] - if not deltas: - continue - net = series[-1] - series[0] - non_decreasing = all(d >= 0 for d in deltas) - growing = [i for i, d in enumerate(deltas) if d > 0] - if ( - non_decreasing - and net >= LEAK_MIN_NET_GROWTH - and growing == [len(deltas) - 1] - ): - suspects.add(t) - return suspects - - -async def _detect_leaks_confirmed(session, samples): - """ - `_detect_leaks`, plus a single confirmation sample when growth is confined - to the final interval (see `_terminal_suspects`). A genuine ongoing leak - keeps climbing -> now grows in >= 2 intervals -> flagged; a one-time - terminal registration plateaus -> still 1 growing interval -> ignored. - Returns `(leaks, samples)` (samples may have one extra entry appended). - """ - leaks = _detect_leaks(samples) - if not leaks and _terminal_suspects(samples): - await asyncio.sleep(SAMPLE_INTERVAL_SECONDS) - _, _, all_cb = await get_active_callbacks(session=session) - samples = samples + [_summarize(all_cb)] - leaks = _detect_leaks(samples) - return leaks, samples - - -def _format_report(samples, leaks) -> str: - lines = ["Callback count per type across samples (time order):"] - all_types = sorted(set().union(*[set(s) for s in samples])) - for t in all_types: - series = [s.get(t, 0) for s in samples] - marker = " <-- LEAK" if t in leaks else "" - lines.append(f" {t}: {series}{marker}") - totals = [sum(s.values()) for s in samples] - lines.append(f"TOTAL callbacks per sample: {totals}") - if leaks: - lines.append( - "Leaking callback types (sustained monotonic growth): " - + ", ".join(sorted(leaks)) - ) - return "\n".join(lines) - - -async def _sample_callbacks(session, num_samples, interval): - """Take `num_samples` callback snapshots `interval`s apart.""" - samples = [] - alerts = [] - for i in range(num_samples): - if i > 0: - await asyncio.sleep(interval) - num_cb, num_alert, all_cb = await get_active_callbacks(session=session) - samples.append(_summarize(all_cb)) - alerts.append(num_alert) - return samples, alerts - - -async def config_update(session, routing_strategy=None): - url = "http://0.0.0.0:4000/config/update" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - print("routing_strategy: ", routing_strategy) - data = { - "router_settings": { - "routing_strategy": routing_strategy, - }, - "general_settings": { - "alert_to_webhook_url": {"llm_exceptions": "example-slack-webhook-url"}, - "alert_types": ["llm_exceptions", "db_exceptions"], - }, - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - -async def get_active_callbacks(session): - url = "http://0.0.0.0:4000/active/callbacks" - headers = { - "Content-Type": "application/json", - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - } - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - print("response from /active/callbacks") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - _json_response = await response.json() - - _num_callbacks = _json_response["num_callbacks"] - _num_alerts = _json_response["num_alerting"] - all_litellm_callbacks = _json_response["all_litellm_callbacks"] - - print("current number of callbacks: ", _num_callbacks) - print("current number of alerts: ", _num_alerts) - return _num_callbacks, _num_alerts, all_litellm_callbacks - - -async def get_current_routing_strategy(session): - url = "http://0.0.0.0:4000/get/config/callbacks" - headers = { - "Content-Type": "application/json", - "Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", - } - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - _json_response = await response.json() - print("JSON response: ", _json_response) - - router_settings = _json_response["router_settings"] - print("Router settings: ", router_settings) - routing_strategy = router_settings["routing_strategy"] - return routing_strategy - - -@pytest.mark.asyncio -@pytest.mark.order1 -@pytest.mark.flaky(reruns=2, reruns_delay=5) -async def test_check_num_callbacks(): - """ - PROD invariant: no callback TYPE should grow without bound over time. - - This suite runs `pytest -n 4` against one shared proxy, so the raw count is - noisy — other workers legitimately add team/key-scoped callbacks that then - plateau. We settle first, then sample several times, and only fail on - *sustained, monotonic* per-type growth (a genuine leak), naming the type. - """ - async with aiohttp.ClientSession() as session: - # Absorb proxy warmup / in-flight parallel registration before baseline. - await asyncio.sleep(SETTLE_SECONDS) - - samples, _ = await _sample_callbacks( - session, NUM_SAMPLES, SAMPLE_INTERVAL_SECONDS - ) - - assert sum(samples[0].values()) > 0, "expected some callbacks registered" - - leaks, samples = await _detect_leaks_confirmed(session, samples) - report = _format_report(samples, leaks) - print(report) - assert not leaks, f"Callback leak detected.\n{report}" - - -@pytest.mark.asyncio -@pytest.mark.order2 -@pytest.mark.flaky(reruns=2, reruns_delay=5) -async def test_check_num_callbacks_on_lowest_latency(): - """ - Same PROD invariant as test_check_num_callbacks, but after switching the - router to latency-based-routing. That switch is a *known, bounded* one-time - registration (it adds the latency strategy handler + Slack alerting); we - settle past it before baselining so only post-switch growth counts as a - leak. Also asserts the alerting count is stable. - """ - async with aiohttp.ClientSession() as session: - await asyncio.sleep(30) - - original_routing_strategy = await get_current_routing_strategy(session=session) - await config_update(session=session, routing_strategy="latency-based-routing") - - try: - # Absorb the deliberate one-time config/update registration step. - await asyncio.sleep(SETTLE_SECONDS) - - samples, alerts = await _sample_callbacks( - session, NUM_SAMPLES, SAMPLE_INTERVAL_SECONDS - ) - - leaks, samples = await _detect_leaks_confirmed(session, samples) - report = _format_report(samples, leaks) - print(report) - assert not leaks, f"Callback leak detected.\n{report}" - assert ( - len(set(alerts)) == 1 - ), f"alerting count changed across samples: {alerts}" - finally: - await config_update( - session=session, routing_strategy=original_routing_strategy - ) diff --git a/tests/test_default_encoding_non_root.py b/tests/test_default_encoding_non_root.py deleted file mode 100644 index 06a5de51976..00000000000 --- a/tests/test_default_encoding_non_root.py +++ /dev/null @@ -1,55 +0,0 @@ -import importlib -import os -from unittest.mock import MagicMock, patch - -import litellm.litellm_core_utils.default_encoding as default_encoding - - -def _reload_default_encoding(monkeypatch, **env_overrides): - """ - Helper to reload default_encoding with a clean TIKTOKEN_CACHE_DIR and - specific environment overrides. - """ - monkeypatch.delenv("TIKTOKEN_CACHE_DIR", raising=False) - monkeypatch.delenv("CUSTOM_TIKTOKEN_CACHE_DIR", raising=False) - for key, value in env_overrides.items(): - monkeypatch.setenv(key, value) - importlib.reload(default_encoding) - - -def test_default_encoding_uses_bundled_tokenizers_by_default(monkeypatch): - """ - TIKTOKEN_CACHE_DIR should point at the bundled tokenizers directory - when no CUSTOM_TIKTOKEN_CACHE_DIR is set, even in non-root environments. - """ - _reload_default_encoding(monkeypatch, LITELLM_NON_ROOT="true") - - assert "TIKTOKEN_CACHE_DIR" in os.environ - cache_dir = os.environ["TIKTOKEN_CACHE_DIR"] - assert "tokenizers" in cache_dir - - -def test_custom_tiktoken_cache_dir_override(monkeypatch, tmp_path): - """ - CUSTOM_TIKTOKEN_CACHE_DIR must override the default bundled directory - and the directory should be created if it does not exist. - Reload with an empty custom dir would otherwise trigger tiktoken to - download the vocab; we patch get_encoding so the test is offline-safe - and does not depend on tiktoken's in-memory cache state. - """ - custom_dir = tmp_path / "tiktoken_cache" - with patch( - "litellm.litellm_core_utils.default_encoding.tiktoken.get_encoding", - return_value=MagicMock(), - ): - _reload_default_encoding(monkeypatch, CUSTOM_TIKTOKEN_CACHE_DIR=str(custom_dir)) - - cache_dir = os.environ.get("TIKTOKEN_CACHE_DIR") - assert cache_dir == str(custom_dir) - assert os.path.isdir(cache_dir) - - # Restore module to a clean state so default_encoding.encoding is a real - # tiktoken Encoding, not the MagicMock, for any test that runs after this. - monkeypatch.delenv("TIKTOKEN_CACHE_DIR", raising=False) - monkeypatch.delenv("CUSTOM_TIKTOKEN_CACHE_DIR", raising=False) - importlib.reload(default_encoding) diff --git a/tests/test_end_users.py b/tests/test_end_users.py deleted file mode 100644 index 03fb4c9e86c..00000000000 --- a/tests/test_end_users.py +++ /dev/null @@ -1,211 +0,0 @@ -import os -# What is this? -## Unit tests for the /end_users/* endpoints -import pytest -import asyncio -import aiohttp -import time -from litellm._uuid import uuid -from openai import AsyncOpenAI -from typing import Optional - -""" -- `/end_user/new` -- `/end_user/info` -""" - - -async def generate_key( - session, - i, - budget=None, - budget_duration=None, - models=["azure-models", "gpt-4", "dall-e-3"], - max_parallel_requests: Optional[int] = None, - user_id: Optional[str] = None, - team_id: Optional[str] = None, - calling_key=os.environ["LITELLM_MASTER_KEY"], -): - url = "http://0.0.0.0:4000/key/generate" - headers = { - "Authorization": f"Bearer {calling_key}", - "Content-Type": "application/json", - } - data = { - "models": models, - "aliases": {"mistral-7b": "gpt-3.5-turbo"}, - "duration": None, - "max_budget": budget, - "budget_duration": budget_duration, - "max_parallel_requests": max_parallel_requests, - "user_id": user_id, - "team_id": team_id, - } - - print(f"data: {data}") - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response {i} (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request {i} did not return a 200 status code: {status}") - - return await response.json() - - -async def new_end_user( - session, - i, - user_id=str(uuid.uuid4()), - model_region=None, - default_model=None, - budget_id=None, -): - url = "http://0.0.0.0:4000/end_user/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = { - "user_id": user_id, - "allowed_model_region": model_region, - "default_model": default_model, - } - - if budget_id is not None: - data["budget_id"] = budget_id - print("end user data: {}".format(data)) - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response {i} (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request {i} did not return a 200 status code: {status}") - - return await response.json() - - -async def new_budget(session, i, budget_id=None): - url = "http://0.0.0.0:4000/budget/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = { - "budget_id": budget_id, - "tpm_limit": 2, - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response {i} (Status code: {status}):") - print(response_text) - print() - - -@pytest.mark.asyncio -async def test_end_user_new(): - """ - Make 20 parallel calls to /user/new. Assert all worked. - """ - async with aiohttp.ClientSession() as session: - tasks = [new_end_user(session, i, str(uuid.uuid4())) for i in range(1, 11)] - await asyncio.gather(*tasks) - - -@pytest.mark.asyncio -async def test_enduser_tpm_limits_non_master_key(): - """ - 1. budget_id = Create Budget with tpm_limit = 10 - 2. create end_user with budget_id - 3. Make /chat/completions calls - 4. Sleep 1 second - 4. Make /chat/completions call -> expect this to fail because rate limit hit - """ - async with aiohttp.ClientSession() as session: - # create a budget with budget_id = "free-tier" - budget_id = f"free-tier-{uuid.uuid4()}" - await new_budget(session, 0, budget_id=budget_id) - await asyncio.sleep(2) - - end_user_id = str(uuid.uuid4()) - - await new_end_user( - session=session, i=0, user_id=end_user_id, budget_id=budget_id - ) - - ## MAKE CALL ## - key_gen = await generate_key(session=session, i=0, models=[]) - - key = key_gen["key"] - - # chat completion 1 - client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000", max_retries=0) - - # chat completion 2 - passed = 0 - for _ in range(10): - try: - result = await client.chat.completions.create( - model="fake-openai-endpoint", - messages=[{"role": "user", "content": "Hey!"}], - user=end_user_id, - ) - passed += 1 - except Exception: - pass - print("Passed requests=", passed) - - assert ( - passed < 5 - ), f"Sent 10 requests and end-user has tpm_limit of 2. Number requests passed: {passed}. Expected less than 5 to pass" - - -@pytest.mark.asyncio -async def test_enduser_tpm_limits_with_master_key(): - """ - 1. budget_id = Create Budget with tpm_limit = 10 - 2. create end_user with budget_id - 3. Make /chat/completions calls - 4. Sleep 1 second - 4. Make /chat/completions call -> expect this to fail because rate limit hit - """ - async with aiohttp.ClientSession() as session: - # create a budget with budget_id = "free-tier" - budget_id = f"free-tier-{uuid.uuid4()}" - await new_budget(session, 0, budget_id=budget_id) - - end_user_id = str(uuid.uuid4()) - - await new_end_user( - session=session, i=0, user_id=end_user_id, budget_id=budget_id - ) - - # chat completion 1 - client = AsyncOpenAI( - api_key=os.environ["LITELLM_MASTER_KEY"], base_url="http://0.0.0.0:4000", max_retries=0 - ) - - # chat completion 2 - passed = 0 - for _ in range(10): - try: - result = await client.chat.completions.create( - model="fake-openai-endpoint", - messages=[{"role": "user", "content": "Hey!"}], - user=end_user_id, - ) - passed += 1 - except Exception: - pass - print("Passed requests=", passed) - - assert ( - passed < 5 - ), f"Sent 10 requests and end-user has tpm_limit of 2. Number requests passed: {passed}. Expected less than 5 to pass" diff --git a/tests/test_fallbacks.py b/tests/test_fallbacks.py deleted file mode 100644 index 8418847d062..00000000000 --- a/tests/test_fallbacks.py +++ /dev/null @@ -1,325 +0,0 @@ -import os -from typing import Final - -# What is this? -## This tests if the proxy fallbacks work as expected -import pytest -import asyncio -import aiohttp -from tests.large_text import text -import time -from typing import Optional -from openai import AsyncOpenAI, PermissionDeniedError - -PROXY_BASE_URL: Final = os.environ.get("LITELLM_PROXY_BASE_URL", "http://0.0.0.0:4000") - - -async def generate_key( - session, - i, - models: list, - calling_key=os.environ["LITELLM_MASTER_KEY"], -): - url: Final = f"{PROXY_BASE_URL}/key/generate" - headers = { - "Authorization": f"Bearer {calling_key}", - "Content-Type": "application/json", - } - data = { - "models": models, - } - - print(f"data: {data}") - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response {i} (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request {i} did not return a 200 status code: {status}") - - return await response.json() - - -async def chat_completion( - session, - key: str, - model: str, - messages: list, - return_headers: bool = False, - extra_headers: Optional[dict] = None, - **kwargs, -): - url: Final = f"{PROXY_BASE_URL}/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - if extra_headers is not None: - headers.update(extra_headers) - data = {"model": model, "messages": messages, **kwargs} - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - if return_headers: - return None, response.headers - else: - raise Exception(f"Request did not return a 200 status code: {status}") - - if return_headers: - return await response.json(), response.headers - else: - return await response.json() - - -@pytest.mark.parametrize("has_access", [True, False]) -@pytest.mark.asyncio -async def test_chat_completion_client_fallbacks(has_access: bool) -> None: - models: Final = ["gpt-3.5-turbo", "gpt-6-luna"] if has_access else ["gpt-3.5-turbo"] - async with aiohttp.ClientSession() as session: - generated_key: Final = await generate_key(session=session, i=0, models=models) - async with AsyncOpenAI(api_key=generated_key["key"], base_url=PROXY_BASE_URL, max_retries=0) as client: - request: Final = { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": "Who was Alexander?"}], - "max_tokens": 32, - "temperature": 0, - "extra_body": { - "mock_testing_fallbacks": True, - "fallbacks": ["gpt-6-luna"], - }, - } - if not has_access: - with pytest.raises(PermissionDeniedError) as denied: - await client.chat.completions.create(**request) - assert denied.value.status_code == 403 - assert "gpt-6-luna" in str(denied.value) - return - response: Final = await client.chat.completions.create(**request) - assert response.model == "gpt-6-luna" - assert response.choices[0].message.content - - -@pytest.mark.asyncio -async def test_chat_completion_with_retries(): - """ - make chat completion call with prompt > context window. expect it to work with fallback - """ - async with aiohttp.ClientSession() as session: - model = "fake-openai-endpoint-4" - messages = [ - {"role": "system", "content": text}, - {"role": "user", "content": "Who was Alexander?"}, - ] - response, headers = await chat_completion( - session=session, - key=os.environ["LITELLM_MASTER_KEY"], - model=model, - messages=messages, - mock_testing_rate_limit_error=True, - return_headers=True, - ) - print(f"headers: {headers}") - assert headers["x-litellm-attempted-retries"] == "1" - assert headers["x-litellm-max-retries"] == "50" - - -@pytest.mark.asyncio -async def test_chat_completion_with_fallbacks(): - """ - make chat completion call with prompt > context window. expect it to work with fallback - """ - async with aiohttp.ClientSession() as session: - model = "badly-configured-openai-endpoint" - messages = [ - {"role": "system", "content": text}, - {"role": "user", "content": "Who was Alexander?"}, - ] - response, headers = await chat_completion( - session=session, - key=os.environ["LITELLM_MASTER_KEY"], - model=model, - messages=messages, - fallbacks=["fake-openai-endpoint-5"], - return_headers=True, - ) - print(f"headers: {headers}") - assert headers["x-litellm-attempted-fallbacks"] == "1" - - -@pytest.mark.asyncio -async def test_chat_completion_with_timeout(): - """ - make chat completion call with low timeout and `mock_timeout`: true. Expect it to fail and correct timeout to be set in headers. - """ - async with aiohttp.ClientSession() as session: - model = "fake-openai-endpoint-5" - messages = [ - {"role": "system", "content": text}, - {"role": "user", "content": "Who was Alexander?"}, - ] - start_time = time.time() - response, headers = await chat_completion( - session=session, - key=os.environ["LITELLM_MASTER_KEY"], - model=model, - messages=messages, - num_retries=0, - mock_timeout=True, - return_headers=True, - ) - end_time = time.time() - print(f"headers: {headers}") - assert ( - headers["x-litellm-timeout"] == "1.0" - ) # assert model-specific timeout used - - -@pytest.mark.asyncio -async def test_chat_completion_with_timeout_from_request(): - """ - make chat completion call with low timeout and `mock_timeout`: true. Expect it to fail and correct timeout to be set in headers. - """ - async with aiohttp.ClientSession() as session: - model = "fake-openai-endpoint-5" - messages = [ - {"role": "system", "content": text}, - {"role": "user", "content": "Who was Alexander?"}, - ] - extra_headers = { - "x-litellm-timeout": "0.001", - } - start_time = time.time() - response, headers = await chat_completion( - session=session, - key=os.environ["LITELLM_MASTER_KEY"], - model=model, - messages=messages, - num_retries=0, - mock_timeout=True, - extra_headers=extra_headers, - return_headers=True, - ) - end_time = time.time() - print(f"headers: {headers}") - assert ( - headers["x-litellm-timeout"] == "0.001" - ) # assert model-specific timeout used - - -@pytest.mark.parametrize("has_access", [True, False]) -@pytest.mark.asyncio -async def test_chat_completion_client_fallbacks_with_custom_message(has_access: bool) -> None: - original_messages: Final = [{"role": "user", "content": "Who was Alexander?"}] - custom_messages: Final = [ - { - "role": "user", - "content": ( - "Describe the weather in a coastal city during winter, including the usual temperature, rain, wind, " - "and the clothing a visitor should bring." - ), - } - ] - models: Final = ["gpt-3.5-turbo", "gpt-6-luna"] if has_access else ["gpt-3.5-turbo"] - async with aiohttp.ClientSession() as session: - generated_key: Final = await generate_key(session=session, i=0, models=models) - async with AsyncOpenAI(api_key=generated_key["key"], base_url=PROXY_BASE_URL, max_retries=0) as client: - request: Final = { - "model": "gpt-3.5-turbo", - "messages": original_messages, - "max_tokens": 32, - "temperature": 0, - "extra_body": { - "mock_testing_fallbacks": True, - "fallbacks": [ - { - "model": "gpt-6-luna", - "messages": custom_messages, - } - ], - }, - } - if not has_access: - with pytest.raises(PermissionDeniedError) as denied: - await client.chat.completions.create(**request) - assert denied.value.status_code == 403 - assert "gpt-6-luna" in str(denied.value) - return - response: Final = await client.chat.completions.create(**request) - assert response.model == "gpt-6-luna" - assert response.choices[0].message.content - custom_control: Final = await client.chat.completions.create( - model="gpt-6-luna", - messages=custom_messages, - max_tokens=32, - temperature=0, - ) - original_control: Final = await client.chat.completions.create( - model="gpt-6-luna", - messages=original_messages, - max_tokens=32, - temperature=0, - ) - assert response.usage is not None - assert custom_control.usage is not None - assert original_control.usage is not None - assert custom_control.usage.completion_tokens > 0 - assert original_control.usage.completion_tokens > 0 - assert custom_control.usage.prompt_tokens != original_control.usage.prompt_tokens - assert response.usage.prompt_tokens == custom_control.usage.prompt_tokens - - -from typing import List - - -async def make_request(client: AsyncOpenAI, model: str) -> bool: - try: - await client.chat.completions.create( - model=model, - messages=[{"role": "user", "content": "Who was Alexander?"}], - ) - return True - except Exception as e: - print(f"Error with {model}: {str(e)}") - return False - - -async def run_good_model_test(client: AsyncOpenAI, num_requests: int) -> bool: - tasks = [make_request(client, "good-model") for _ in range(num_requests)] - good_results = await asyncio.gather(*tasks) - return all(good_results) - - -@pytest.mark.asyncio -async def test_chat_completion_bad_and_good_model(): - """ - Prod test - ensure even if bad model is down, good model is still working. - """ - client = AsyncOpenAI(api_key=os.environ["LITELLM_MASTER_KEY"], base_url="http://0.0.0.0:4000") - num_requests = 100 - num_iterations = 3 - - for iteration in range(num_iterations): - print(f"\nIteration {iteration + 1}/{num_iterations}") - start_time = time.time() - - # Fire and forget bad model requests - for _ in range(num_requests): - asyncio.create_task(make_request(client, "bad-model")) - - # Wait only for good model requests - success = await run_good_model_test(client, num_requests) - print( - f"Iteration {iteration + 1}: {'✓' if success else '✗'} ({time.time() - start_time:.2f}s)" - ) - assert success, "Not all good model requests succeeded" diff --git a/tests/test_gpt5_azure_temperature_support.py b/tests/test_gpt5_azure_temperature_support.py deleted file mode 100644 index 025b921236a..00000000000 --- a/tests/test_gpt5_azure_temperature_support.py +++ /dev/null @@ -1,102 +0,0 @@ -""" -Test that Azure GPT-5 models support temperature parameter in Responses API. -""" - -import pytest -from litellm.utils import ProviderConfigManager -from litellm.types.utils import LlmProviders - - -def test_azure_gpt5_supports_temperature(): - """Test that Azure GPT-5 uses the correct config that supports temperature.""" - config = ProviderConfigManager.get_provider_responses_api_config( - provider=LlmProviders.AZURE, model="gpt-5" - ) - - # Should use AzureOpenAIResponsesAPIConfig, not AzureOpenAIOSeriesResponsesAPIConfig - assert type(config).__name__ == "AzureOpenAIResponsesAPIConfig" - - # Should support temperature parameter - supported_params = config.get_supported_openai_params("gpt-5") - assert ( - "temperature" in supported_params - ), "Azure GPT-5 should support temperature parameter" - - -def test_azure_o_series_does_not_support_temperature(): - """Test that Azure O-series models still use the correct O-series config.""" - test_models = ["o1", "o3"] - - for model in test_models: - config = ProviderConfigManager.get_provider_responses_api_config( - provider=LlmProviders.AZURE, model=model - ) - - # Should use AzureOpenAIOSeriesResponsesAPIConfig - assert ( - type(config).__name__ == "AzureOpenAIOSeriesResponsesAPIConfig" - ), f"Azure {model} should use O-series config" - - # Should NOT support temperature parameter - supported_params = config.get_supported_openai_params(model) - assert ( - "temperature" not in supported_params - ), f"Azure {model} should NOT support temperature parameter" - - -def test_openai_gpt5_supports_temperature(): - """Test that OpenAI GPT-5 supports temperature parameter.""" - config = ProviderConfigManager.get_provider_responses_api_config( - provider=LlmProviders.OPENAI, model="gpt-5" - ) - - # Should use OpenAIResponsesAPIConfig - assert type(config).__name__ == "OpenAIResponsesAPIConfig" - - # Should support temperature parameter - supported_params = config.get_supported_openai_params("gpt-5") - assert ( - "temperature" in supported_params - ), "OpenAI GPT-5 should support temperature parameter" - - -def test_azure_gpt5_variants_support_temperature(): - """Test that various GPT-5 model name variants support temperature.""" - gpt5_variants = ["gpt-5", "gpt-5-turbo", "GPT-5", "azure/gpt-5"] - - for model in gpt5_variants: - config = ProviderConfigManager.get_provider_responses_api_config( - provider=LlmProviders.AZURE, model=model - ) - - # All GPT-5 variants should use the base config, not O-series config - assert ( - type(config).__name__ == "AzureOpenAIResponsesAPIConfig" - ), f"Model '{model}' should not use O-series config" - - # All should support temperature - supported_params = config.get_supported_openai_params(model) - assert ( - "temperature" in supported_params - ), f"Model '{model}' should support temperature parameter" - - -def test_azure_gpt_models_support_temperature(): - """Test that all GPT models (gpt-3.5, gpt-4, gpt-5, etc.) support temperature.""" - gpt_models = ["gpt-3.5-turbo", "gpt-4", "gpt-4-turbo", "gpt-4o", "gpt-5"] - - for model in gpt_models: - config = ProviderConfigManager.get_provider_responses_api_config( - provider=LlmProviders.AZURE, model=model - ) - - # All GPT models should use the base config, not O-series config - assert ( - type(config).__name__ == "AzureOpenAIResponsesAPIConfig" - ), f"Model '{model}' should not use O-series config" - - # All should support temperature - supported_params = config.get_supported_openai_params(model) - assert ( - "temperature" in supported_params - ), f"Model '{model}' should support temperature parameter" diff --git a/tests/test_health.py b/tests/test_health.py deleted file mode 100644 index a92c57314b1..00000000000 --- a/tests/test_health.py +++ /dev/null @@ -1,84 +0,0 @@ -import os - -# What this tests? -## Tests /health + /routes endpoints. - -import pytest -import asyncio -import aiohttp - - -async def health(session, call_key): - url = "http://0.0.0.0:4000/health" - headers = { - "Authorization": f"Bearer {call_key}", - "Content-Type": "application/json", - } - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - - print(f"Response (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - return await response.json() - - -async def generate_key(session): - url = "http://0.0.0.0:4000/key/generate" - headers = { - "Authorization": "Bearer " + os.environ["LITELLM_MASTER_KEY"], - "Content-Type": "application/json", - } - data = { - "models": ["gpt-4", "text-embedding-ada-002", "gpt-image-1"], - "duration": None, - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - -@pytest.mark.asyncio -async def test_health(): - """ - - Call /health - """ - async with aiohttp.ClientSession() as session: - # as admin # - all_healthy_models = await health(session=session, call_key=os.environ["LITELLM_MASTER_KEY"]) - total_model_count = ( - all_healthy_models["healthy_count"] + all_healthy_models["unhealthy_count"] - ) - assert total_model_count > 0 - - -@pytest.mark.asyncio -async def test_routes(): - """ - Check if 200 - """ - async with aiohttp.ClientSession() as session: - url = "http://0.0.0.0:4000/routes" - async with session.get(url) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") diff --git a/tests/test_keys.py b/tests/test_keys.py index 835c1f12250..c1862563b0f 100644 --- a/tests/test_keys.py +++ b/tests/test_keys.py @@ -2,57 +2,11 @@ ## Tests /key endpoints. import pytest -import asyncio, uuid +import asyncio import aiohttp -from openai import AsyncOpenAI -import sys, os +import os from typing import Optional -import litellm -from litellm.proxy._types import LitellmUserRoles - - -async def generate_team( - session, models: Optional[list] = None, team_id: Optional[str] = None -): - url = "http://0.0.0.0:4000/team/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - if team_id is None: - team_id = "litellm-dashboard" - data = {"team_id": team_id, **({"models": models} if models is not None else {})} - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response (Status code: {status}):") - print(response_text) - print() - _json_response = await response.json() - return _json_response - - -async def generate_user( - session, - user_role="app_owner", -): - url = "http://0.0.0.0:4000/user/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = { - "user_role": user_role, - "team_id": "litellm-dashboard", - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response (Status code: {status}):") - print(response_text) - print() - _json_response = await response.json() - return _json_response - async def generate_key( session, @@ -99,25 +53,6 @@ async def generate_key( return await response.json() -@pytest.mark.asyncio -async def test_key_gen(): - async with aiohttp.ClientSession() as session: - tasks = [generate_key(session, i) for i in range(1, 11)] - await asyncio.gather(*tasks) - - -@pytest.mark.asyncio -async def test_simple_key_gen(): - async with aiohttp.ClientSession() as session: - key_data = await generate_key(session, i=0) - key = key_data["key"] - assert key_data["token"] is not None - assert key_data["token"] != key - assert key_data["token_id"] is not None - assert key_data["created_at"] is not None - assert key_data["updated_at"] is not None - - @pytest.mark.asyncio async def test_key_gen_bad_key(): """ @@ -147,91 +82,6 @@ async def test_key_gen_bad_key(): pass -async def chat_completion(session, key, model="gpt-4"): - url = "http://0.0.0.0:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": model, - "messages": [ - {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": "Hello!"}, - ], - } - - for i in range(3): - try: - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception( - f"Request did not return a 200 status code: {status}. Response: {response_text}" - ) - - return await response.json() - except Exception as e: - if "Request did not return a 200 status code" in str(e): - raise e - else: - pass - - -async def chat_completion_streaming(session, key, model="gpt-4"): - client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000") - messages = [ - {"role": "system", "content": "You are a helpful assistant"}, - {"role": "user", "content": "Hello!"}, - ] - prompt_tokens = litellm.token_counter(model="gpt-35-turbo", messages=messages) - data = { - "model": model, - "messages": messages, - "stream": True, - } - response = await client.chat.completions.create(**data) - - content = "" - async for chunk in response: - content += chunk.choices[0].delta.content or "" - - print(f"content: {content}") - - completion_tokens = litellm.token_counter( - model="gpt-35-turbo", text=content, count_response_tokens=True - ) - - return prompt_tokens, completion_tokens - - -async def delete_key(session, get_key, auth_key=os.environ["LITELLM_MASTER_KEY"]): - """ - Delete key - """ - url = "http://0.0.0.0:4000/key/delete" - headers = { - "Authorization": f"Bearer {auth_key}", - "Content-Type": "application/json", - } - data = {"keys": [get_key]} - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - async def get_key_info(session, call_key, get_key=None): """ Make sure only models user has access to are returned @@ -262,152 +112,6 @@ async def get_key_info(session, call_key, get_key=None): return await response.json() -async def get_model_list(session, call_key, endpoint: str = "/v1/models"): - """ - Make sure only models user has access to are returned - """ - url = "http://0.0.0.0:4000" + endpoint - headers = { - "Authorization": f"Bearer {call_key}", - "Content-Type": "application/json", - } - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - print(response_text) - print() - - if status != 200: - raise Exception( - f"Request did not return a 200 status code: {status}. Responses {response_text}" - ) - return await response.json() - - -async def get_model_info(session, call_key): - """ - Make sure only models user has access to are returned - """ - url = "http://0.0.0.0:4000/model/info" - headers = { - "Authorization": f"Bearer {call_key}", - "Content-Type": "application/json", - } - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - print(response_text) - print() - - if status != 200: - raise Exception( - f"Request did not return a 200 status code: {status}. Responses {response_text}" - ) - return await response.json() - - -@pytest.mark.asyncio -async def test_key_info(): - """ - Get key info - - as admin -> 200 - - as key itself -> 200 - - as non existent key -> 404 - """ - async with aiohttp.ClientSession() as session: - key_gen = await generate_key(session=session, i=0) - key = key_gen["key"] - # as admin # - await get_key_info(session=session, get_key=key, call_key=os.environ["LITELLM_MASTER_KEY"]) - # as key itself # - await get_key_info(session=session, get_key=key, call_key=key) - - # as key itself, use the auth param, and no query key needed - await get_key_info(session=session, call_key=key) - # as random key # - random_key = f"sk-{uuid.uuid4()}" - status = await get_key_info(session=session, get_key=random_key, call_key=key) - assert status == 404 - - -@pytest.mark.asyncio -async def test_model_info(): - """ - Get model info for models key has access to - """ - async with aiohttp.ClientSession() as session: - key_gen = await generate_key(session=session, i=0) - key = key_gen["key"] - # as admin # - admin_models = await get_model_info(session=session, call_key=os.environ["LITELLM_MASTER_KEY"]) - admin_models = admin_models["data"] - # as key itself # - user_models = await get_model_info(session=session, call_key=key) - user_models = user_models["data"] - - assert len(admin_models) > len(user_models) - assert len(user_models) > 0 - - -async def get_spend_logs(session, request_id): - url = f"http://0.0.0.0:4000/spend/logs?request_id={request_id}" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - -@pytest.mark.asyncio -@pytest.mark.flaky(retries=6, delay=2) -@pytest.mark.skip( - reason="Temporarily skipping due to model change. Will be updated soon." -) -async def test_aaaaakey_info_spend_values_streaming(): - """ - Test to ensure spend is correctly calculated. - - create key - - make completion call - - assert cost is expected value - """ - async with aiohttp.ClientSession() as session: - ## streaming - azure - key_gen = await generate_key(session=session, i=0) - new_key = key_gen["key"] - prompt_tokens, completion_tokens = await chat_completion_streaming( - session=session, key=new_key - ) - print(f"prompt_tokens: {prompt_tokens}, completion_tokens: {completion_tokens}") - prompt_cost, completion_cost = litellm.cost_per_token( - model="azure/gpt-4o", - prompt_tokens=prompt_tokens, - completion_tokens=completion_tokens, - ) - response_cost = prompt_cost + completion_cost - await asyncio.sleep(8) # allow db log to be updated - print(f"new_key: {new_key}") - key_info = await get_key_info( - session=session, get_key=new_key, call_key=new_key - ) - print( - f"response_cost: {response_cost}; key_info spend: {key_info['info']['spend']}" - ) - rounded_response_cost = round(response_cost, 8) - rounded_key_info_spend = round(key_info["info"]["spend"], 8) - assert ( - rounded_response_cost == rounded_key_info_spend - ), f"Expected={rounded_response_cost}, Got={rounded_key_info_spend}" - - @pytest.mark.skip(reason="Frequent check on ci/cd leads to read timeout issue.") @pytest.mark.asyncio async def test_key_with_budgets(): @@ -455,88 +159,3 @@ async def test_key_with_budgets(): # assert rounded_response_cost == rounded_key_info_spend - - -@pytest.mark.asyncio -async def test_key_delete_ui(): - """ - Admin UI flow - DO NOT DELETE - -> Create a key with user_id = "ishaan" - -> Log on Admin UI, delete the key for user "ishaan" - -> This should work, since we're on the admin UI and role == "proxy_admin - """ - async with aiohttp.ClientSession() as session: - key_gen = await generate_key(session=session, i=0, user_id="ishaan-smart") - key = key_gen["key"] - - # generate a admin UI key - team = await generate_team(session=session) - admin_ui_key = await generate_user( - session=session, user_role=LitellmUserRoles.PROXY_ADMIN.value - ) - print( - "trying to delete key=", - key, - "using key=", - admin_ui_key["key"], - " to auth in", - ) - - await delete_key( - session=session, - get_key=key, - auth_key=admin_ui_key["key"], - ) - - -@pytest.mark.parametrize("model_access", ["all-team-models", "gpt-3.5-turbo"]) -@pytest.mark.parametrize("model_access_level", ["key", "team"]) -@pytest.mark.parametrize("model_endpoint", ["/v1/models", "/model/info"]) -@pytest.mark.asyncio -async def test_key_model_list(model_access, model_access_level, model_endpoint): - """ - Test if `/v1/models` works as expected. - """ - async with aiohttp.ClientSession() as session: - _models = [] if model_access == "all-team-models" else [model_access] - team_id = "litellm_dashboard_{}".format(uuid.uuid4()) - new_team = await generate_team( - session=session, - models=_models if model_access_level == "team" else None, - team_id=team_id, - ) - assert new_team["team_id"] == team_id - key_gen = await generate_key( - session=session, - i=0, - team_id=team_id, - models=_models if model_access_level == "key" else [], - ) - key = key_gen["key"] - print(f"key: {key}") - - model_list = await get_model_list( - session=session, call_key=key, endpoint=model_endpoint - ) - print(f"model_list: {model_list}") - - if model_access == "all-team-models": - if model_endpoint == "/v1/models": - assert not isinstance(model_list["data"][0]["id"], list) - assert isinstance(model_list["data"][0]["id"], str) - elif model_endpoint == "/model/info": - assert isinstance(model_list["data"], list) - assert len(model_list["data"]) > 0 - if model_access == "gpt-3.5-turbo": - if model_endpoint == "/v1/models": - assert {entry["id"] for entry in model_list["data"]} == { - model_access, - "mistral-7b", - }, "generate_key sets alias mistral-7b -> gpt-3.5-turbo, so /v1/models lists both; model_access={}, model_access_level={}".format( - model_access, model_access_level - ) - elif model_endpoint == "/model/info": - assert isinstance(model_list["data"], list) - assert len(model_list["data"]) == 1 - - diff --git a/tests/test_litellm_proxy_responses_config.py b/tests/test_litellm_proxy_responses_config.py index 0743565874a..7bfb89a8711 100644 --- a/tests/test_litellm_proxy_responses_config.py +++ b/tests/test_litellm_proxy_responses_config.py @@ -2,61 +2,7 @@ Unit test for LiteLLM Proxy Responses API configuration. """ -import pytest - from litellm.types.utils import LlmProviders -from litellm.utils import ProviderConfigManager - - -def test_litellm_proxy_responses_api_config(): - """Test that litellm_proxy provider returns correct Responses API config""" - from litellm.llms.litellm_proxy.responses.transformation import ( - LiteLLMProxyResponsesAPIConfig, - ) - - config = ProviderConfigManager.get_provider_responses_api_config( - model="litellm_proxy/gpt-5.5", - provider=LlmProviders.LITELLM_PROXY, - ) - print(f"config: {config}") - assert config is not None, "Config should not be None for litellm_proxy provider" - assert isinstance( - config, LiteLLMProxyResponsesAPIConfig - ), f"Expected LiteLLMProxyResponsesAPIConfig, got {type(config)}" - assert ( - config.custom_llm_provider == LlmProviders.LITELLM_PROXY - ), "custom_llm_provider should be LITELLM_PROXY" - - -def test_litellm_proxy_responses_api_config_get_complete_url(): - """Test that get_complete_url works correctly""" - import os - from litellm.llms.litellm_proxy.responses.transformation import ( - LiteLLMProxyResponsesAPIConfig, - ) - - config = LiteLLMProxyResponsesAPIConfig() - - # Test with explicit api_base - url = config.get_complete_url( - api_base="https://my-proxy.example.com", - litellm_params={}, - ) - assert url == "https://my-proxy.example.com/responses" - - # Test with trailing slash - url = config.get_complete_url( - api_base="https://my-proxy.example.com/", - litellm_params={}, - ) - assert url == "https://my-proxy.example.com/responses" - - # Test that it raises error when api_base is None and env var is not set - if "LITELLM_PROXY_API_BASE" in os.environ: - del os.environ["LITELLM_PROXY_API_BASE"] - - with pytest.raises(ValueError, match="api_base not set"): - config.get_complete_url(api_base=None, litellm_params={}) def test_litellm_proxy_responses_api_config_inherits_from_openai(): @@ -70,15 +16,6 @@ def test_litellm_proxy_responses_api_config_inherits_from_openai(): config = LiteLLMProxyResponsesAPIConfig() - # Should inherit from OpenAI config assert isinstance(config, OpenAIResponsesAPIConfig) - # Should have the correct provider set assert config.custom_llm_provider == LlmProviders.LITELLM_PROXY - - -if __name__ == "__main__": - test_litellm_proxy_responses_api_config() - test_litellm_proxy_responses_api_config_get_complete_url() - test_litellm_proxy_responses_api_config_inherits_from_openai() - print("All tests passed!") diff --git a/tests/test_litellm_rust/cache/test_python_cache.py b/tests/test_litellm_rust/cache/test_python_cache.py index 490643191ea..139859bdf55 100644 --- a/tests/test_litellm_rust/cache/test_python_cache.py +++ b/tests/test_litellm_rust/cache/test_python_cache.py @@ -446,7 +446,7 @@ def test_sync_rust_messages_calls_python_cache(recording_server: RecordingServer request: Final = NativeCall( args=(), kwargs=arguments, - bound={ + base={ "model": MESSAGES_MODEL, "messages": list(MESSAGES), "max_tokens": 32, diff --git a/tests/test_litellm_rust/messages/test_request_shaping.py b/tests/test_litellm_rust/messages/test_request_shaping.py index a1c99f245b4..3584a2a8593 100644 --- a/tests/test_litellm_rust/messages/test_request_shaping.py +++ b/tests/test_litellm_rust/messages/test_request_shaping.py @@ -251,7 +251,7 @@ async def test_native_messages_observes_runtime_capabilities_and_separate_caller request: Final = NativeCall( args=(), kwargs={"temperature": 0.2, "drop_params": True}, - bound={ + base={ "model": model, "messages": MESSAGES, "max_tokens": 16, diff --git a/tests/test_litellm_rust/ocr/test_requests.py b/tests/test_litellm_rust/ocr/test_requests.py index 1eae999ff70..8767b273b68 100644 --- a/tests/test_litellm_rust/ocr/test_requests.py +++ b/tests/test_litellm_rust/ocr/test_requests.py @@ -624,7 +624,7 @@ def test_native_projection_errors_never_select_python( request: Final = NativeCall( args=(), kwargs={}, - bound={ + base={ "model": "mistral/mistral-ocr-latest", "document": OCR_DOCUMENT, "api_key": "test-key", diff --git a/tests/test_litellm_rust/support/cache.py b/tests/test_litellm_rust/support/cache.py index 352722f3269..5deacc7869f 100644 --- a/tests/test_litellm_rust/support/cache.py +++ b/tests/test_litellm_rust/support/cache.py @@ -51,7 +51,7 @@ async def invoke( request: Final = NativeCall( args=(), kwargs=arguments, - bound={ + base={ "model": RESPONSES_MODEL, "input": "hello", "stream": None, @@ -78,7 +78,7 @@ async def invoke( if route == "chat": if not native: return await litellm.acompletion(**parameters) - chat: Final = NativeCall(args=(), kwargs=parameters, bound=parameters) + chat: Final = NativeCall(args=(), kwargs=parameters, base={}) return await runtime.arun( RouteContext(Route.CHAT_COMPLETIONS), binding=NATIVE_ACOMPLETION, @@ -91,7 +91,7 @@ async def invoke( messages: Final = NativeCall( args=(), kwargs=parameters, - bound={ + base={ "model": MESSAGES_MODEL, "messages": list(MESSAGES), "max_tokens": 32, diff --git a/tests/test_litellm_rust/test_inference.py b/tests/test_litellm_rust/test_inference.py index 5e4b7b284bb..8211b50cb8d 100644 --- a/tests/test_litellm_rust/test_inference.py +++ b/tests/test_litellm_rust/test_inference.py @@ -66,7 +66,7 @@ def native_call( "max_tokens": 32, **options, } - request: Final = NativeCall(args=(), kwargs=kwargs, bound=kwargs) + request: Final = NativeCall(args=(), kwargs=kwargs, base={}) return (_native.acompletion if asynchronous else _native.completion)(request) response_kwargs: Final = { "model": RESPONSES_MODEL, @@ -79,7 +79,7 @@ def native_call( response_request: Final = NativeCall( args=(), kwargs=response_kwargs, - bound={ + base={ "model": RESPONSES_MODEL, "input": "hello", "stream": None, diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index f07bc776bce..3596d1d343c 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -174,10 +174,11 @@ async def test_schema_setup_uses_configured_retention(recording_server: Recordin request.raw_body for request in recording_server.requests if b"MODIFY TTL" in request.raw_body ) assert all(b"INTERVAL 7 DAY" in statement for statement in ttl_statements) - assert tuple(request.raw_body.strip() for request in recording_server.requests[-3:]) == ( + assert tuple(request.raw_body.strip() for request in recording_server.requests[-4:]) == ( b"ALTER TABLE `trace_test`.otel_traces MODIFY TTL toDateTime(Timestamp) + INTERVAL 7 DAY", b"ALTER TABLE `trace_test`.agent_traces_by_key MODIFY TTL toDateTime(StartTs) + INTERVAL 7 DAY", b"ALTER TABLE `trace_test`.spend_logs MODIFY TTL toDateTime(start_time) + INTERVAL 7 DAY", + b"ALTER TABLE `trace_test`.lens_feedback MODIFY TTL toDateTime(CreatedAt) + INTERVAL 7 DAY", ) diff --git a/tests/test_models.py b/tests/test_models.py index 14706098f27..8ae9d8edcd7 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -5,8 +5,6 @@ import pytest import asyncio import aiohttp import os -import dotenv -from typing import Final from dotenv import load_dotenv load_dotenv() @@ -32,46 +30,6 @@ async def generate_key(session, models=[]): return await response.json() -async def get_models(session, key, only_model_access_groups=False): - url = "http://0.0.0.0:4000/models" - if only_model_access_groups: - url += "?only_model_access_groups=True" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - print("response from /models") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - -@pytest.mark.asyncio -async def test_get_models_multiple_tests(): - async with aiohttp.ClientSession() as session: - key_gen = await generate_key(session=session) - key = key_gen["key"] - models = await get_models(session=session, key=key) - print(f"\n\nmodels: {models}") - assert len(models["data"]) > 0 - - ## Test only_model_access_groups - new_response = await get_models( - session=session, key=key, only_model_access_groups=True - ) - print(f"\n\nnew_response: {new_response}") - assert ( - len(new_response["data"]) == 0 - ) # no model access groups set on config.yaml - - async def add_models( session, model_id="123", model_name="azure-gpt-3.5", key=os.environ["LITELLM_MASTER_KEY"], team_id=None ): @@ -106,48 +64,6 @@ async def add_models( return response_json -async def get_model_info(session, key, litellm_model_id=None): - """ - Make sure only models user has access to are returned - """ - if litellm_model_id: - url = f"http://0.0.0.0:4000/model/info?litellm_model_id={litellm_model_id}" - else: - url = "http://0.0.0.0:4000/model/info" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - -async def get_model_group_info(session, key): - url = "http://0.0.0.0:4000/model_group/info" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - async def chat_completion(session, key, model="azure-gpt-3.5"): url = "http://0.0.0.0:4000/chat/completions" headers = { @@ -173,49 +89,6 @@ async def chat_completion(session, key, model="azure-gpt-3.5"): raise Exception(f"Request did not return a 200 status code: {status}") -@pytest.mark.asyncio -async def test_get_models(): - """ - Get models user has access to - """ - async with aiohttp.ClientSession() as session: - key_gen = await generate_key(session=session, models=["gpt-4"]) - key = key_gen["key"] - response = await get_model_info(session=session, key=key) - models = [m["model_name"] for m in response["data"]] - for m in models: - assert m == "gpt-4" - - -@pytest.mark.asyncio -async def test_get_specific_model(): - """ - Return specific model info - - Ensure value of model_info is same as on `/model/info` (no id set) - """ - async with aiohttp.ClientSession() as session: - key_gen = await generate_key(session=session, models=["gpt-4"]) - key = key_gen["key"] - response = await get_model_info(session=session, key=key) - models = [m["model_name"] for m in response["data"]] - model_specific_info = None - for idx, m in enumerate(models): - assert m == "gpt-4" - litellm_model_id = response["data"][idx]["model_info"]["id"] - model_specific_info = response["data"][idx] - assert litellm_model_id is not None - response = await get_model_info( - session=session, key=key, litellm_model_id=litellm_model_id - ) - assert response["data"][0]["model_info"]["id"] == litellm_model_id - assert ( - response["data"][0] == model_specific_info - ), "Model info is not the same. Got={}, Expected={}".format( - response["data"][0], model_specific_info - ) - - async def delete_model(session, model_id="123", key=os.environ["LITELLM_MASTER_KEY"]): """ Make sure only models user has access to are returned @@ -270,45 +143,3 @@ async def test_add_and_delete_models(): pass -@pytest.mark.asyncio -async def test_get_personal_models_for_user(): - """ - Test /models endpoint with team - """ - from tests.test_users import new_user - - async with aiohttp.ClientSession() as session: - # Creat a user - user_data = await new_user(session=session, i=0, models=["gpt-3.5-turbo"]) - user_id = user_data["user_id"] - user_api_key = user_data["key"] - - model_group_info = await get_model_group_info(session=session, key=user_api_key) - print(model_group_info) - - assert len(model_group_info["data"]) == 1 - assert model_group_info["data"][0]["model_group"] == "gpt-3.5-turbo" - - -@pytest.mark.asyncio -async def test_model_group_info_e2e(): - """ - Test /model/group/info endpoint - """ - async with aiohttp.ClientSession() as session: - models = await get_models(session=session, key=os.environ["LITELLM_MASTER_KEY"]) - print(models) - - model_group_info = await get_model_group_info(session=session, key=os.environ["LITELLM_MASTER_KEY"]) - print(model_group_info) - - model_groups: Final = [m["model_group"] for m in model_group_info["data"]] - - assert "anthropic/*" not in model_groups, ( - f"Expected 'anthropic/*' to be expanded, but it was returned verbatim: {model_groups}" - ) - assert any(m.startswith("anthropic/") for m in model_groups), ( - f"Expected concrete anthropic models from the 'anthropic/*' config entry, got: {model_groups}" - ) - - diff --git a/tests/test_new_vector_store_endpoints.py b/tests/test_new_vector_store_endpoints.py index c44723937ac..8cd866a8e3c 100644 --- a/tests/test_new_vector_store_endpoints.py +++ b/tests/test_new_vector_store_endpoints.py @@ -4,189 +4,12 @@ Tests both basic functionality and complex scenarios including target_model_name """ import asyncio -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, patch import pytest import litellm -from litellm.proxy._types import UserAPIKeyAuth - - -@pytest.mark.asyncio -async def test_vector_store_retrieve_basic(): - """Test basic vector store retrieve functionality.""" - mock_response = { - "id": "vs_test123", - "object": "vector_store", - "created_at": 1699061776, - "name": "Test Vector Store", - "file_counts": { - "in_progress": 0, - "completed": 5, - "failed": 0, - "cancelled": 0, - "total": 5, - }, - "status": "completed", - "usage_bytes": 12345, - } - - with patch( - "litellm.vector_stores.main.aretrieve", - new=AsyncMock(return_value=mock_response), - ) as mock_retrieve: - router = litellm.Router(model_list=[]) - result = await router.avector_store_retrieve( - vector_store_id="vs_test123", - custom_llm_provider="openai", - ) - - assert result["id"] == "vs_test123" - assert result["object"] == "vector_store" - assert result["status"] == "completed" - mock_retrieve.assert_called_once() - - -@pytest.mark.asyncio -async def test_vector_store_list_basic(): - """Test basic vector store list functionality.""" - mock_response = { - "object": "list", - "data": [ - { - "id": "vs_test1", - "object": "vector_store", - "created_at": 1699061776, - "name": "Store 1", - }, - { - "id": "vs_test2", - "object": "vector_store", - "created_at": 1699061777, - "name": "Store 2", - }, - ], - "first_id": "vs_test1", - "last_id": "vs_test2", - "has_more": False, - } - - with patch( - "litellm.vector_stores.main.alist", - new=AsyncMock(return_value=mock_response), - ) as mock_list: - router = litellm.Router(model_list=[]) - result = await router.avector_store_list( - limit=20, - order="desc", - custom_llm_provider="openai", - ) - - assert result["object"] == "list" - assert len(result["data"]) == 2 - assert result["data"][0]["id"] == "vs_test1" - mock_list.assert_called_once() - - -@pytest.mark.asyncio -async def test_vector_store_update_basic(): - """Test basic vector store update functionality.""" - mock_response = { - "id": "vs_test123", - "object": "vector_store", - "created_at": 1699061776, - "name": "Updated Name", - "metadata": {"key": "value"}, - "status": "completed", - } - - with patch( - "litellm.vector_stores.main.aupdate", - new=AsyncMock(return_value=mock_response), - ) as mock_update: - router = litellm.Router(model_list=[]) - result = await router.avector_store_update( - vector_store_id="vs_test123", - name="Updated Name", - metadata={"key": "value"}, - custom_llm_provider="openai", - ) - - assert result["id"] == "vs_test123" - assert result["name"] == "Updated Name" - assert result["metadata"]["key"] == "value" - mock_update.assert_called_once() - - -@pytest.mark.asyncio -async def test_vector_store_delete_basic(): - """Test basic vector store delete functionality.""" - mock_response = { - "id": "vs_test123", - "object": "vector_store.deleted", - "deleted": True, - } - - with patch( - "litellm.vector_stores.main.adelete", - new=AsyncMock(return_value=mock_response), - ) as mock_delete: - router = litellm.Router(model_list=[]) - result = await router.avector_store_delete( - vector_store_id="vs_test123", - custom_llm_provider="openai", - ) - - assert result["id"] == "vs_test123" - assert result["deleted"] is True - assert result["object"] == "vector_store.deleted" - mock_delete.assert_called_once() - - -@pytest.mark.asyncio -async def test_async_vector_store_retrieve(): - """Test async vector store retrieve.""" - mock_response = { - "id": "vs_async123", - "object": "vector_store", - "name": "Async Test Store", - } - - with patch( - "litellm.vector_stores.main.aretrieve", - new=AsyncMock(return_value=mock_response), - ) as mock_aretrieve: - router = litellm.Router(model_list=[]) - result = await router.avector_store_retrieve( - vector_store_id="vs_async123", - custom_llm_provider="openai", - ) - - assert result["id"] == "vs_async123" - mock_aretrieve.assert_called_once() - - -@pytest.mark.asyncio -async def test_async_vector_store_list(): - """Test async vector store list.""" - mock_response = { - "object": "list", - "data": [{"id": "vs_1"}, {"id": "vs_2"}], - } - - with patch( - "litellm.vector_stores.main.alist", - new=AsyncMock(return_value=mock_response), - ) as mock_alist: - router = litellm.Router(model_list=[]) - result = await router.avector_store_list( - limit=10, - custom_llm_provider="openai", - ) - - assert len(result["data"]) == 2 - mock_alist.assert_called_once() @pytest.mark.asyncio @@ -212,93 +35,6 @@ async def test_async_vector_store_update(): mock_aupdate.assert_called_once() -@pytest.mark.asyncio -async def test_async_vector_store_delete(): - """Test async vector store delete.""" - mock_response = { - "id": "vs_async123", - "deleted": True, - } - - with patch( - "litellm.vector_stores.main.adelete", - new=AsyncMock(return_value=mock_response), - ) as mock_adelete: - router = litellm.Router(model_list=[]) - result = await router.avector_store_delete( - vector_store_id="vs_async123", - custom_llm_provider="openai", - ) - - assert result["deleted"] is True - mock_adelete.assert_called_once() - - -@pytest.mark.asyncio -async def test_vector_store_list_with_pagination(): - """Test vector store list with pagination parameters.""" - mock_response = { - "object": "list", - "data": [{"id": f"vs_{i}"} for i in range(5)], - "has_more": True, - "first_id": "vs_0", - "last_id": "vs_4", - } - - with patch( - "litellm.vector_stores.main.list", - return_value=mock_response, - ) as mock_list: - router = litellm.Router(model_list=[]) - result = router.vector_store_list( - limit=5, - after="vs_previous", - order="asc", - custom_llm_provider="openai", - ) - - assert result["has_more"] is True - assert len(result["data"]) == 5 - - # Verify pagination params were passed - call_kwargs = mock_list.call_args.kwargs - assert call_kwargs["limit"] == 5 - assert call_kwargs["after"] == "vs_previous" - assert call_kwargs["order"] == "asc" - - -@pytest.mark.asyncio -async def test_vector_store_update_with_expires_after(): - """Test vector store update with expiration policy.""" - expires_after = { - "anchor": "last_active_at", - "days": 7, - } - - mock_response = { - "id": "vs_test123", - "expires_after": expires_after, - "expires_at": 1699668576, - } - - with patch( - "litellm.vector_stores.main.update", - return_value=mock_response, - ) as mock_update: - router = litellm.Router(model_list=[]) - result = router.vector_store_update( - vector_store_id="vs_test123", - expires_after=expires_after, - custom_llm_provider="openai", - ) - - assert result["expires_after"]["days"] == 7 - assert result["expires_at"] is not None - - call_kwargs = mock_update.call_args.kwargs - assert call_kwargs["expires_after"] == expires_after - - def test_router_initializes_new_endpoints(): """Test that router properly initializes the new vector store endpoints.""" router = litellm.Router(model_list=[]) @@ -333,20 +69,9 @@ if __name__ == "__main__": test_router_initializes_new_endpoints() print("✓ Router initialization successful") - # Test basic sync operations - print("✓ Testing basic sync operations...") - asyncio.run(test_vector_store_retrieve_basic()) - asyncio.run(test_vector_store_list_basic()) - asyncio.run(test_vector_store_update_basic()) - asyncio.run(test_vector_store_delete_basic()) - print("✓ Basic sync operations successful") - # Test async operations print("✓ Testing async operations...") - asyncio.run(test_async_vector_store_retrieve()) - asyncio.run(test_async_vector_store_list()) asyncio.run(test_async_vector_store_update()) - asyncio.run(test_async_vector_store_delete()) print("✓ Async operations successful") print("\n✅ All smoke tests passed!") diff --git a/tests/test_openai_endpoints.py b/tests/test_openai_endpoints.py index 6b92373a98c..691ed37ac9f 100644 --- a/tests/test_openai_endpoints.py +++ b/tests/test_openai_endpoints.py @@ -3,10 +3,7 @@ from typing import Final # What this tests ? ## Tests /chat/completions by generating a key and then making a chat completions-request import pytest -import asyncio -import aiohttp, openai -from openai import OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI -from typing import Optional, List, Union +from openai import AsyncOpenAI LITELLM_MASTER_KEY = os.environ["LITELLM_MASTER_KEY"] @@ -19,64 +16,6 @@ def response_header_check(response): assert headers_size < 4096, "Response headers exceed the 4kb limit" -async def generate_key( - session, - models=[ - "gpt-4", - "text-embedding-ada-002", - "gpt-image-1", - "fake-openai-endpoint-2", - "mistral-embed", - ], -): - url = "http://0.0.0.0:4000/key/generate" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = { - "models": models, - "duration": None, - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - response_header_check( - response - ) # calling the function to check response headers - - return await response.json() - - -async def new_user(session): - url = "http://0.0.0.0:4000/user/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = { - "models": ["gpt-4", "text-embedding-ada-002", "gpt-image-1"], - "duration": None, - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - response_header_check( - response - ) # calling the function to check response headers - return await response.json() - - async def moderation(session, key): url = "http://0.0.0.0:4000/moderations" headers = { @@ -98,141 +37,6 @@ async def moderation(session, key): return await response.json() -async def chat_completion(session, key, model: Union[str, List] = "gpt-4"): - url = "http://0.0.0.0:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": model, - "messages": [ - {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": "Hello!"}, - ], - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception( - f"Request did not return a 200 status code: {status}, response text={response_text}" - ) - - response_header_check( - response - ) # calling the function to check response headers - - return await response.json() - - -async def queue_chat_completion( - session, key, priority: int, model: Union[str, List] = "gpt-4" -): - url = "http://0.0.0.0:4000/queue/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": model, - "messages": [ - {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": "Hello!"}, - ], - "priority": priority, - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - return response.raw_headers - - -async def chat_completion_with_headers(session, key, model="gpt-4"): - url = "http://0.0.0.0:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": model, - "messages": [ - {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": "Hello!"}, - ], - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - response_header_check( - response - ) # calling the function to check response headers - - raw_headers = response.raw_headers - raw_headers_json = {} - - for ( - item - ) in ( - response.raw_headers - ): # ((b'date', b'Fri, 19 Apr 2024 21:17:29 GMT'), (), ) - raw_headers_json[item[0].decode("utf-8")] = item[1].decode("utf-8") - - return raw_headers_json - - -async def chat_completion_with_model_from_route(session, key, route): - url = "http://0.0.0.0:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - - -async def completion(session, key): - url = "http://0.0.0.0:4000/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = {"model": "gpt-4", "prompt": "Hello!"} - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - response_header_check( - response - ) # calling the function to check response headers - - response = await response.json() - - return response - - async def embeddings(session, key, model="text-embedding-ada-002"): url = "http://0.0.0.0:4000/embeddings" headers = { @@ -258,120 +62,6 @@ async def embeddings(session, key, model="text-embedding-ada-002"): ) # calling the function to check response headers -async def image_generation(session, key): - url = "http://0.0.0.0:4000/images/generations" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": "gpt-image-1", - "prompt": "A cute baby sea otter", - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - if ( - "Connection error" in response_text - ): # OpenAI endpoint returns a connection error - return - raise Exception(f"Request did not return a 200 status code: {status}") - - response_header_check( - response - ) # calling the function to check response headers - - -@pytest.mark.asyncio -async def test_chat_completion(): - """ - - Create key - Make chat completion call - - Create user - make chat completion call - """ - async with aiohttp.ClientSession() as session: - key_gen = await generate_key(session=session, models=["gpt-3.5-turbo"]) - azure_client = AsyncAzureOpenAI( - azure_endpoint="http://0.0.0.0:4000", - azure_deployment="random-model", - api_key=key_gen["key"], - api_version="2024-02-15-preview", - ) - with pytest.raises(openai.PermissionDeniedError) as e: - response = await azure_client.chat.completions.create( - model="gpt-4", - messages=[{"role": "user", "content": "Hello!"}], - ) - assert "is not available for this API key" in str(e.value) - - -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=1) -@pytest.mark.skip(reason="Flaky test, this works locally but not on CI") -async def test_chat_completion_ratelimit(): - """ - - call model with rpm 1 - - make 2 parallel calls - - make sure 1 fails - """ - async with aiohttp.ClientSession() as session: - # key_gen = await generate_key(session=session) - key = os.environ["LITELLM_MASTER_KEY"] - tasks = [] - tasks.append( - chat_completion(session=session, key=key, model="fake-openai-endpoint-2") - ) - tasks.append( - chat_completion(session=session, key=key, model="fake-openai-endpoint-2") - ) - try: - await asyncio.gather(*tasks) - pytest.fail("Expected at least 1 call to fail") - except Exception as e: - if "Request did not return a 200 status code: 429" in str(e): - pass - else: - pytest.fail(f"Wrong error received - {str(e)}") - - -@pytest.mark.asyncio -@pytest.mark.skip(reason="Flaky test") -async def test_chat_completion_different_deployments(): - """ - - call model group with 2 deployments - - make 5 calls - - expect 2 unique deployments - """ - async with aiohttp.ClientSession() as session: - # key_gen = await generate_key(session=session) - key = os.environ["LITELLM_MASTER_KEY"] - results = [] - for _ in range(20): - results.append( - await chat_completion_with_headers( - session=session, key=key, model="fake-openai-endpoint-3" - ) - ) - try: - print(f"results: {results}") - init_model_id = results[0]["x-litellm-model-id"] - deployments_shuffled = False - for result in results[1:]: - if init_model_id != result["x-litellm-model-id"]: - deployments_shuffled = True - if deployments_shuffled == False: - pytest.fail("Expected at least 1 shuffled call") - except Exception as e: - pass - - @pytest.mark.asyncio async def test_chat_completion_streaming(): """ @@ -425,46 +115,3 @@ async def test_completion_streaming_usage_metrics(): assert last_chunk.usage.total_tokens > 0, "Total tokens should be greater than 0" -@pytest.mark.asyncio -async def test_proxy_all_models(): - """ - - proxy_server_config.yaml has model = * / * - - Make chat completion call - - groq is NOT defined on /models - - - """ - async with aiohttp.ClientSession() as session: - # call chat/completions with a model that the key was not created for + the model is not on the config.yaml - await chat_completion( - session=session, key=LITELLM_MASTER_KEY, model="groq/openai/gpt-oss-120b" - ) - - await chat_completion( - session=session, - key=LITELLM_MASTER_KEY, - model="anthropic/claude-sonnet-4-5-20250929", - ) - - -@pytest.mark.asyncio -async def test_batch_chat_completions(): - """ - - Make chat completion call using - - """ - async with aiohttp.ClientSession() as session: - - # call chat/completions with a model that the key was not created for + the model is not on the config.yaml - response = await chat_completion( - session=session, - key=os.environ["LITELLM_MASTER_KEY"], - model="gpt-3.5-turbo,fake-openai-endpoint", - ) - - print(f"response: {response}") - - assert len(response) == 2 - assert isinstance(response, list) - - diff --git a/tests/test_otel_thread_leak.py b/tests/test_otel_thread_leak.py deleted file mode 100644 index cb5f54eefa4..00000000000 --- a/tests/test_otel_thread_leak.py +++ /dev/null @@ -1,91 +0,0 @@ -import sys -import os -import threading -import time -import pytest - -# Add the project root to the path -sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))) - -from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig -from litellm.types.utils import StandardCallbackDynamicParams - - -def get_thread_count() -> int: - """Helper to get active thread count""" - return threading.active_count() - - -@pytest.fixture -def otel_logger(): - """Fixture to provide a clean OTEL logger for each test""" - config = OpenTelemetryConfig( - exporter="console", enable_metrics=False, service_name="litellm-unit-test" - ) - return OpenTelemetry(config=config) - - -def test_otel_thread_leak_dynamic_headers(otel_logger): - """ - Unit test to verify that calling get_tracer_to_use_for_request with - dynamic headers doesn't cause a linear thread leak. - - This test reproduces the issue where each unique team/key credential - set causes a new TracerProvider (and its background threads) to be - spawned but never closed. - """ - - # 1. Setup dynamic header simulation (monkey-patch) - # This simulates what LangfuseOtelLogger does for per-team keys - def mock_construct_dynamic_headers(standard_callback_dynamic_params): - if standard_callback_dynamic_params: - return {"Authorization": "Bearer fake_token"} - return None - - otel_logger.construct_dynamic_otel_headers = mock_construct_dynamic_headers - - # 2. Establish Baseline - initial_threads = get_thread_count() - - # 3. Simulate requests - num_requests = 10 - latencies = [] - - print("\n🚀 Simulating requests with dynamic headers:") - for i in range(num_requests): - kwargs = { - "standard_callback_dynamic_params": StandardCallbackDynamicParams( - langfuse_public_key=f"key_{i}", - langfuse_secret_key=f"secret_{i}", - ) - } - - # Measure latency - start_time = time.perf_counter() - tracer = otel_logger.get_tracer_to_use_for_request(kwargs) - end_time = time.perf_counter() - - latency_ms = (end_time - start_time) * 1000 - latencies.append(latency_ms) - print(f" Request {i+1:2d}: Latency = {latency_ms:6.2f} ms") - - # Verify a tracer was actually returned - assert tracer is not None - - avg_latency = sum(latencies) / len(latencies) - print(f"\n📊 Average Latency: {avg_latency:.2f} ms") - - # 4. Check for leaks - # Allow for a small constant increase (OTEL might start a few shared threads) - # but a linear leak would result in +10 or more threads here. - final_threads = get_thread_count() - thread_delta = final_threads - initial_threads - - print(f"\nThread growth: {thread_delta} threads across {num_requests} requests") - - # ASSERTION: The growth should be significantly less than 1 thread per request. - # If the bug exists, thread_delta will be >= num_requests. - assert thread_delta < (num_requests / 2), ( - f"Thread leak detected! Threads grew by {thread_delta} over {num_requests} requests. " - "Each request with dynamic headers appears to be leaking background threads." - ) diff --git a/tests/test_presidio_latency.py b/tests/test_presidio_latency.py deleted file mode 100644 index bb676ca051e..00000000000 --- a/tests/test_presidio_latency.py +++ /dev/null @@ -1,77 +0,0 @@ -import asyncio -import aiohttp -import pytest -from unittest.mock import MagicMock, patch -from litellm.proxy.guardrails.guardrail_hooks.presidio import ( - OPTIONAL_PresidioPIIMasking, -) - - -@pytest.mark.asyncio -async def test_sanity_presidio_session_reuse_main_thread(): - """ - SANITY CHECK: - Verify that Presidio guardrail reuses sessions in the main thread. - This ensures we don't break existing session pooling functionality. - """ - presidio = OPTIONAL_PresidioPIIMasking( - mock_testing=True, - presidio_analyzer_api_base="http://mock-analyzer", - presidio_anonymizer_api_base="http://mock-anonymizer", - ) - - session_creations = 0 - original_init = aiohttp.ClientSession.__init__ - - def mocked_init(self, *args, **kwargs): - nonlocal session_creations - session_creations += 1 - original_init(self, *args, **kwargs) - - with patch.object(aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True): - for _ in range(10): - async with presidio._get_session_iterator() as session: - pass - - # Expected: Only 1 session created for all 10 calls. - assert session_creations == 1 - - await presidio._close_http_session() - - -@pytest.mark.asyncio -async def test_bug_presidio_session_explosion_background_thread_causes_latency(): - """ - BUG REPRODUCTION: - Verify that background threads (like logging hooks) REUSE sessions. - Previously, each call in a background loop created a NEW ephemeral session, - leading to socket exhaustion and the reported 97s latency spike. - """ - import threading - - presidio = OPTIONAL_PresidioPIIMasking( - mock_testing=True, - presidio_analyzer_api_base="http://mock-analyzer", - presidio_anonymizer_api_base="http://mock-anonymizer", - ) - - # Force the code to think it's in a background thread - presidio._main_thread_id = threading.get_ident() + 1 - - session_creations = 0 - original_init = aiohttp.ClientSession.__init__ - - def mocked_init(self, *args, **kwargs): - nonlocal session_creations - session_creations += 1 - original_init(self, *args, **kwargs) - - with patch.object(aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True): - for _ in range(10): - async with presidio._get_session_iterator() as session: - pass - - # FIX VERIFICATION: Should now be 1 session (reused) instead of 10. - assert session_creations == 1 - - await presidio._close_http_session() diff --git a/tests/test_ratelimit.py b/tests/test_ratelimit.py deleted file mode 100644 index 94d48f0accf..00000000000 --- a/tests/test_ratelimit.py +++ /dev/null @@ -1,170 +0,0 @@ -# %% -import asyncio -import os -import pytest -import random -from typing import Any -from dotenv import load_dotenv - -load_dotenv() - - -import litellm -from pydantic import BaseModel -from litellm import utils, Router - -COMPLETION_TOKENS = 5 -base_model_list = [ - { - "model_name": "gpt-5-mini", - "litellm_params": { - "model": "gpt-5-mini", - "api_key": os.getenv("OPENAI_API_KEY"), - "max_tokens": COMPLETION_TOKENS, - }, - } -] - - -class RouterConfig(BaseModel): - rpm: int - tpm: int - - -@pytest.fixture(scope="function") -def router_factory(): - def create_router(rpm, tpm, routing_strategy): - model_list = base_model_list.copy() - model_list[0]["rpm"] = rpm - model_list[0]["tpm"] = tpm - return Router( - model_list=model_list, - routing_strategy=routing_strategy, - enable_pre_call_checks=True, - debug_level="DEBUG", - ) - - return create_router - - -def generate_list_of_messages(num_messages): - """ - create num_messages new chat conversations - """ - return [ - [{"role": "user", "content": f"{i}. Hey, how's it going? {random.random()}"}] - for i in range(num_messages) - ] - - -def calculate_limits(list_of_messages): - """ - Return the min rpm and tpm level that would let all messages in list_of_messages be sent this minute - """ - rpm = len(list_of_messages) - tpm = sum( - (utils.token_counter(messages=m) + COMPLETION_TOKENS for m in list_of_messages) - ) - return rpm, tpm - - -async def async_call(router: Router, list_of_messages) -> Any: - tasks = [ - router.acompletion(model="gpt-5-mini", messages=m) for m in list_of_messages - ] - return await asyncio.gather(*tasks) - - -def sync_call(router: Router, list_of_messages) -> Any: - return [ - router.completion(model="gpt-5-mini", messages=m) for m in list_of_messages - ] - - -class ExpectNoException(Exception): - pass - - -@pytest.mark.parametrize( - "num_try_send, num_allowed_send", - [ - (2, 3), # sending as many as allowed, ExpectNoException - # (10, 10), # sending as many as allowed, ExpectNoException - (3, 2), # Sending more than allowed, ValueError - # (10, 9), # Sending more than allowed, ValueError - ], -) -@pytest.mark.parametrize( - "sync_mode", [True, False] -) # Use parametrization for sync/async -@pytest.mark.parametrize( - "routing_strategy", - [ - "usage-based-routing", - # "simple-shuffle", # dont expect to rate limit - # "least-busy", # dont expect to rate limit - # "latency-based-routing", - ], -) -def test_async_rate_limit( - router_factory, num_try_send, num_allowed_send, sync_mode, routing_strategy -): - """ - Check if router.completion and router.acompletion can send more messages than they've been limited to. - Args: - router_factory: makes new router object, without any shared Global state - num_try_send (int): number of messages to try to send - num_allowed_send (int): max number of messages allowed to be sent in 1 minute - sync_mode (bool): if making sync (router.completion) or async (router.acompletion) - Raises: - ValueError: Error router throws when it hits rate limits - ExpectNoException: Signfies that no other error has happened. A NOP - """ - # Can send more messages then we're going to; so don't expect a rate limit error - litellm.logging_callback_manager._reset_all_callbacks() - args = locals() - print(f"args: {args}") - expected_exception = ( - ExpectNoException if num_try_send <= num_allowed_send else ValueError - ) - - # usage-based-routing tracks RPM in log_success_event which runs in a - # background ThreadPoolExecutor. The cache update races with the next - # call's routing check, so over-limit detection is non-deterministic in - # both sync tight-loops and async concurrent gathers. - if num_try_send > num_allowed_send: - pytest.skip( - "RPM tracking via background thread is racy; " - "RPM over-limit rejection is tested for usage-based-routing-v2 in " - "tests/unit/router_strategy/test_router_routing_groups.py" - ) - - list_of_messages = generate_list_of_messages(max(num_try_send, num_allowed_send)) - rpm, tpm = calculate_limits(list_of_messages[:num_allowed_send]) - list_of_messages = list_of_messages[:num_try_send] - router: Router = router_factory(rpm, tpm, routing_strategy) - - print(f"router: {router.model_list}") - received = [] - - def _send_and_check(): - results = ( - sync_call(router, list_of_messages) - if sync_mode - else asyncio.run(async_call(router, list_of_messages)) - ) - received.extend(results) - print(results) - if len([i for i in results if i is not None]) != num_try_send: - # since not all results got returned, raise rate limit error - raise ValueError("No deployments available for selected model") - raise ExpectNoException - - with pytest.raises(expected_exception) as excinfo: # asserts correct type raised - _send_and_check() - - print(expected_exception, excinfo) - if expected_exception is ValueError: - assert "No deployments available for selected model" in str(excinfo.value) - else: - assert len([i for i in received if i is not None]) == num_try_send diff --git a/tests/test_resource_cleanup.py b/tests/test_resource_cleanup.py deleted file mode 100644 index d205b739915..00000000000 --- a/tests/test_resource_cleanup.py +++ /dev/null @@ -1,117 +0,0 @@ -""" -Test that async HTTP clients are properly cleaned up to prevent resource leaks. -Issue: https://github.com/BerriAI/litellm/issues/12107 -""" - -import asyncio -import os -import warnings - -import pytest - -import litellm - - -@pytest.mark.asyncio -async def test_acompletion_resource_cleanup(): - """Test that acompletion doesn't leave unclosed client sessions.""" - # Suppress warnings to check for them later - with warnings.catch_warnings(record=True) as w: - warnings.simplefilter("always") - - # Make an async completion call - response = await litellm.acompletion( - model="gemini/gemini-2.0-flash-lite-001", - messages=[{"role": "user", "content": "Hello"}], - mock_response="Hi there! How can I help you today?", - ) - - # Check that response was received - assert ( - response.choices[0].message.content == "Hi there! How can I help you today?" - ) - - # Manually close async clients - await litellm.close_litellm_async_clients() - - # Give a small delay for any warnings to appear - await asyncio.sleep(0.1) - - # Check for resource warnings - resource_warnings = [ - warning - for warning in w - if "Unclosed" in str(warning.message) - and ( - "client session" in str(warning.message) - or "connector" in str(warning.message) - ) - ] - - # Should be no unclosed resource warnings - assert ( - len(resource_warnings) == 0 - ), f"Found unclosed resources: {[str(w.message) for w in resource_warnings]}" - - -@pytest.mark.asyncio -async def test_multiple_acompletion_calls_cleanup(): - """Test that multiple acompletion calls reuse clients and don't leak resources.""" - with warnings.catch_warnings(record=True) as w: - warnings.simplefilter("always") - - # Make multiple async completion calls - for i in range(3): - response = await litellm.acompletion( - model="gemini/gemini-2.0-flash-lite-001", - messages=[{"role": "user", "content": f"Hello {i}"}], - mock_response=f"Response {i}", - ) - assert response.choices[0].message.content == f"Response {i}" - - # Clean up - await litellm.close_litellm_async_clients() - - # Give a small delay for any warnings to appear - await asyncio.sleep(0.1) - - # Check for resource warnings - resource_warnings = [ - warning - for warning in w - if "Unclosed" in str(warning.message) - and ( - "client session" in str(warning.message) - or "connector" in str(warning.message) - ) - ] - - assert ( - len(resource_warnings) == 0 - ), f"Found unclosed resources: {[str(w.message) for w in resource_warnings]}" - - -@pytest.mark.asyncio -async def test_cleanup_function_is_safe_to_call_multiple_times(): - """Test that the cleanup function can be called multiple times safely.""" - # This should not raise any errors - await litellm.close_litellm_async_clients() - await litellm.close_litellm_async_clients() - await litellm.close_litellm_async_clients() - - # Should still work after multiple cleanups - response = await litellm.acompletion( - model="gemini/gemini-2.0-flash-lite-001", - messages=[{"role": "user", "content": "Hello"}], - mock_response="Hi!", - ) - assert response.choices[0].message.content == "Hi!" - - # Clean up again - await litellm.close_litellm_async_clients() - - -if __name__ == "__main__": - # Run the test - asyncio.run(test_acompletion_resource_cleanup()) - print("✅ All tests passed!") diff --git a/tests/test_rust_python_harness.py b/tests/test_rust_python_harness.py index 9af38941684..e115cd4839c 100644 --- a/tests/test_rust_python_harness.py +++ b/tests/test_rust_python_harness.py @@ -1,119 +1,11 @@ from __future__ import annotations import importlib -from typing import Final -import pytest - -models = importlib.import_module("tests.rust-python-harness.shared.reporting.models") -strategy_module = importlib.import_module("tests.rust-python-harness.shared.reporting.strategy") -ui = importlib.import_module("tests.rust-python-harness.shared.reporting.ui") contracts = importlib.import_module("tests.rust-python-harness.shared.unit_runners.contracts") -cli = importlib.import_module("tests.rust-python-harness.cli") UNIT_TEST_CONTRACTS = contracts.UNIT_TEST_CONTRACTS -CaseResult = models.CaseResult -Coverage = models.Coverage -HarnessCase = models.HarnessCase -HarnessRun = models.HarnessRun -RunStatus = models.RunStatus -ModuleCaseSpec = strategy_module.ModuleCaseSpec -NotImplementedCaseSpec = strategy_module.NotImplementedCaseSpec -SkippedCaseSpec = strategy_module.SkippedCaseSpec -_format_duration = ui._format_duration -_summary = ui._summary - - -def _case(module: str = "tests.example") -> HarnessCase: - return HarnessCase( - strategy_id="example", - strategy_label="Example", - sdk_function="messages", - spec=ModuleCaseSpec(coverage=Coverage.COMPLETE, module=module), - ) - - -@pytest.mark.parametrize( - "module", - [ - "tests.rust-python-harness.strategies.trace_parity.sdk.messages.case", - "tests.rust-python-harness.strategies.trace_parity.sdk.chat_completions.case", - "tests.rust-python-harness.strategies.trace_parity.sdk.transcription.case", - ], -) -def test_implemented_namespace_case_modules_remain_importable(module: str) -> None: - assert importlib.import_module(module) - - -def test_should_mark_not_implemented_and_skipped_cases_without_running() -> None: - not_implemented: Final = CaseResult( - case=HarnessCase( - strategy_id="example", - strategy_label="Example", - sdk_function="messages", - spec=NotImplementedCaseSpec(reason="No case is registered."), - ) - ) - skipped: Final = CaseResult( - case=HarnessCase( - strategy_id="example", - strategy_label="Example", - sdk_function="messages", - spec=SkippedCaseSpec(reason="The surface does not apply."), - ) - ) - - not_implemented.set_initial_status() - skipped.set_initial_status() - - assert not_implemented.status is RunStatus.NOT_IMPLEMENTED - assert skipped.status is RunStatus.SKIPPED - - -def test_should_finalize_a_fully_passing_case() -> None: - result = CaseResult(case=_case()) - result.set_initial_status() - result.collected.update({"one", "two"}) - result.completed.update({"one", "two"}) - result.passed = 2 - - result.finalize() - - assert result.status is RunStatus.PASSED - - -def test_should_replace_a_pass_with_a_teardown_error() -> None: - result = CaseResult(case=_case()) - result.set_initial_status() - result.collected.add("one") - - result.record("one", RunStatus.PASSED, 0.1) - result.record("one", RunStatus.ERROR, 0.2) - - assert result.status is RunStatus.ERROR - assert result.passed == 0 - assert result.errors == 1 - assert result.duration == pytest.approx(0.3) - - -def test_should_format_developer_facing_run_context() -> None: - run = HarnessRun.from_cases((_case(),)) - result = next(iter(run.results.values())) - result.collected.add("tests/test_parity.py::test_one") - result.record("tests/test_parity.py::test_one", RunStatus.PASSED, 1.25) - - assert _summary(run) == (1, 0, 0, 0) - assert _format_duration(1.25) == "1.2s" def test_should_leave_functions_without_unit_test_contracts_unimplemented() -> None: assert "messages" not in UNIT_TEST_CONTRACTS - - -def test_strategy_subcommand_accepts_function_filter(capsys: pytest.CaptureFixture[str]) -> None: - exit_code: Final = cli.main(["run", "unit_tests_rust", "--function", "messages"]) - - captured: Final = capsys.readouterr() - assert exit_code == 0 - assert "- messages: not_implemented" in captured.out - assert "unit_tests_rust:messages: not_implemented" not in captured.out diff --git a/tests/test_service_logger_otel.py b/tests/test_service_logger_otel.py index 044d37d6781..8152a45f945 100644 --- a/tests/test_service_logger_otel.py +++ b/tests/test_service_logger_otel.py @@ -1,18 +1,13 @@ import os import sys import unittest -from datetime import datetime -from unittest.mock import patch, AsyncMock, MagicMock +from unittest.mock import patch # Add the project root to sys.path sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../.."))) import litellm from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger -from litellm.integrations.opentelemetry import OpenTelemetry -from litellm.types.services import ServiceTypes -from litellm._service_logger import ServiceLogging -from litellm.types.utils import StandardCallbackDynamicParams class TestServiceLoggerOTEL(unittest.IsolatedAsyncioTestCase): @@ -43,110 +38,6 @@ class TestServiceLoggerOTEL(unittest.IsolatedAsyncioTestCase): "LangfuseOtelLogger.async_service_failure_hook", ) - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_tracing") - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_metrics") - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_logs") - async def test_langfuse_otel_does_not_create_proxy_request_span( - self, mock_logs, mock_metrics, mock_tracing - ): - """ - Test that LangfuseOtelLogger returns None for create_litellm_proxy_request_started_span. - - This prevents empty proxy request spans from being sent to Langfuse when - requests don't result in actual LLM calls (e.g., auth failures, health checks). - """ - logger = LangfuseOtelLogger() - - # Verify the method is overridden - self.assertEqual( - logger.create_litellm_proxy_request_started_span.__qualname__, - "LangfuseOtelLogger.create_litellm_proxy_request_started_span", - ) - - # Verify it returns None - result = logger.create_litellm_proxy_request_started_span( - start_time=datetime.now(), - headers={"Authorization": "Bearer test"}, - ) - self.assertIsNone(result) - - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_tracing") - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_metrics") - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_logs") - async def test_service_logging_shadowing_fix( - self, mock_logs, mock_metrics, mock_tracing - ): - """ - Test the architectural fix: multiple OTEL loggers should receive logs independently. - """ - # 1. Initialize two loggers - langfuse_logger = LangfuseOtelLogger() - otel_logger = OpenTelemetry() - - # 2. Setup service_callback list - litellm.service_callback = [langfuse_logger, otel_logger] - - service_logging = ServiceLogging() - - # 3. Mock the base OpenTelemetry hook - with patch.object( - OpenTelemetry, "async_service_success_hook", new_callable=AsyncMock - ) as mock_base_hook: - # Trigger a service event - await service_logging.async_service_success_hook( - service=ServiceTypes.DB, - call_type="success", - duration=0.1, - parent_otel_span=MagicMock(), - start_time=0.0, - end_time=1.0, - ) - - # The architectural fix ensures we call each correctly. - self.assertEqual( - mock_base_hook.call_count, - 1, - "Generic OTEL logger should have received the log exactly once.", - ) - - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_tracing") - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_metrics") - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_logs") - async def test_langfuse_otel_env_config_includes_v4_ingestion_header( - self, mock_logs, mock_metrics, mock_tracing - ): - logger = LangfuseOtelLogger() - - headers = OpenTelemetry._get_headers_dictionary(logger.config.headers) - - self.assertEqual( - headers["x-langfuse-ingestion-version"], - "4", - ) - self.assertTrue(headers["Authorization"].startswith("Basic ")) - - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_tracing") - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_metrics") - @patch("litellm.integrations.opentelemetry.OpenTelemetry._init_logs") - async def test_langfuse_otel_dynamic_headers_include_v4_ingestion_header( - self, mock_logs, mock_metrics, mock_tracing - ): - logger = LangfuseOtelLogger() - - headers = logger.construct_dynamic_otel_headers( - StandardCallbackDynamicParams( - langfuse_public_key="pk-lf-dynamic", - langfuse_secret_key="sk-lf-dynamic", - ) - ) - - self.assertIsNotNone(headers) - self.assertEqual( - headers["x-langfuse-ingestion-version"], - "4", - ) - self.assertTrue(headers["Authorization"].startswith("Basic ")) - if __name__ == "__main__": unittest.main() diff --git a/tests/test_spend_logs.py b/tests/test_spend_logs.py index 57109ea01e5..7e641f1f301 100644 --- a/tests/test_spend_logs.py +++ b/tests/test_spend_logs.py @@ -2,157 +2,10 @@ import os # What this tests? ## Tests /spend endpoints. -import pytest, uuid, json -import asyncio +import pytest import aiohttp -async def generate_key(session, models=[], team_id=None): - url = "http://0.0.0.0:4000/key/generate" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = { - "models": models, - "duration": None, - } - if team_id is not None: - data["team_id"] = team_id - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - -async def chat_completion(session, key, model="gpt-3.5-turbo"): - url = "http://0.0.0.0:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": model, - "messages": [ - {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": f"Hello! {uuid.uuid4()}"}, - ], - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - return await response.json() - - -async def get_spend_logs(session, request_id=None, api_key=None): - if api_key is not None: - url = f"http://0.0.0.0:4000/spend/logs?api_key={api_key}" - else: - url = f"http://0.0.0.0:4000/spend/logs?request_id={request_id}" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - -async def generate_org(session: aiohttp.ClientSession) -> dict: - """ - Generate a new organization using the API. - - Args: - session: aiohttp client session - - Returns: - dict: Response containing org_id - """ - url = "http://0.0.0.0:4000/organization/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - - request_body = { - "organization_alias": f"test-org-{uuid.uuid4()}", - } - - async with session.post(url, headers=headers, json=request_body) as response: - return await response.json() - - -async def generate_team(session: aiohttp.ClientSession, org_id: str) -> dict: - """ - Generate a new team within an organization using the API. - - Args: - session: aiohttp client session - org_id: Organization ID to create the team in - - Returns: - dict: Response containing team_id - """ - url = "http://0.0.0.0:4000/team/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = {"organization_id": org_id} - - async with session.post(url, headers=headers, json=data) as response: - return await response.json() - - -@pytest.mark.skip( - reason="Flaky in CI: /spend/logs?request_id=... returns 500 even after a 20s wait for the spend log to be written. Spend-log accuracy is covered by tests/unit/proxy/spend_tracking/ and the proxy_spend_accuracy_tests CircleCI job." -) -@pytest.mark.asyncio -async def test_spend_logs_with_org_id(): - """ - - Create Organization - - Create Team in organization - - Create Key in organization - - Make call (makes sure it's in spend logs) - - Get request id from logs - - Assert spend logs have correct org_id and team_id - """ - async with aiohttp.ClientSession() as session: - org_gen = await generate_org(session=session) - print("org_gen: ", json.dumps(org_gen, indent=4, default=str)) - org_id = org_gen["organization_id"] - team_gen = await generate_team(session=session, org_id=org_id) - print("team_gen: ", json.dumps(team_gen, indent=4, default=str)) - team_id = team_gen["team_id"] - key_gen = await generate_key(session=session, team_id=team_id) - print("key_gen: ", json.dumps(key_gen, indent=4, default=str)) - key = key_gen["key"] - response = await chat_completion(session=session, key=key) - await asyncio.sleep(20) - spend_logs_response = await get_spend_logs( - session=session, request_id=response["id"] - ) - print( - "spend_logs_response: ", - json.dumps(spend_logs_response, indent=4, default=str), - ) - spend_logs_response = spend_logs_response[0] - assert spend_logs_response["metadata"]["user_api_key_org_id"] == org_id - assert spend_logs_response["metadata"]["user_api_key_team_id"] == team_id - assert spend_logs_response["team_id"] == team_id - - async def get_predict_spend_logs(session): url = "http://0.0.0.0:4000/global/predict/spend/logs" headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} diff --git a/tests/test_team.py b/tests/test_team.py deleted file mode 100644 index e7b5d0cac6d..00000000000 --- a/tests/test_team.py +++ /dev/null @@ -1,912 +0,0 @@ -import os -# What this tests ? -## Tests /team endpoints. -import pytest -import asyncio -import aiohttp -import time, uuid -from openai import AsyncOpenAI -from typing import Optional -import openai -from unittest.mock import MagicMock, patch - - -async def get_user_info(session, get_user, call_user, view_all: Optional[bool] = None): - """ - Make sure only models user has access to are returned - """ - if view_all is True: - url = "http://localhost:4000/user/info" - else: - url = f"http://localhost:4000/user/info?user_id={get_user}" - headers = { - "Authorization": f"Bearer {call_user}", - "Content-Type": "application/json", - } - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - print(response_text) - print() - - if status != 200: - if call_user != get_user: - return status - else: - print(f"call_user: {call_user}; get_user: {get_user}") - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - -async def wait_for_team_member_spend_update( - session, user_id, team_id, expected_min_spend, max_wait=10 -): - """ - Wait for the team member spend update to be committed to the database. - Polls the user info endpoint until the spend is updated. - This is needed because spend updates are queued asynchronously and committed periodically. - """ - start_time = time.time() - initial_spend = None - while time.time() - start_time < max_wait: - try: - user_info = await get_user_info(session, user_id, call_user=os.environ["LITELLM_MASTER_KEY"]) - if user_info.get("teams"): - for team in user_info["teams"]: - if team.get("team_id") == team_id: - for membership in team.get("team_memberships", []): - spend = membership.get("spend", 0.0) - if initial_spend is None: - initial_spend = spend - print(f"Initial team member spend: {spend}") - - if spend >= expected_min_spend: - print( - f"[OK] Team member spend updated: {spend} >= {expected_min_spend}" - ) - return True - - print( - f"[WAITING] Team member spend: {spend}, expected >= {expected_min_spend}, elapsed: {time.time() - start_time:.1f}s" - ) - await asyncio.sleep(0.5) - except Exception as e: - print(f"Error checking team member spend: {e}") - await asyncio.sleep(0.5) - print( - f"[TIMEOUT] Timeout waiting for team member spend update (expected >= {expected_min_spend})" - ) - return False - - -async def new_user( - session, - i, - user_id=None, - budget=None, - budget_duration=None, - models=["azure-models"], - team_id=None, - user_email=None, -): - url = "http://localhost:4000/user/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = { - "models": models, - "aliases": {"mistral-7b": "gpt-3.5-turbo"}, - "duration": None, - "max_budget": budget, - "budget_duration": budget_duration, - "user_email": user_email, - } - - if user_id is not None: - data["user_id"] = user_id - - if team_id is not None: - data["team_id"] = team_id - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response {i} (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception( - f"Request {i} did not return a 200 status code: {status}, response: {response_text}" - ) - - return await response.json() - - -async def add_member( - session, i, team_id, user_id=None, user_email=None, max_budget=None, members=None -): - url = "http://localhost:4000/team/member_add" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = {"team_id": team_id, "member": {"role": "user"}} - if user_email is not None: - data["member"]["user_email"] = user_email - elif user_id is not None: - data["member"]["user_id"] = user_id - elif members is not None: - data["member"] = members - - if max_budget is not None: - data["max_budget_in_team"] = max_budget - - print("sent data: {}".format(data)) - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"ADD MEMBER Response {i} (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request {i} did not return a 200 status code: {status}") - - return await response.json() - - -async def update_member( - session, - i, - team_id, - user_id=None, - user_email=None, - max_budget=None, -): - url = "http://localhost:4000/team/member_update" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = {"team_id": team_id} - if user_id is not None: - data["user_id"] = user_id - elif user_email is not None: - data["user_email"] = user_email - - if max_budget is not None: - data["max_budget_in_team"] = max_budget - - print("sent data: {}".format(data)) - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"ADD MEMBER Response {i} (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception( - f"Request {i} did not return a 200 status code: {status}, response: {response_text}" - ) - - return await response.json() - - -async def delete_member(session, i, team_id, user_id=None, user_email=None): - url = "http://localhost:4000/team/member_delete" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = {"team_id": team_id} - if user_id is not None: - data["user_id"] = user_id - elif user_email is not None: - data["user_email"] = user_email - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response {i} (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request {i} did not return a 200 status code: {status}") - - return await response.json() - - -async def generate_key( - session, - i, - budget=None, - budget_duration=None, - models=["azure-models", "gpt-4", "dall-e-3"], - team_id=None, -): - url = "http://localhost:4000/key/generate" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = { - "models": models, - "duration": None, - "max_budget": budget, - "budget_duration": budget_duration, - } - if team_id is not None: - data["team_id"] = team_id - - print(f"data: {data}") - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response {i} (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request {i} did not return a 200 status code: {status}") - - return await response.json() - - -async def chat_completion(session, key, model="gpt-4"): - url = "http://localhost:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": model, - "messages": [ - {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": "Hello!"}, - ], - } - - for i in range(3): - try: - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception( - f"Request did not return a 200 status code: {status}. Response: {response_text}" - ) - - return await response.json() - except Exception as e: - if "Request did not return a 200 status code" in str(e): - raise e - else: - pass - - -async def new_team(session, i, user_id=None, member_list=None, model_aliases=None): - import json - - url = "http://localhost:4000/team/new" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = {"team_alias": "my-new-team"} - if user_id is not None: - data["members_with_roles"] = [{"role": "user", "user_id": user_id}] - elif member_list is not None: - data["members_with_roles"] = member_list - - if model_aliases is not None: - data["model_aliases"] = model_aliases - - print(f"data: {data}") - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response {i} (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request {i} did not return a 200 status code: {status}") - - return await response.json() - - -async def update_team(session, i, team_id, user_id=None, member_list=None, **kwargs): - url = "http://localhost:4000/team/update" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = {"team_id": team_id, **kwargs} - if user_id is not None: - data["members_with_roles"] = [{"role": "user", "user_id": user_id}] - elif member_list is not None: - data["members_with_roles"] = member_list - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response {i} (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request {i} did not return a 200 status code: {status}") - - return await response.json() - - -async def delete_team( - session, - i, - team_id, -): - url = "http://localhost:4000/team/delete" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - data = { - "team_ids": [team_id], - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(f"Response {i} (Status code: {status}):") - print(response_text) - print() - - if status != 200: - raise Exception(f"Request {i} did not return a 200 status code: {status}") - - return await response.json() - - -async def list_teams( - session, - i, -): - url = "http://localhost:4000/team/list" - headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"} - - async with session.get(url, headers=headers) as response: - status = response.status - if status != 200: - raise Exception(f"Request {i} did not return a 200 status code: {status}") - - return await response.json() - - -@pytest.mark.asyncio -async def test_team_new(): - """ - Make 20 parallel calls to /user/new. Assert all worked. - """ - user_id = f"{uuid.uuid4()}" - async with aiohttp.ClientSession() as session: - new_user(session=session, i=0, user_id=user_id) - tasks = [new_team(session, i, user_id=user_id) for i in range(1, 11)] - await asyncio.gather(*tasks) - - -async def get_team_info(session, get_team, call_key): - url = f"http://localhost:4000/team/info?team_id={get_team}" - headers = { - "Authorization": f"Bearer {call_key}", - "Content-Type": "application/json", - } - - async with session.get(url, headers=headers) as response: - status = response.status - response_text = await response.text() - print(response_text) - print() - - if status == 404: - raise openai.NotFoundError( - message="404 received", response=MagicMock(), body=None - ) - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - return await response.json() - - -@pytest.mark.asyncio -async def test_team_info(): - """ - Scenario 1: - - test with admin key -> expect to work - Scenario 2: - - test with team key -> expect to work - Scenario 3: - - test with non-team key -> expect to fail - """ - async with aiohttp.ClientSession() as session: - """ - Scenario 1 - as admin - """ - new_team_data = await new_team( - session, - 0, - ) - team_id = new_team_data["team_id"] - ## as admin ## - await get_team_info(session=session, get_team=team_id, call_key=os.environ["LITELLM_MASTER_KEY"]) - """ - Scenario 2 - as team key - """ - key_gen = await generate_key(session=session, i=0, team_id=team_id) - key = key_gen["key"] - - await get_team_info(session=session, get_team=team_id, call_key=key) - - """ - Scenario 3 - as non-team key - """ - key_gen = await generate_key(session=session, i=0) - key = key_gen["key"] - - try: - await get_team_info(session=session, get_team=team_id, call_key=key) - pytest.fail("Expected call to fail") - except Exception as e: - pass - - -""" -- Create team -- Add user (user exists in db) -- Update team -- Check if it works -""" - -""" -- Create team -- Add user (user doesn't exist in db) -- Update team -- Check if it works -""" - - -@pytest.mark.asyncio -async def test_team_update_sc_2(): - """ - - Create team - - Add 3 users (doesn't exist in db) - - Change team alias - - Check if it works - - Assert team object unchanged besides team alias - """ - async with aiohttp.ClientSession() as session: - ## Create admin - admin_user = f"{uuid.uuid4()}" - await new_user(session=session, i=0, user_id=admin_user) - ## Create team with 1 admin and 1 user - member_list = [ - {"role": "admin", "user_id": admin_user}, - ] - team_data = await new_team(session=session, i=0, member_list=member_list) - ## Create 10 normal users - members = [ - {"role": "user", "user_id": f"krrish_{uuid.uuid4()}@berri.ai"} - for _ in range(10) - ] - await add_member( - session=session, i=0, team_id=team_data["team_id"], members=members - ) - ## ASSERT TEAM SIZE - team_info = await get_team_info( - session=session, get_team=team_data["team_id"], call_key=os.environ["LITELLM_MASTER_KEY"] - ) - - assert len(team_info["team_info"]["members_with_roles"]) == 12 - - ## CHANGE TEAM ALIAS - - new_team_data = await update_team( - session=session, i=0, team_id=team_data["team_id"], team_alias="my-new-team" - ) - - assert new_team_data["data"]["team_alias"] == "my-new-team" - print(f"team_data: {team_data}") - ## assert rest of object is the same - for k, v in new_team_data["data"].items(): - if k == "members_with_roles": - assert len(new_team_data["data"][k]) == len( - team_info["team_info"]["members_with_roles"] - ) - elif ( - k == "created_at" - or k == "updated_at" - or k == "model_spend" - or k == "model_max_budget" - or k == "model_id" - or k == "litellm_organization_table" - or k == "object_permission_id" - or k == "object_permission" - or k == "litellm_model_table" - or k == "policies" - or k == "allow_team_guardrail_config" - or k == "projects" - ): - pass - else: - assert new_team_data["data"][k] == team_data[k] - - -@pytest.mark.asyncio -async def test_team_member_add_email(): - from tests.test_users import get_user_info - - async with aiohttp.ClientSession() as session: - ## Create admin - admin_user = f"{uuid.uuid4()}" - await new_user(session=session, i=0, user_id=admin_user) - ## Create team with 1 admin and 1 user - member_list = [ - {"role": "admin", "user_id": admin_user}, - ] - team_data = await new_team(session=session, i=0, member_list=member_list) - ## Add 1 user via email - user_email = "krrish{}@berri.ai".format(uuid.uuid4()) - new_user_info = await new_user(session=session, i=0, user_email=user_email) - new_member = {"role": "user", "user_email": user_email} - await add_member( - session=session, i=0, team_id=team_data["team_id"], members=[new_member] - ) - - ## check user info to confirm user is in team - updated_user_info = await get_user_info( - session=session, get_user=new_user_info["user_id"], call_user=os.environ["LITELLM_MASTER_KEY"] - ) - - print(updated_user_info) - - ## check if team in user table - is_team_in_list: bool = False - for team in updated_user_info["teams"]: - if team_data["team_id"] == team["team_id"]: - is_team_in_list = True - assert is_team_in_list - - -@pytest.mark.asyncio -async def test_team_delete(): - """ - - Create team - - Create key for team - - Check if key works - - Delete team - """ - async with aiohttp.ClientSession() as session: - ## Create admin - admin_user = f"{uuid.uuid4()}" - await new_user(session=session, i=0, user_id=admin_user) - ## Create normal user - normal_user = f"{uuid.uuid4()}" - await new_user(session=session, i=0, user_id=normal_user) - ## Create team with 1 admin and 1 user - member_list = [ - {"role": "admin", "user_id": admin_user}, - {"role": "user", "user_id": normal_user}, - ] - team_data = await new_team(session=session, i=0, member_list=member_list) - - ## ASSERT USER MEMBERSHIP IS CREATED - user_info = await get_user_info( - session=session, get_user=normal_user, call_user=os.environ["LITELLM_MASTER_KEY"] - ) - assert len(user_info["teams"]) == 1 - - ## Create key - key_gen = await generate_key(session=session, i=0, team_id=team_data["team_id"]) - key = key_gen["key"] - ## Test key - # response = await chat_completion(session=session, key=key) - ## Delete team - await delete_team(session=session, i=0, team_id=team_data["team_id"]) - - ## ASSERT USER MEMBERSHIP IS DELETED - user_info = await get_user_info( - session=session, get_user=normal_user, call_user=os.environ["LITELLM_MASTER_KEY"] - ) - assert len(user_info["teams"]) == 0 - - ## ASSERT TEAM INFO NOW RETURNS A 404 - with pytest.raises(openai.NotFoundError): - await get_team_info( - session=session, get_team=team_data["team_id"], call_key=os.environ["LITELLM_MASTER_KEY"] - ) - - -@pytest.mark.parametrize("dimension", ["user_id", "user_email"]) -@pytest.mark.asyncio -async def test_member_delete(dimension): - """ - - Create team - - Add member - - Get team info (check if member in team) - - Delete member - - Get team info (check if member in team) - """ - async with aiohttp.ClientSession() as session: - # Create Team - ## Create admin - admin_user = f"{uuid.uuid4()}" - await new_user(session=session, i=0, user_id=admin_user) - ## Create normal user - normal_user = f"{uuid.uuid4()}" - normal_user_email = "{}@berri.ai".format(normal_user) - print(f"normal_user: {normal_user}") - await new_user( - session=session, i=0, user_id=normal_user, user_email=normal_user_email - ) - ## Create team with 1 admin and 1 user - member_list = [ - {"role": "admin", "user_id": admin_user}, - ] - if dimension == "user_id": - member_list.append({"role": "user", "user_id": normal_user}) - elif dimension == "user_email": - member_list.append({"role": "user", "user_email": normal_user_email}) - team_data = await new_team(session=session, i=0, member_list=member_list) - - user_in_team = False - for member in team_data["members_with_roles"]: - if dimension == "user_id" and member["user_id"] == normal_user: - user_in_team = True - elif ( - dimension == "user_email" and member["user_email"] == normal_user_email - ): - user_in_team = True - - assert ( - user_in_team is True - ), "User not in team. Team list={}, User details - id={}, email={}. Dimension={}".format( - team_data["members_with_roles"], normal_user, normal_user_email, dimension - ) - # Delete member - if dimension == "user_id": - updated_team_data = await delete_member( - session=session, i=0, team_id=team_data["team_id"], user_id=normal_user - ) - elif dimension == "user_email": - updated_team_data = await delete_member( - session=session, - i=0, - team_id=team_data["team_id"], - user_email=normal_user_email, - ) - print(f"updated_team_data: {updated_team_data}") - user_in_team = False - for member in team_data["members_with_roles"]: - if dimension == "user_id" and member["user_id"] == normal_user: - user_in_team = True - elif ( - dimension == "user_email" and member["user_email"] == normal_user_email - ): - user_in_team = True - - assert user_in_team is True - - -@pytest.mark.asyncio -async def test_users_in_team_budget(): - """ - - Create User - - Create Team with User - - Add User to team with budget = 0.0000001 - - Make Call 1 -> pass - - Make Call 2 -> fail - """ - get_user = f"krrish_{time.time()}@berri.ai" - async with aiohttp.ClientSession() as session: - # IMPORTANT: Create team first, then create user with team_id. - # This order is critical for the test to work correctly: - # - When a user is created with team_id, the API key gets team_id set from the start - # - This ensures spend tracking and budget enforcement work correctly - # - If we create the user first (without team_id) and then add them to a team, - # the key's team_id remains None, breaking team budget tracking - # DO NOT change this order - it's testing the intended flow where keys are - # associated with teams at creation time. - team = await new_team(session, 0, user_id=None) - print(f"[DEBUG] Created team: {team['team_id']}") - print(f"[DEBUG] Full team data: {team}") - - # Create user with team_id so the key is associated with the team from the start - key_gen = await new_user( - session, - 0, - user_id=get_user, - budget=10, - budget_duration="5s", - models=["fake-openai-endpoint"], - team_id=team["team_id"], - ) - key = key_gen["key"] - print(f"[DEBUG] Created user '{get_user}' with key: {key}") - print(f"[DEBUG] User budget: 10, budget_duration: 5s") - print(f"[DEBUG] Key team_id: {team['team_id']}") - - # Check user info BEFORE updating member budget - user_info_before = await get_user_info(session, get_user, call_user=os.environ["LITELLM_MASTER_KEY"]) - print(f"[DEBUG] User info BEFORE update_member:") - print(f" - User budget: {user_info_before.get('max_budget')}") - print(f" - User spend: {user_info_before.get('spend')}") - if user_info_before.get("teams"): - for team_info in user_info_before["teams"]: - if team_info.get("team_id") == team["team_id"]: - print(f" - Team memberships: {team_info.get('team_memberships')}") - - # update user to have budget = 0.0000001 - update_result = await update_member( - session, 0, team_id=team["team_id"], user_id=get_user, max_budget=0.0000001 - ) - print(f"[DEBUG] Updated member budget to 0.0000001") - print(f"[DEBUG] Update result: {update_result}") - - # Check user info AFTER updating member budget - user_info_after = await get_user_info(session, get_user, call_user=os.environ["LITELLM_MASTER_KEY"]) - print(f"[DEBUG] User info AFTER update_member:") - print(f" - User budget: {user_info_after.get('max_budget')}") - print(f" - User spend: {user_info_after.get('spend')}") - if user_info_after.get("teams"): - for team_info in user_info_after["teams"]: - if team_info.get("team_id") == team["team_id"]: - print(f" - Team: {team_info.get('team_id')}") - for membership in team_info.get("team_memberships", []): - print(f" - Membership: {membership}") - if "litellm_budget_table" in membership: - budget_table = membership["litellm_budget_table"] - print(f" - Max budget: {budget_table.get('max_budget')}") - print(f" - Current spend: {membership.get('spend', 0)}") - - # Call 1 - print("\n[DEBUG] ===== Making Call 1 =====") - result = await chat_completion(session, key, model="fake-openai-endpoint") - print(f"[DEBUG] Call 1 PASSED (expected)") - print(f"[DEBUG] Call 1 result: {result}") - # Extract cost from result if available - if isinstance(result, dict): - usage = result.get("usage", {}) - print(f"[DEBUG] Call 1 usage: {usage}") - - # Wait for spend to be committed to database before checking budget - # Spend updates are queued asynchronously and committed periodically (every minute), - # so we need to wait for the spend from Call 1 to be persisted - print("\n[DEBUG] ===== Waiting for spend to be committed =====") - print("Waiting for team member spend to be committed to database...") - print( - "Note: Spend updates are flushed periodically, this may take up to 90 seconds..." - ) - spend_updated = await wait_for_team_member_spend_update( - session, get_user, team["team_id"], 0.0000001, max_wait=90 - ) - if not spend_updated: - pytest.fail( - "Team member spend was not updated within 90s. " - "The spend update queue may not have flushed, or the model may have 0 cost." - ) - - # Check user info BEFORE Call 2 - user_info_before_call2 = await get_user_info( - session, get_user, call_user=os.environ["LITELLM_MASTER_KEY"] - ) - print(f"\n[DEBUG] User info BEFORE Call 2:") - print(f" - User budget: {user_info_before_call2.get('max_budget')}") - print(f" - User spend: {user_info_before_call2.get('spend')}") - if user_info_before_call2.get("teams"): - for team_info in user_info_before_call2["teams"]: - if team_info.get("team_id") == team["team_id"]: - print(f" - Team: {team_info.get('team_id')}") - for membership in team_info.get("team_memberships", []): - if "litellm_budget_table" in membership: - budget_table = membership["litellm_budget_table"] - current_spend = membership.get("spend", 0) - max_budget = budget_table.get("max_budget") - print(f" - Max budget in team: {max_budget}") - print(f" - Current spend in team: {current_spend}") - print( - f" - Budget remaining: {max_budget - current_spend}" - ) - print(f" - Should fail?: {current_spend >= max_budget}") - - # Call 2 - print("\n[DEBUG] ===== Making Call 2 =====") - call2_failed = False - call2_error = None - call2_status = None - try: - # Capture the response to check status code - url = "http://localhost:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": "fake-openai-endpoint", - "messages": [ - {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": "Hello!"}, - ], - } - async with session.post(url, headers=headers, json=data) as response: - call2_status = response.status - response_text = await response.text() - print(f"[DEBUG] Call 2 status code: {call2_status}") - print(f"[DEBUG] Call 2 response: {response_text}") - - if call2_status != 200: - call2_failed = True - call2_error = f"Status {call2_status}: {response_text}" - raise Exception(call2_error) - else: - # Call succeeded when it should have failed - print(f"[ERROR] Call 2 PASSED when it should have FAILED!") - print(f"[ERROR] Response was 200 OK") - - except Exception as e: - if call2_failed: - print(f"[DEBUG] Call 2 FAILED (expected): {e}") - print(f"[DEBUG] Checking if error message indicates budget exceeded...") - else: - call2_error = str(e) - print(f"[DEBUG] Call 2 raised exception: {e}") - - # Check user info AFTER Call 2 - user_info_after_call2 = await get_user_info( - session, get_user, call_user=os.environ["LITELLM_MASTER_KEY"] - ) - print(f"\n[DEBUG] User info AFTER Call 2:") - print(f" - User budget: {user_info_after_call2.get('max_budget')}") - print(f" - User spend: {user_info_after_call2.get('spend')}") - if user_info_after_call2.get("teams"): - for team_info in user_info_after_call2["teams"]: - if team_info.get("team_id") == team["team_id"]: - print(f" - Team: {team_info.get('team_id')}") - for membership in team_info.get("team_memberships", []): - if "litellm_budget_table" in membership: - budget_table = membership["litellm_budget_table"] - print(f" - Max budget: {budget_table.get('max_budget')}") - print(f" - Current spend: {membership.get('spend', 0)}") - - # Assert Call 2 failed - if not call2_failed: - error_msg = ( - f"\n[FAILURE] Call 2 should have failed but it passed!\n" - f"Expected: Budget enforcement to block the call\n" - f"Actual: Call returned status {call2_status}\n" - f"Team member budget: 0.0000001\n" - f"User budget: {user_info_before_call2.get('max_budget')}\n" - f"User spend before call: {user_info_before_call2.get('spend')}\n" - ) - # Add team member info if available - if user_info_before_call2.get("teams"): - for team_info in user_info_before_call2["teams"]: - if team_info.get("team_id") == team["team_id"]: - for membership in team_info.get("team_memberships", []): - if "litellm_budget_table" in membership: - error_msg += f"Team member spend before call: {membership.get('spend', 0)}\n" - error_msg += f"Team member max budget: {membership['litellm_budget_table'].get('max_budget')}\n" - pytest.fail(error_msg) - - # Check the error message contains budget exceeded - if call2_error and "Budget has been exceeded" not in call2_error: - pytest.fail( - f"Call 2 failed but not with expected error message.\n" - f"Expected error to contain: 'Budget has been exceeded'\n" - f"Actual error: {call2_error}" - ) - - print("[DEBUG] Call 2 failed as expected with budget exceeded error") - - ## Check user info - user_info = await get_user_info(session, get_user, call_user=os.environ["LITELLM_MASTER_KEY"]) - - assert ( - user_info["teams"][0]["team_memberships"][0]["litellm_budget_table"][ - "max_budget" - ] - == 0.0000001 - ) diff --git a/tests/test_team_members.py b/tests/test_team_members.py deleted file mode 100644 index 72db6bfe7ae..00000000000 --- a/tests/test_team_members.py +++ /dev/null @@ -1,202 +0,0 @@ -import os -import pytest -import requests -import time -from typing import Dict, List -import logging -from litellm._uuid import uuid - -# Configure logging -logging.basicConfig(level=logging.INFO) -logger = logging.getLogger(__name__) - - -class TeamAPI: - def __init__(self, base_url: str, auth_token: str): - self.base_url = base_url - self.headers = { - "Authorization": f"Bearer {auth_token}", - "Content-Type": "application/json", - } - - def create_team(self, team_alias: str, models: List[str] = None) -> Dict: - """Create a new team""" - # Generate a unique team_id using uuid - team_id = f"test_team_{uuid.uuid4().hex[:8]}" - - data = { - "team_id": team_id, - "team_alias": team_alias, - "models": models or ["o3-mini"], - } - - response = requests.post( - f"{self.base_url}/team/new", headers=self.headers, json=data - ) - response.raise_for_status() - logger.info(f"Created new team: {team_id}") - return response.json(), team_id - - def get_team_info(self, team_id: str) -> Dict: - """Get current team information""" - response = requests.get( - f"{self.base_url}/team/info", - headers=self.headers, - params={"team_id": team_id}, - ) - response.raise_for_status() - return response.json() - - def add_team_member(self, team_id: str, user_email: str, role: str) -> Dict: - """Add a single team member""" - data = {"team_id": team_id, "member": [{"role": role, "user_id": user_email}]} - response = requests.post( - f"{self.base_url}/team/member_add", headers=self.headers, json=data - ) - response.raise_for_status() - return response.json() - - def delete_team_member(self, team_id: str, user_id: str) -> Dict: - """Delete a team member - - Args: - team_id (str): ID of the team - user_id (str): User ID to remove from team - - Returns: - Dict: Response from the API - """ - data = {"team_id": team_id, "user_id": user_id} - response = requests.post( - f"{self.base_url}/team/member_delete", headers=self.headers, json=data - ) - response.raise_for_status() - return response.json() - - -@pytest.fixture -def api_client(): - """Fixture for TeamAPI client""" - base_url = "http://localhost:4000" - auth_token = os.environ["LITELLM_MASTER_KEY"] # Replace with your token - return TeamAPI(base_url, auth_token) - - -@pytest.fixture -def new_team(api_client): - """Fixture that creates a new team for each test""" - team_alias = f"Test Team {uuid.uuid4().hex[:6]}" - team_response, team_id = api_client.create_team(team_alias) - logger.info(f"Created test team: {team_id} ({team_alias})") - return team_id - - -def verify_member_in_team(team_info: Dict, user_email: str) -> bool: - """Verify if a member exists in team""" - return any( - member["user_id"] == user_email - for member in team_info["team_info"]["members_with_roles"] - ) - - -def test_team_creation(api_client): - """Test team creation""" - team_alias = f"Test Team {uuid.uuid4().hex[:6]}" - team_response, team_id = api_client.create_team(team_alias) - - # Verify team was created - team_info = api_client.get_team_info(team_id) - assert team_info["team_id"] == team_id - assert team_info["team_info"]["team_alias"] == team_alias - assert "o3-mini" in team_info["team_info"]["models"] - - -def test_add_single_member(api_client, new_team): - """Test adding a single member to a new team""" - # Get initial team info - initial_info = api_client.get_team_info(new_team) - initial_size = len(initial_info["team_info"]["members_with_roles"]) - - # Add new member - test_email = f"pytest_user_{uuid.uuid4().hex[:6]}@mycompany.com" - api_client.add_team_member(new_team, test_email, "user") - - # Allow time for system to process - time.sleep(1) - - # Verify addition - updated_info = api_client.get_team_info(new_team) - updated_size = len(updated_info["team_info"]["members_with_roles"]) - - # Assertions - assert verify_member_in_team( - updated_info, test_email - ), f"Member {test_email} not found in team" - assert ( - updated_size == initial_size + 1 - ), f"Team size did not increase by 1 (was {initial_size}, now {updated_size})" - - -def test_member_deletion(api_client, new_team): - """Test that member deletion works correctly and removes all instances of a user""" - # Add a test user - user_id = f"pytest_user_{uuid.uuid4().hex[:6]}" - api_client.add_team_member(new_team, user_id, "user") - time.sleep(1) - - # Verify user was added - team_info_before = api_client.get_team_info(new_team) - assert verify_member_in_team( - team_info_before, user_id - ), "User was not added successfully" - - initial_size = len(team_info_before["team_info"]["members_with_roles"]) - - # Attempt to delete the same user multiple times (5 times) - for i in range(5): - logger.info(f"Attempting deletion {i+1}/5") - if i == 0: - # First deletion should succeed - api_client.delete_team_member(new_team, user_id) - time.sleep(1) - else: - # Subsequent deletions should raise an error - try: - api_client.delete_team_member(new_team, user_id) - pytest.fail("Expected HTTPError for duplicate deletion") - except requests.exceptions.HTTPError as e: - logger.info( - f"Expected error received on deletion attempt {i+1}: {str(e)}" - ) - - # Verify final state - final_info = api_client.get_team_info(new_team) - final_size = len(final_info["team_info"]["members_with_roles"]) - - # Verify user is completely removed - assert not verify_member_in_team( - final_info, user_id - ), "User still exists in team after deletion" - - # Verify only one member was removed - assert ( - final_size == initial_size - 1 - ), f"Team size changed unexpectedly (was {initial_size}, now {final_size})" - - -def test_delete_nonexistent_member(api_client, new_team): - """Test that attempting to delete a nonexistent member raises appropriate error""" - nonexistent_user = f"nonexistent_{uuid.uuid4().hex[:6]}" - - # Verify user doesn't exist first - team_info = api_client.get_team_info(new_team) - assert not verify_member_in_team( - team_info, nonexistent_user - ), "Test setup error: nonexistent user somehow exists" - - # Attempt to delete nonexistent user - with pytest.raises(requests.exceptions.HTTPError) as exc_info: - api_client.delete_team_member(new_team, nonexistent_user) - e = exc_info.value - logger.info(f"Expected error received: {str(e)}") - assert e.response.status_code == 400 diff --git a/tests/test_users.py b/tests/test_users.py index 55caeb592c8..9b19848694a 100644 --- a/tests/test_users.py +++ b/tests/test_users.py @@ -5,10 +5,7 @@ import pytest import asyncio import aiohttp import time -from openai import AsyncOpenAI -from tests.test_team import list_teams from typing import Optional -from fastapi import HTTPException async def new_user( @@ -41,16 +38,6 @@ async def new_user( return await response.json() -@pytest.mark.asyncio -async def test_user_new(): - """ - Make 20 parallel calls to /user/new. Assert all worked. - """ - async with aiohttp.ClientSession() as session: - tasks = [new_user(session, i) for i in range(1, 11)] - await asyncio.gather(*tasks) - - async def get_user_info(session, get_user, call_user, view_all: Optional[bool] = None): """ Make sure only models user has access to are returned @@ -79,37 +66,6 @@ async def get_user_info(session, get_user, call_user, view_all: Optional[bool] = return await response.json() -@pytest.mark.asyncio -async def test_user_info(): - """ - Get user info - - as admin - - as user themself - - as random - """ - get_user = f"krrish_{time.time()}@berri.ai" - async with aiohttp.ClientSession() as session: - key_gen = await new_user(session, 0, user_id=get_user) - key = key_gen["key"] - ## as admin ## - resp = await get_user_info( - session=session, get_user=get_user, call_user=os.environ["LITELLM_MASTER_KEY"] - ) - assert isinstance(resp["user_info"], dict) - assert len(resp["user_info"]) > 0 - ## as user themself ## - resp = await get_user_info(session=session, get_user=get_user, call_user=key) - assert isinstance(resp["user_info"], dict) - assert len(resp["user_info"]) > 0 - # as random user # - key_gen = await new_user(session=session, i=0) - random_key = key_gen["key"] - status = await get_user_info( - session=session, get_user=get_user, call_user=random_key - ) - assert status == 403 - - @pytest.mark.skip(reason="Frequent check on ci/cd leads to read timeout issue.") @pytest.mark.asyncio async def test_users_budgets_reset(): @@ -145,168 +101,6 @@ async def test_users_budgets_reset(): assert reset_at_init_value != reset_at_new_value -async def chat_completion(session, key, model="gpt-4"): - client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000") - messages = [ - {"role": "system", "content": "You are a helpful assistant"}, - {"role": "user", "content": f"Hello! {time.time()}"}, - ] - - data = { - "model": model, - "messages": messages, - } - response = await client.chat.completions.create(**data) - - -async def chat_completion_streaming(session, key, model="gpt-4"): - client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000") - messages = [ - {"role": "system", "content": "You are a helpful assistant"}, - {"role": "user", "content": f"Hello! {time.time()}"}, - ] - - data = {"model": model, "messages": messages, "stream": True} - response = await client.chat.completions.create(**data) - async for chunk in response: - continue - - - - -import json -from litellm._uuid import uuid import pytest -from typing import Dict, Tuple -async def setup_test_users(session: aiohttp.ClientSession) -> Tuple[Dict, Dict]: - """ - Create two test users and an additional key for the first user. - Returns tuple of (user1_data, user2_data) where each contains user info and keys. - """ - # Create two test users - user1 = await new_user( - session=session, - i=0, - budget=100, - budget_duration="30d", - models=["anthropic.claude-haiku-4-5-20251001-v1:0"], - ) - - user2 = await new_user( - session=session, - i=1, - budget=100, - budget_duration="30d", - models=["anthropic.claude-haiku-4-5-20251001-v1:0"], - ) - - print("\nCreated two test users:") - print(f"User 1 ID: {user1['user_id']}") - print(f"User 2 ID: {user2['user_id']}") - - # Create an additional key for user1 - headers = { - "Content-Type": "application/json", - "Authorization": f"Bearer {user1['key']}", - } - - key_payload = { - "user_id": user1["user_id"], - "duration": "7d", - "key_alias": f"test_key_{uuid.uuid4()}", - "models": ["anthropic.claude-haiku-4-5-20251001-v1:0"], - } - - print("\nGenerating additional key for user1...") - key_response = await session.post( - f"http://0.0.0.0:4000/key/generate", headers=headers, json=key_payload - ) - - assert key_response.status == 200, "Failed to generate additional key for user1" - user1_additional_key = await key_response.json() - - print(f"\nGenerated key details:") - print(json.dumps(user1_additional_key, indent=2)) - - # Return both users' data including the additional key - return { - "user_data": user1, - "additional_key": user1_additional_key, - "headers": headers, - }, { - "user_data": user2, - "headers": { - "Content-Type": "application/json", - "Authorization": f"Bearer {user2['key']}", - }, - } - - -async def print_response_details(response: aiohttp.ClientResponse) -> None: - """Helper function to print response details""" - print("\nResponse Details:") - print(f"Status Code: {response.status}") - print("\nResponse Content:") - try: - formatted_json = json.dumps(await response.json(), indent=2) - print(formatted_json) - except json.JSONDecodeError: - print(await response.text()) - - -@pytest.mark.asyncio -async def test_key_update_user_isolation(): - """Test that a user cannot update a key that belongs to another user""" - async with aiohttp.ClientSession() as session: - user1_data, user2_data = await setup_test_users(session) - - # Try to update the key to belong to user2 - update_payload = { - "key": user1_data["additional_key"]["key"], - "user_id": user2_data["user_data"][ - "user_id" - ], # Attempting to change ownership - "metadata": {"purpose": "testing_user_isolation", "environment": "test"}, - } - - print("\nAttempting to update key ownership to user2...") - update_response = await session.post( - f"http://0.0.0.0:4000/key/update", - headers=user1_data["headers"], # Using user1's headers - json=update_payload, - ) - - await print_response_details(update_response) - - # Verify update attempt was rejected - assert ( - update_response.status == 403 - ), "Request should have been rejected with 403 status code" - - -@pytest.mark.asyncio -async def test_key_delete_user_isolation(): - """Test that a user cannot delete a key that belongs to another user""" - async with aiohttp.ClientSession() as session: - user1_data, user2_data = await setup_test_users(session) - - # Try to delete user1's additional key using user2's credentials - delete_payload = { - "keys": [user1_data["additional_key"]["key"]], - } - - print("\nAttempting to delete user1's key using user2's credentials...") - delete_response = await session.post( - f"http://0.0.0.0:4000/key/delete", - headers=user2_data["headers"], - json=delete_payload, - ) - - await print_response_details(delete_response) - - # Verify delete attempt was rejected - assert ( - delete_response.status == 403 - ), "Request should have been rejected with 403 status code" diff --git a/tests/unit/batches/test_main.py b/tests/unit/batches/test_main.py index 3303c13b6a3..fb97e0e4919 100644 --- a/tests/unit/batches/test_main.py +++ b/tests/unit/batches/test_main.py @@ -34,6 +34,16 @@ import pytest import litellm import litellm.batches.main as bm +import asyncio +import datetime +import json +from collections.abc import Mapping +from typing import Final +import httpx +import respx +from pydantic import TypeAdapter +from typing_extensions import ReadOnly, TypedDict +from litellm.integrations.custom_logger import CustomLogger # --------------------------------------------------------------------------- # @@ -986,3 +996,211 @@ async def test_batch_logging_azure_credentials_regression(): print("✓ Batch output files can be fetched with Azure credentials") print("✓ Cost and usage tracking works for Azure batches") print("✓ Backwards compatibility maintained\n") + + +_OPENAI_FILE_JSON: Final = MappingProxyType( + { + "id": "file-abc123", + "object": "file", + "purpose": "batch", + "filename": "batch.jsonl", + "bytes": 416, + "created_at": 1739598666, + "status": "processed", + } +) + + +_OPENAI_BATCH_JSON: Final = MappingProxyType( + { + "id": "batch_abc123", + "object": "batch", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-abc123", + "status": "validating", + "completion_window": "24h", + "created_at": 1739598666, + } +) + + +class _KeyAliasMetadata(TypedDict): + user_api_key_alias: ReadOnly[str | None] + user_api_key_team_alias: ReadOnly[str | None] + + +class _LoggedCall(TypedDict): + call_type: ReadOnly[str] + metadata: ReadOnly[_KeyAliasMetadata] + + +_LOGGED_CALL: Final = TypeAdapter(_LoggedCall) + + +class _SuccessPayloadRecorder(CustomLogger): + def __init__(self, call_type: str) -> None: + super().__init__() + self._call_type: Final = call_type + self.logged: Final = asyncio.Event() + self.payload: _LoggedCall | None = None + + async def async_log_success_event( + self, + kwargs: Mapping[str, object], + response_obj: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + payload: Final = _LOGGED_CALL.validate_python(kwargs["standard_logging_object"]) + if payload["call_type"] != self._call_type: + return + self.payload = payload + self.logged.set() + + +@pytest.mark.asyncio +async def test_acreate_batch_full_crud_and_logging_metadata( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.logging_callback_manager._reset_all_callbacks() + recorder: Final = _SuccessPayloadRecorder("acreate_batch") + monkeypatch.setattr(litellm, "callbacks", [recorder]) + + upload_route: Final = respx_mock.post("https://api.openai.com/v1/files").mock( + return_value=httpx.Response(200, json=dict(_OPENAI_FILE_JSON)) + ) + create_route: Final = respx_mock.post("https://api.openai.com/v1/batches").mock( + return_value=httpx.Response(200, json=dict(_OPENAI_BATCH_JSON)) + ) + retrieve_route: Final = respx_mock.get("https://api.openai.com/v1/batches/batch_abc123").mock( + return_value=httpx.Response(200, json=dict(_OPENAI_BATCH_JSON)) + ) + list_batches_route: Final = respx_mock.get("https://api.openai.com/v1/batches").mock( + return_value=httpx.Response(200, json={"object": "list", "data": [dict(_OPENAI_BATCH_JSON)]}) + ) + respx_mock.get("https://api.openai.com/v1/files/file-abc123/content").mock( + return_value=httpx.Response(200, content=b'{"custom_id": "request-1"}\n') + ) + respx_mock.get("https://api.openai.com/v1/files/file-abc123").mock( + return_value=httpx.Response(200, json=dict(_OPENAI_FILE_JSON)) + ) + respx_mock.delete("https://api.openai.com/v1/files/file-abc123").mock( + return_value=httpx.Response(200, json={"id": "file-abc123", "object": "file", "deleted": True}) + ) + list_files_route: Final = respx_mock.get("https://api.openai.com/v1/files").mock( + return_value=httpx.Response(200, json={"object": "list", "data": [dict(_OPENAI_FILE_JSON)]}) + ) + cancel_route: Final = respx_mock.post("https://api.openai.com/v1/batches/batch_abc123/cancel").mock( + return_value=httpx.Response(200, json={**_OPENAI_BATCH_JSON, "status": "cancelling"}) + ) + + batch_file: Final = ( + "batch.jsonl", + b'{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", ' + b'"body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}}\n', + "application/jsonl", + ) + file_obj: Final = await litellm.acreate_file( + file=batch_file, purpose="batch", custom_llm_provider="openai", api_key="fake-key" + ) + assert file_obj.id == "file-abc123" + upload_body: Final = upload_route.calls.last.request.content + assert b'name="purpose"\r\n\r\nbatch' in upload_body + assert batch_file[1] in upload_body + + extra_metadata_field: Final = { + "user_api_key_alias": "special_api_key_alias", + "user_api_key_team_alias": "special_team_alias", + } + create_batch_response: Final = await litellm.acreate_batch( + completion_window="24h", + endpoint="/v1/chat/completions", + input_file_id=file_obj.id, + custom_llm_provider="openai", + api_key="fake-key", + metadata={"key1": "value1", "key2": "value2"}, + litellm_metadata=extra_metadata_field, + ) + + assert json.loads(create_route.calls.last.request.content) == { + "completion_window": "24h", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-abc123", + "metadata": {"key1": "value1", "key2": "value2"}, + } + assert create_batch_response.id == "batch_abc123" + assert create_batch_response.endpoint == "/v1/chat/completions" + assert create_batch_response.input_file_id == file_obj.id + + await asyncio.wait_for(recorder.logged.wait(), timeout=10) + assert recorder.payload is not None + standard_logging_object: Final = recorder.payload + assert standard_logging_object["metadata"]["user_api_key_alias"] == extra_metadata_field["user_api_key_alias"] + assert ( + standard_logging_object["metadata"]["user_api_key_team_alias"] + == extra_metadata_field["user_api_key_team_alias"] + ) + + retrieved_batch: Final = await litellm.aretrieve_batch( + batch_id=create_batch_response.id, custom_llm_provider="openai", api_key="fake-key" + ) + assert retrieve_route.called + assert retrieved_batch.id == create_batch_response.id + + list_batches: Final = await litellm.alist_batches(custom_llm_provider="openai", limit=2, api_key="fake-key") + assert list_batches_route.calls.last.request.url.params["limit"] == "2" + assert [batch.id for batch in list_batches.data] == ["batch_abc123"] + + file_content: Final = await litellm.afile_content( + file_id=file_obj.id, custom_llm_provider="openai", api_key="fake-key" + ) + assert file_content.content == b'{"custom_id": "request-1"}\n' + + retrieved_file: Final = await litellm.afile_retrieve( + file_id=file_obj.id, custom_llm_provider="openai", api_key="fake-key" + ) + assert retrieved_file.id == file_obj.id + + delete_file_response: Final = await litellm.afile_delete( + file_id=file_obj.id, custom_llm_provider="openai", api_key="fake-key" + ) + assert delete_file_response.id == file_obj.id + + all_files_list: Final = await litellm.afile_list(custom_llm_provider="openai", api_key="fake-key") + assert list_files_route.called + assert [file.id for file in all_files_list.data] == ["file-abc123"] + + cancel_batch_response: Final = await litellm.acancel_batch( + batch_id=create_batch_response.id, custom_llm_provider="openai", api_key="fake-key" + ) + assert cancel_route.called + assert cancel_batch_response.id == create_batch_response.id + + +@pytest.mark.asyncio +async def test_delete_batch_output_file(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + batch_with_output: Final = { + **_OPENAI_BATCH_JSON, + "status": "completed", + "output_file_id": "file-output123", + } + respx_mock.get("https://api.openai.com/v1/batches/batch_abc123").mock( + return_value=httpx.Response(200, json=batch_with_output) + ) + delete_route: Final = respx_mock.delete("https://api.openai.com/v1/files/file-output123").mock( + return_value=httpx.Response(200, json={"id": "file-output123", "object": "file", "deleted": True}) + ) + + batch: Final = await litellm.aretrieve_batch( + batch_id="batch_abc123", custom_llm_provider="openai", api_key="fake-key" + ) + assert batch.output_file_id == "file-output123" + + delete_response: Final = await litellm.afile_delete( + file_id=batch.output_file_id, custom_llm_provider="openai", api_key="fake-key" + ) + assert delete_route.call_count == 1 + assert delete_response.id == "file-output123" + assert delete_response.deleted is True diff --git a/tests/unit/caching/test_caching.py b/tests/unit/caching/test_caching.py index a799c45e4c7..c94ee836f91 100644 --- a/tests/unit/caching/test_caching.py +++ b/tests/unit/caching/test_caching.py @@ -953,6 +953,7 @@ def test_redis_caching_multiple_namespaces(): _TOOL_TURN_ITEM: Final = {"role": "user", "content": "hi"} +_FILE_BLOCK_ITEM: Final = {"type": "file", "file": {"file_data": "data:video/mp4;base64,AAAA", "format": "video/mp4"}} @pytest.mark.parametrize( @@ -964,6 +965,7 @@ _TOOL_TURN_ITEM: Final = {"role": "user", "content": "hi"} pytest.param({"input": [_TOOL_TURN_ITEM] * 5}, False, id="five-responses-items-skip-the-cache"), pytest.param({"input": "one prompt"}, True, id="string-input-is-one-message"), pytest.param({"input": ["a", "b", "c", "d", "e"]}, True, id="embedding-strings-are-not-messages"), + pytest.param({"input": [_FILE_BLOCK_ITEM] * 5}, True, id="embedding-file-blocks-are-not-messages"), ], ) def test_should_use_cache_stops_past_the_default_max_messages(kwargs: dict[str, object], expected: bool) -> None: diff --git a/tests/unit/caching/test_caching_handler.py b/tests/unit/caching/test_caching_handler.py index 22eea9935f6..d98c02667ee 100644 --- a/tests/unit/caching/test_caching_handler.py +++ b/tests/unit/caching/test_caching_handler.py @@ -41,6 +41,7 @@ from litellm.types.utils import ( from litellm.types.llms.openai import ResponsesAPIResponse from collections.abc import Awaitable, Callable from datetime import timedelta, datetime +from typing import Final from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm._logging import verbose_logger @@ -2714,3 +2715,88 @@ async def test_async_get_cache_forgets_the_worker_copy_of_a_stored_response_with assert lookup.cached_result is None assert await handler.dual_cache.async_get_cache(key) is None + + +@pytest.mark.asyncio +async def test_async_get_cache_partial_hit_keeps_file_block_items_uncached() -> None: + setup_cache() + fixed_start: Final = datetime(2026, 1, 1) + caching_handler: Final = LLMCachingHandler(original_function=aembedding, request_kwargs={}, start_time=fixed_start) + model: Final = "gemini/gemini-embedding-2-preview" + logging_obj: Final = LiteLLMLogging( + litellm_call_id=str(uuid.uuid4()), + call_type=CallTypes.aembedding.value, + model=model, + messages=[], + function_id=str(uuid.uuid4()), + stream=False, + start_time=fixed_start, + ) + await caching_handler.async_set_cache( + result=EmbeddingResponse(model=model, data=[Embedding(embedding=[0.1, 0.2], index=0, object="embedding")]), + original_function=aembedding, + kwargs={"model": model, "input": ["a red bus"], "caching": True}, + ) + clip_block: Final = { + "type": "file", + "file": { + "file_data": "data:video/mp4;base64,AAAA", + "format": "video/mp4", + "video_metadata": {"fps": 1, "start_offset": "0s", "end_offset": "1s"}, + }, + "detail": "left for the provider transformation to judge", + } + + cached_response: Final = await caching_handler.async_get_cache( + model=model, + original_function=aembedding, + logging_obj=logging_obj, + start_time=fixed_start, + call_type=CallTypes.aembedding.value, + kwargs={"model": model, "input": [clip_block, "a red bus"], "caching": True}, + ) + + assert cached_response.embedding_all_elements_cache_hit is False + assert cached_response.embedding_uncached_input == [clip_block] + assert cached_response.final_embedding_cached_response is not None + assert cached_response.final_embedding_cached_response.data[1].embedding == [0.1, 0.2] + assert cached_response.final_embedding_cached_response.data[0] is None + + +def test_handle_kwargs_input_answers_400_for_a_single_object_input() -> None: + caching_handler: Final = LLMCachingHandler( + original_function=aembedding, request_kwargs={}, start_time=datetime(2026, 1, 1) + ) + clip_block: Final = {"type": "file", "file": {"file_data": "data:video/mp4;base64,AAAA"}} + with pytest.raises(litellm.BadRequestError, match="string or a list"): + caching_handler.handle_kwargs_input_list_or_str( + {"model": "gemini/gemini-embedding-2-preview", "custom_llm_provider": "gemini", "input": clip_block} + ) + + +@pytest.mark.asyncio +async def test_async_get_cache_answers_400_for_a_single_object_embedding_input() -> None: + setup_cache() + fixed_start: Final = datetime(2026, 1, 1) + caching_handler: Final = LLMCachingHandler(original_function=aembedding, request_kwargs={}, start_time=fixed_start) + model: Final = "gemini/gemini-embedding-2-preview" + logging_obj: Final = LiteLLMLogging( + litellm_call_id=str(uuid.uuid4()), + call_type=CallTypes.aembedding.value, + model=model, + messages=[], + function_id=str(uuid.uuid4()), + stream=False, + start_time=fixed_start, + ) + clip_block: Final = {"type": "file", "file": {"file_data": "data:video/mp4;base64,AAAA"}} + + with pytest.raises(litellm.BadRequestError, match="string or a list"): + await caching_handler.async_get_cache( + model=model, + original_function=aembedding, + logging_obj=logging_obj, + start_time=fixed_start, + call_type=CallTypes.aembedding.value, + kwargs={"model": model, "custom_llm_provider": "gemini", "input": clip_block, "caching": True}, + ) diff --git a/tests/unit/chat_completions/test_dispatch.py b/tests/unit/chat_completions/test_dispatch.py index fbfc4875aa1..4ed75ea5f95 100644 --- a/tests/unit/chat_completions/test_dispatch.py +++ b/tests/unit/chat_completions/test_dispatch.py @@ -134,13 +134,13 @@ def test_native_receives_bound_request_and_original_call_shape() -> None: ) request, call_args, call_kwargs = captured[0] - assert request.bound["model"] == "anthropic/claude-sonnet-4-5" - assert request.bound["messages"] is MESSAGES - assert request.bound["stream"] is True - assert request.bound["api_key"] == "sk-test" - assert request.bound["base_url"] == "https://example.invalid" - assert request.bound["custom_llm_provider"] == "anthropic" - assert request.bound["extra_headers"] == {"x-test": "1"} + assert request.resolved["model"] == "anthropic/claude-sonnet-4-5" + assert request.resolved["messages"] is MESSAGES + assert request.resolved["stream"] is True + assert request.resolved["api_key"] == "sk-test" + assert request.resolved["base_url"] == "https://example.invalid" + assert request.resolved["custom_llm_provider"] == "anthropic" + assert request.resolved["extra_headers"] == {"x-test": "1"} assert request.kwargs is kwargs assert call_args == args assert call_kwargs == kwargs @@ -218,7 +218,7 @@ def test_public_completion_routes_through_dispatch(monkeypatch: pytest.MonkeyPat finally: NATIVE_COMPLETION.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["gpt-4o"] + assert [request.resolved["model"] for request in captured] == ["gpt-4o"] @pytest.mark.asyncio @@ -238,7 +238,7 @@ async def test_public_acompletion_routes_through_dispatch(monkeypatch: pytest.Mo finally: NATIVE_ACOMPLETION.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["gpt-4o"] + assert [request.resolved["model"] for request in captured] == ["gpt-4o"] @pytest.mark.asyncio @@ -257,10 +257,10 @@ def test_sync_completion_request_projects_public_arguments() -> None: expected: Final = ModelResponse() def native(request: NativeCall) -> ModelResponse: - assert request.bound["model"] == "test-model" - assert request.bound["messages"] == MESSAGES - assert request.bound["custom_llm_provider"] == "openai" - assert request.bound["stream"] is True + assert request.resolved["model"] == "test-model" + assert request.resolved["messages"] == MESSAGES + assert request.resolved["custom_llm_provider"] == "openai" + assert request.resolved["stream"] is True return expected binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None) @@ -335,6 +335,6 @@ def test_internal_acompletion_marker_bypasses_native() -> None: def test_positional_parameters_remain_available_to_native_projection() -> None: request: Final = _DISPATCH.request(("anthropic/test-model", MESSAGES, 12.0, 0.25), {}) assert request is not None - assert request.bound["timeout"] == 12.0 - assert request.bound["temperature"] == 0.25 - assert request.bound["messages"] is MESSAGES + assert request.resolved["timeout"] == 12.0 + assert request.resolved["temperature"] == 0.25 + assert request.resolved["messages"] is MESSAGES diff --git a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 5b2187211df..8c6c4059a22 100644 --- a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -3,6 +3,7 @@ import datetime import json import os import unittest +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Final, List, Literal, Optional, Tuple, cast, get_args from unittest.mock import ANY, MagicMock, Mock, patch @@ -4738,6 +4739,188 @@ def test_every_bridged_chunk_after_response_created_carries_the_served_service_t assert relayed == ["default"] * len(events), relayed +_AUDIO_PART: Final = {"type": "input_audio", "input_audio": {"data": "Zm9v", "format": "wav"}} +_TEXT_PART: Final = {"type": "text", "text": "Transcribe this"} + + +def test_convert_chat_completion_messages_to_responses_api_maps_input_audio_block(): + handler: Final = LiteLLMResponsesTransformationHandler() + messages: Final = [ + { + "role": "user", + "content": [_TEXT_PART, {**_AUDIO_PART, "prompt_cache_breakpoint": {"mode": "explicit"}}], + } + ] + + items, _ = handler.convert_chat_completion_messages_to_responses_api(messages, keep_prompt_cache_breakpoints=True) + + assert items[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"}, + { + "type": "input_audio", + "input_audio": {"data": "Zm9v", "format": "wav"}, + "prompt_cache_breakpoint": {"mode": "explicit"}, + }, + ] + + +def test_convert_chat_completion_messages_to_responses_api_drops_malformed_input_audio_breakpoint_under_drop_params(): + handler: Final = LiteLLMResponsesTransformationHandler() + messages: Final = [{"role": "user", "content": [{**_AUDIO_PART, "prompt_cache_breakpoint": ["explicit"]}]}] + + items, _ = handler.convert_chat_completion_messages_to_responses_api( + messages, drop_params=True, keep_prompt_cache_breakpoints=True + ) + + assert items[0]["content"] == [_AUDIO_PART] + + +@pytest.fixture +def registered_audio_models(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setitem( + litellm.model_cost, + "unit-audio-capable", + {"litellm_provider": "openai", "mode": "chat", "supports_audio_input": True}, + ) + monkeypatch.setitem( + litellm.model_cost, + "unit-text-only", + { + "litellm_provider": "openai", + "mode": "chat", + "supports_audio_input": False, + "supports_prompt_cache_breakpoint": True, + }, + ) + + +def _bridge_input( + model: str, + drop_params: bool | None, + messages: Sequence[Mapping[str, object]] | None = None, + **extra_litellm_params: object, +) -> list[dict[str, object]]: + chat_messages: Final = cast( + List[AllMessageValues], list(messages or [{"role": "user", "content": [_TEXT_PART, _AUDIO_PART]}]) + ) # cast-ok: the tests build chat messages as plain mappings + request: Final = LiteLLMResponsesTransformationHandler().transform_request( + model=model, + messages=chat_messages, + optional_params={}, + litellm_params={"custom_llm_provider": "openai", "drop_params": drop_params, **extra_litellm_params}, + headers={}, + litellm_logging_obj=Mock(), + ) + return cast(list[dict[str, object]], request["input"]) # cast-ok: the bridge emits message item mappings + + +def test_transform_request_forwards_input_audio_without_drop_params( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + + assert _bridge_input("unit-text-only", drop_params=False)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"}, + _AUDIO_PART, + ] + + +def test_transform_request_drops_input_audio_under_drop_params_when_model_lacks_audio_input( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + + assert _bridge_input("unit-text-only", drop_params=True)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"} + ] + + +def test_transform_request_keeps_input_audio_under_drop_params_when_model_supports_audio_input( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + + assert _bridge_input("unit-audio-capable", drop_params=True)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"}, + _AUDIO_PART, + ] + + +def test_transform_request_keeps_input_audio_under_drop_params_when_base_model_supports_audio_input( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + + assert _bridge_input("my-audio-deployment", drop_params=True, base_model="unit-audio-capable")[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"}, + _AUDIO_PART, + ] + + +def test_transform_request_drops_input_audio_under_global_drop_params( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", True) + + assert _bridge_input("unit-text-only", drop_params=None)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this"} + ] + + +def test_transform_request_drops_input_audio_from_tool_output_under_drop_params( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + messages: Final = [ + {"role": "user", "content": "Describe the recording"}, + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "record", "arguments": "{}"}}], + }, + {"role": "tool", "tool_call_id": "call_1", "content": [_TEXT_PART, _AUDIO_PART]}, + ] + + forwarded: Final = _bridge_input("unit-text-only", drop_params=False, messages=messages) + dropped: Final = _bridge_input("unit-text-only", drop_params=True, messages=messages) + + assert forwarded[-1]["type"] == "function_call_output" + assert forwarded[-1]["output"] == [{"type": "input_text", "text": "Transcribe this"}, _AUDIO_PART] + assert dropped[-1]["output"] == [{"type": "input_text", "text": "Transcribe this"}] + + +def test_transform_request_moves_the_dropped_audio_part_breakpoint_to_the_preceding_part( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + messages: Final = [ + {"role": "user", "content": [_TEXT_PART, {**_AUDIO_PART, "prompt_cache_breakpoint": {"mode": "explicit"}}]} + ] + + assert _bridge_input("unit-text-only", drop_params=True, messages=messages)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this", "prompt_cache_breakpoint": {"mode": "explicit"}} + ] + + +def test_transform_request_keeps_the_preceding_part_breakpoint_over_the_dropped_audio_part_breakpoint( + monkeypatch: pytest.MonkeyPatch, registered_audio_models: None +): + monkeypatch.setattr(litellm, "drop_params", False) + messages: Final = [ + { + "role": "user", + "content": [ + {**_TEXT_PART, "prompt_cache_breakpoint": {"mode": "explicit", "ttl": "30m"}}, + {**_AUDIO_PART, "prompt_cache_breakpoint": {"mode": "explicit"}}, + ], + } + ] + + assert _bridge_input("unit-text-only", drop_params=True, messages=messages)[0]["content"] == [ + {"type": "input_text", "text": "Transcribe this", "prompt_cache_breakpoint": {"mode": "explicit", "ttl": "30m"}} + ] + + def test_convert_chat_completion_messages_to_responses_api_keeps_prompt_cache_breakpoint_on_unknown_block(): """The hook marks the last block of its target message, so a message ending in a block the bridge cannot map reaches the stringify path and has to keep the marker there.""" @@ -4754,8 +4937,8 @@ def test_convert_chat_completion_messages_to_responses_api_keeps_prompt_cache_br "content": [ {"type": "text", "text": "describe this"}, { - "type": "input_audio", - "input_audio": {"data": "Zm9v", "format": "wav"}, + "type": "video_url", + "video_url": {"url": "https://example.com/clip.mp4"}, "prompt_cache_breakpoint": breakpoint_marker, }, ], @@ -4836,4 +5019,4 @@ def test_transform_request_drop_params_in_litellm_params_gates_the_prompt_cache_ litellm_logging_obj=Mock(), ) - assert "prompt_cache_breakpoint" not in result["input"][0]["content"][0] + assert "prompt_cache_breakpoint" not in result["input"][0]["content"][0] \ No newline at end of file diff --git a/tests/unit/embeddings/test_dispatch.py b/tests/unit/embeddings/test_dispatch.py index 88a2e7532c2..750c88cd235 100644 --- a/tests/unit/embeddings/test_dispatch.py +++ b/tests/unit/embeddings/test_dispatch.py @@ -33,9 +33,9 @@ def test_sync_embedding_request_projects_public_arguments() -> None: expected: Final = EmbeddingResponse(model="test-model", data=[]) def native(request: NativeCall) -> EmbeddingResponse: - assert request.bound["model"] == "test-model" - assert request.bound["input"] == "hello" - assert request.bound["custom_llm_provider"] == "openai" + assert request.resolved["model"] == "test-model" + assert request.resolved["input"] == "hello" + assert request.resolved["custom_llm_provider"] == "openai" return expected binding: Final[NativeBinding[Callable[[NativeCall], EmbeddingResponse]]] = NativeBinding( diff --git a/tests/unit/enterprise/proxy/hooks/test_managed_files.py b/tests/unit/enterprise/proxy/hooks/test_managed_files.py index 11d9dbec49c..ed80bd9b928 100644 --- a/tests/unit/enterprise/proxy/hooks/test_managed_files.py +++ b/tests/unit/enterprise/proxy/hooks/test_managed_files.py @@ -1,19 +1,268 @@ import base64 import json -from typing import cast +import logging +import time +import asyncio +from collections.abc import Awaitable, Callable, Mapping +from types import SimpleNamespace +from typing import TYPE_CHECKING, Final, cast from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from fastapi import HTTPException -from litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles +from openai import APIConnectionError +if TYPE_CHECKING: + from litellm.types.utils import LiteLLMBatch + +from litellm_enterprise.proxy.hooks.managed_files import ( + PROXY_LiteLLMManagedFiles, + _provider_file_retrieve_credentials, +) from litellm.caching import DualCache -from litellm.proxy._types import CallTypes +from litellm.proxy._types import CallTypes, LiteLLM_ManagedFileTable from litellm.proxy.openai_files_endpoints.common_utils import ( BATCH_CREATE_HIDDEN_PARAM, is_base64_encoded_unified_file_id, encode_file_id_with_model, ) +from litellm.types.llms.openai import OpenAIFileObject + + +class _InMemoryManagedFileTable: + def __init__(self, *rows: LiteLLM_ManagedFileTable) -> None: + self.rows: dict[str, LiteLLM_ManagedFileTable] = { + row.unified_file_id: row for row in rows + } + self.find_first_calls: list[Mapping[str, object]] = [] + self.upsert_calls: list[ + tuple[Mapping[str, object], Mapping[str, Mapping[str, object]]] + ] = [] + self.update_many_calls: list[ + tuple[Mapping[str, object], Mapping[str, object]] + ] = [] + + async def find_first( + self, where: Mapping[str, object] + ) -> LiteLLM_ManagedFileTable | None: + self.find_first_calls.append(where) + unified_file_id: Final = where.get("unified_file_id") + if isinstance(unified_file_id, str): + return self.rows.get(unified_file_id) + raw_file_filter: Final = where.get("flat_model_file_ids") + if isinstance(raw_file_filter, Mapping): + raw_file_id: Final = raw_file_filter.get("has") + if isinstance(raw_file_id, str): + return next( + (row for row in self.rows.values() if raw_file_id in row.flat_model_file_ids), + None, + ) + return None + + async def upsert( + self, + where: Mapping[str, object], + data: Mapping[str, Mapping[str, object]], + ) -> LiteLLM_ManagedFileTable: + self.upsert_calls.append((where, data)) + unified_file_id: Final = cast(str, where["unified_file_id"]) + previous_row: Final = self.rows.get(unified_file_id) + values: Final = { + **( + cast(dict[str, object], previous_row.model_dump()) + if previous_row is not None + else {} + ), + **data["update" if previous_row is not None else "create"], + } + raw_model_mappings: Final = values.get("model_mappings", "{}") + model_mappings: Final = ( + cast(dict[str, str], json.loads(raw_model_mappings)) + if isinstance(raw_model_mappings, str) + else cast(dict[str, str], raw_model_mappings) + ) + raw_file_object: Final = values.get("file_object") + file_object: Final = cast( + dict[str, object] | None, + json.loads(raw_file_object) + if isinstance(raw_file_object, str) + else raw_file_object, + ) + row: Final = LiteLLM_ManagedFileTable.model_validate( + { + **values, + "file_object": file_object, + "model_mappings": model_mappings, + } + ) + self.rows[unified_file_id] = row + return row + + async def update_many( + self, + where: Mapping[str, object], + data: Mapping[str, object], + ) -> int: + self.update_many_calls.append((where, data)) + unified_file_id: Final = cast(str, where["unified_file_id"]) + existing_row: Final = self.rows.get(unified_file_id) + if existing_row is None: + return 0 + file_object: Final = OpenAIFileObject.model_validate( + json.loads(cast(str, data["file_object"])) + ) + self.rows[unified_file_id] = existing_row.model_copy( + update={"file_object": file_object} + ) + return 1 + + async def delete( + self, where: Mapping[str, object] + ) -> LiteLLM_ManagedFileTable | None: + return self.rows.pop(cast(str, where["unified_file_id"]), None) + + async def find_many( + self, + where: Mapping[str, object], + **_query: object, + ) -> list[LiteLLM_ManagedFileTable]: + return list(self.rows.values()) + + +class _InMemoryManagedObjectTable: + async def find_first(self, where: Mapping[str, object]) -> None: + return None + + async def update_many( + self, where: Mapping[str, object], data: Mapping[str, object] + ) -> int: + return 0 + + async def upsert( + self, + where: Mapping[str, object], + data: Mapping[str, object], + ) -> None: + return None + + +class _InMemoryManagedFilesDatabase: + def __init__(self, managed_file_table: _InMemoryManagedFileTable) -> None: + self.litellm_managedfiletable: Final = managed_file_table + self.litellm_managedobjecttable: Final = _InMemoryManagedObjectTable() + + +class _InMemoryManagedFilesPrismaClient: + def __init__(self, managed_file_table: _InMemoryManagedFileTable) -> None: + self.db: Final = _InMemoryManagedFilesDatabase(managed_file_table) + + +def _managed_files_with_fake_prisma( + *rows: LiteLLM_ManagedFileTable, + sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, +) -> tuple[PROXY_LiteLLMManagedFiles, _InMemoryManagedFileTable]: + managed_file_table: Final = _InMemoryManagedFileTable(*rows) + prisma_client: Final = _InMemoryManagedFilesPrismaClient(managed_file_table) + managed_files: Final = PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=prisma_client, sleep=sleep + ) + return managed_files, managed_file_table + + +def _marked_fallback_file_row( + *, + size_bytes: int = 0, + created_at: int = 123, +) -> LiteLLM_ManagedFileTable: + from litellm.types.llms.openai import OpenAIFileObject + + file_object: Final = OpenAIFileObject( + id="unified-output", + object="file", + bytes=size_bytes, + created_at=created_at, + filename="output.jsonl", + purpose="batch_output", + status="processed", + litellm_details_fallback=True, + ) + return LiteLLM_ManagedFileTable( + unified_file_id="unified-output", + file_object=file_object, + model_mappings={"model-123": "provider-output"}, + flat_model_file_ids=["provider-output"], + created_by="user-123", + team_id="team-123", + ) + + +def _batch_response_for_output_listing() -> "LiteLLMBatch": + from openai.types.batch import BatchRequestCounts + from litellm.types.utils import LiteLLMBatch + + response: Final = LiteLLMBatch( + id="batch-123", + completion_window="24h", + created_at=1, + endpoint="/v1/chat/completions", + input_file_id="input-file", + object="batch", + status="completed", + output_file_id="provider-output", + request_counts=BatchRequestCounts(completed=1, failed=0, total=1), + ) + response.hidden_params = { + "model_id": "model-123", + "model_name": "bedrock/model-x", + } + return response + + +async def _resolve_batch_for_output_listing( + managed_files: PROXY_LiteLLMManagedFiles, + response: "LiteLLMBatch", +) -> "LiteLLMBatch | None": + from litellm.proxy._types import UserAPIKeyAuth + + return await managed_files._resolve_listed_batch( + row=SimpleNamespace( + unified_object_id="batch-row", + created_by="user-123", + team_id="team-123", + ), + batch_obj=response, + unified_id_by_raw_id={}, + user_api_key_dict=UserAPIKeyAuth(user_id="user-123", team_id="team-123"), + ) + + +def test_provider_credentials_warning_sanitizes_newlines( + caplog: pytest.LogCaptureFixture, +) -> None: + model_id: Final = "deployment-123\nforged warning" + router: Final = MagicMock() + router.get_deployment_credentials_with_provider.side_effect = RuntimeError( + "credential lookup failed\nforged warning" + ) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + credentials: Final = _provider_file_retrieve_credentials( + llm_router=router, + model_id=model_id, + ) + + warnings: Final = tuple( + record.getMessage() + for record in caplog.records + if "Failed to retrieve credentials for provider file" in record.getMessage() + ) + assert credentials is None + assert warnings == ( + "Failed to retrieve credentials for provider file " + "model_id=deployment-123forged warning: credential lookup failedforged warning", + ) + assert all("\n" not in warning and "\r" not in warning for warning in warnings) def test_get_file_ids_from_messages(): @@ -480,13 +729,11 @@ async def test_router_acreate_batch_only_selects_from_file_id_mapping(monkeypatc @pytest.mark.asyncio async def test_output_file_id_for_batch_retrieve(): - """ - Test that the output file id is the same as the input file id - """ - from typing import cast - + import litellm.proxy.proxy_server as proxy_server_module from openai.types.batch import BatchRequestCounts + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.llms.openai import OpenAIFileObject from litellm.types.utils import LiteLLMBatch batch = LiteLLMBatch( @@ -522,18 +769,42 @@ async def test_output_file_id_for_batch_retrieve(): "litellm_model_name": "gpt-5.5", "unified_batch_id": "litellm_proxy;model_id:12345679;llm_batch_id:batch_685c5e5d63988190b85bdb2147ba131d", } - proxy_managed_files = PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=AsyncMock() + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + provider_file_object = OpenAIFileObject( + id="file-provider-output", + object="file", + bytes=123, + created_at=456, + filename="output.jsonl", + purpose="batch_output", ) + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma() - response = await proxy_managed_files.async_post_call_success_hook( - data={}, - user_api_key_dict=MagicMock(), - response=batch, + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + return_value=provider_file_object, + ), + ): + response = await proxy_managed_files.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), + response=batch, + ) + + output_file_id = cast(str, cast(LiteLLMBatch, response).output_file_id) + assert not output_file_id.startswith("file-") + stored_file_object = managed_file_table.rows[output_file_id].file_object + assert stored_file_object is not None + assert stored_file_object.id == output_file_id + assert ( + managed_file_table.upsert_calls[0][1]["create"]["file_object"] + == stored_file_object.model_dump_json() ) - assert not cast(LiteLLMBatch, response).output_file_id.startswith("file-") - @pytest.mark.asyncio async def test_output_file_id_preserves_target_model_names_when_model_name_missing(): @@ -542,6 +813,7 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi (e.g. Vertex batch retrieve), unified output_file_id should still include target_model_names from the managed input file ID. """ + import litellm.proxy.proxy_server as proxy_server_module from openai.types.batch import BatchRequestCounts from litellm.proxy._types import UserAPIKeyAuth @@ -582,10 +854,6 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi # Intentionally omit model_name to mimic Vertex issue. } - proxy_managed_files = PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=AsyncMock() - ) - provider_output_file = OpenAIFileObject( id="file-provider-output-id", object="file", @@ -594,9 +862,18 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi filename="predictions.jsonl", purpose="batch_output", ) + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma() - with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_retrieve: - mock_retrieve.return_value = provider_output_file + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + return_value=provider_output_file, + ) as mock_retrieve, + ): response = await proxy_managed_files.async_post_call_success_hook( data={}, user_api_key_dict=UserAPIKeyAuth(user_id="test-user"), @@ -608,6 +885,15 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi ) assert decoded_output_file_id assert "target_model_names,gemini-2.5-pro" in cast(str, decoded_output_file_id) + mock_retrieve.assert_awaited_once() + output_file_id = cast(str, cast(LiteLLMBatch, response).output_file_id) + stored_file_object = managed_file_table.rows[output_file_id].file_object + assert stored_file_object is not None + assert stored_file_object.id == output_file_id + assert ( + managed_file_table.upsert_calls[0][1]["create"]["file_object"] + == stored_file_object.model_dump_json() + ) @pytest.mark.asyncio @@ -615,8 +901,7 @@ async def test_error_file_id_for_failed_batch(): """ Test that the error_file_id is properly managed when a batch fails """ - from typing import cast - + import litellm.proxy.proxy_server as proxy_server_module from openai.types.batch import BatchRequestCounts from litellm.proxy._types import UserAPIKeyAuth @@ -658,9 +943,7 @@ async def test_error_file_id_for_failed_batch(): "unified_batch_id": "litellm_proxy;model_id:test-model-id;llm_batch_id:batch_abc123", } - proxy_managed_files = PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=AsyncMock() - ) + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma() # Create a proper OpenAIFileObject for the error file error_file_object = OpenAIFileObject( @@ -672,14 +955,20 @@ async def test_error_file_id_for_failed_batch(): purpose="batch_output", status="processed", ) + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} # Mock the afile_retrieve to simulate retrieving error file metadata - with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_retrieve: - mock_retrieve.return_value = error_file_object + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + return_value=error_file_object, + ) as mock_retrieve, + ): - user_api_key_dict = UserAPIKeyAuth( - user_id="test-user-123", parent_otel_span=MagicMock() - ) + user_api_key_dict = UserAPIKeyAuth(user_id="test-user-123") response = await proxy_managed_files.async_post_call_success_hook( data={}, @@ -691,8 +980,15 @@ async def test_error_file_id_for_failed_batch(): assert cast(LiteLLMBatch, response).error_file_id is not None assert not cast(LiteLLMBatch, response).error_file_id.startswith("error-") # Verify it's a base64 encoded managed file ID - assert is_base64_encoded_unified_file_id( - cast(LiteLLMBatch, response).error_file_id + error_file_id = cast(str, cast(LiteLLMBatch, response).error_file_id) + assert is_base64_encoded_unified_file_id(error_file_id) + mock_retrieve.assert_awaited_once() + stored_file_object = managed_file_table.rows[error_file_id].file_object + assert stored_file_object is not None + assert stored_file_object.id == error_file_id + assert ( + managed_file_table.upsert_calls[0][1]["create"]["file_object"] + == stored_file_object.model_dump_json() ) @@ -1514,6 +1810,891 @@ async def test_store_unified_file_id_updates_file_metadata_on_existing_row(): assert second_update["storage_url"] == "s3://bucket/output.jsonl" +@pytest.mark.asyncio +async def test_store_batch_output_file_skips_existing_file_object(): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.llms.openai import OpenAIFileObject + + existing_file_object = OpenAIFileObject( + id="unified-output", + object="file", + bytes=1, + created_at=1, + filename="output.jsonl", + purpose="batch_output", + ) + existing_file_row = LiteLLM_ManagedFileTable( + unified_file_id="unified-output", + file_object=existing_file_object, + model_mappings={"model-123": "provider-output"}, + flat_model_file_ids=["provider-output"], + created_by="user-123", + ) + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + existing_file_row + ) + + with patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve: + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + ) + + retrieve.assert_not_called() + assert managed_file_table.find_first_calls == [ + {"unified_file_id": "unified-output"} + ] + assert managed_file_table.upsert_calls == [] + + +@pytest.mark.asyncio +async def test_store_batch_output_file_stores_provider_object_with_unified_id(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.llms.openai import OpenAIFileObject + + provider_file_object = OpenAIFileObject( + id="provider-output", + object="file", + bytes=123, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + ) + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + sleep_calls: list[float] = [] + + async def record_sleep(delay: float) -> None: + sleep_calls.append(delay) + + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + sleep=record_sleep + ) + + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=[ + HTTPException(status_code=500), + HTTPException(status_code=500), + provider_file_object, + ], + ) as retrieve, + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + ) + + assert retrieve.await_count == 3 + assert [call.kwargs["max_retries"] for call in retrieve.call_args_list] == [0, 0, 0] + assert sleep_calls == [0.5, 1.0] + stored_object = managed_file_table.rows["unified-output"].file_object + assert stored_object is not None + assert stored_object.id == "unified-output" + assert stored_object.bytes == 123 + assert ( + managed_file_table.upsert_calls[0][1]["create"]["file_object"] + == stored_object.model_dump_json() + ) + assert managed_file_table.rows["unified-output"].model_mappings == { + "model-123": "provider-output" + } + + +@pytest.mark.asyncio +async def test_store_batch_output_file_retries_model_name_provider_fetch(): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.llms.openai import OpenAIFileObject + + provider_file_object = OpenAIFileObject( + id="provider-output", + object="file", + bytes=123, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + ) + sleep_calls: list[float] = [] + + async def record_sleep(delay: float) -> None: + sleep_calls.append(delay) + + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + sleep=record_sleep + ) + with patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=[ + HTTPException(status_code=500), + HTTPException(status_code=500), + provider_file_object, + ], + ) as retrieve: + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id=None, + model_name="bedrock/model-x", + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + ) + + assert retrieve.await_count == 3 + assert all( + call.kwargs["custom_llm_provider"] == "bedrock" + and call.kwargs["max_retries"] == 0 + and "_litellm_internal_model_credentials" not in call.kwargs + for call in retrieve.call_args_list + ) + assert sleep_calls == [0.5, 1.0] + stored_object = managed_file_table.rows["unified-output"].file_object + assert stored_object is not None + assert stored_object.id == "unified-output" + assert stored_object.bytes == 123 + assert stored_object.litellm_details_fallback is None + + +@pytest.mark.asyncio +async def test_store_batch_output_file_marks_model_name_fallback_after_four_transient_failures(): + from litellm.proxy._types import UserAPIKeyAuth + + sleep_calls: list[float] = [] + + async def record_sleep(delay: float) -> None: + sleep_calls.append(delay) + + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + sleep=record_sleep + ) + with patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=[HTTPException(status_code=500) for _ in range(4)], + ) as retrieve: + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id=None, + model_name="bedrock/model-x", + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + size_bytes=836, + ) + + assert retrieve.await_count == 4 + assert all( + call.kwargs["custom_llm_provider"] == "bedrock" + and call.kwargs["max_retries"] == 0 + and "_litellm_internal_model_credentials" not in call.kwargs + for call in retrieve.call_args_list + ) + assert sleep_calls == [0.5, 1.0, 2.0] + stored_object = managed_file_table.rows["unified-output"].file_object + assert stored_object is not None + assert stored_object.bytes == 836 + assert stored_object.litellm_details_fallback is True + + +@pytest.mark.asyncio +async def test_store_batch_output_file_falls_back_when_provider_retrieve_raises(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma() + + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=RuntimeError("provider unavailable"), + ) as retrieve, + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="s3://bucket/output.jsonl", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + size_bytes=321, + ) + + retrieve.assert_awaited_once() + stored_object = managed_file_table.rows["unified-output"].file_object + assert stored_object is not None + assert stored_object.id == "unified-output" + assert stored_object.filename == "output.jsonl" + assert stored_object.bytes == 321 + assert stored_object.purpose == "batch_output" + assert ( + managed_file_table.upsert_calls[0][1]["create"]["file_object"] + == stored_object.model_dump_json() + ) + assert managed_file_table.rows["unified-output"].model_mappings == { + "model-123": "s3://bucket/output.jsonl" + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model_id", "model_name"), + [("model-123", None), (None, "openai/gpt-4o-mini")], + ids=["deployment-credentials", "provider-from-model-name"], +) +async def test_store_batch_output_file_falls_back_when_provider_retrieve_hangs( + model_id: str | None, model_name: str | None +): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + async def never_answers(**_: object) -> None: + await asyncio.Event().wait() + + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma() + + with ( + patch.object(proxy_server_module, "llm_router", router), + patch("litellm.afile_retrieve", side_effect=never_answers), + patch("litellm_enterprise.proxy.hooks.managed_files.BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS", 0.01), + ): + await asyncio.wait_for( + proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="s3://bucket/output.jsonl", + model_id=model_id, + model_name=model_name, + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + size_bytes=321, + ), + timeout=5, + ) + + stored_object = managed_file_table.rows["unified-output"].file_object + assert stored_object is not None + assert (stored_object.id, stored_object.bytes, stored_object.purpose) == ("unified-output", 321, "batch_output") + assert stored_object.litellm_details_fallback is True + + +@pytest.mark.asyncio +async def test_afile_retrieve_returns_marked_fallback_when_refresh_hangs(): + async def never_answers(**_: object) -> None: + await asyncio.Event().wait() + + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(_marked_fallback_file_row()) + + with ( + patch("litellm.afile_retrieve", side_effect=never_answers), + patch("litellm_enterprise.proxy.hooks.managed_files.BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS", 0.01), + ): + response = await asyncio.wait_for( + proxy_managed_files.afile_retrieve( + file_id="unified-output", + litellm_parent_otel_span=None, + llm_router=router, + ), + timeout=5, + ) + + assert (response.id, response.bytes) == ("unified-output", 0) + assert managed_file_table.upsert_calls == [] + + +@pytest.mark.asyncio +async def test_store_batch_output_file_falls_back_without_model_id(): + from litellm.proxy._types import UserAPIKeyAuth + + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma() + + with patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve: + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-error", + provider_file_id="error.jsonl", + model_id=None, + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + ) + + retrieve.assert_not_called() + stored_object = managed_file_table.rows["unified-error"].file_object + assert stored_object is not None + assert stored_object.id == "unified-error" + assert stored_object.filename == "error.jsonl" + assert stored_object.bytes == 0 + assert stored_object.purpose == "batch_output" + assert ( + managed_file_table.upsert_calls[0][1]["create"]["file_object"] + == stored_object.model_dump_json() + ) + assert managed_file_table.rows["unified-error"].model_mappings == {} + assert stored_object.litellm_details_fallback is None + assert "litellm_details_fallback" not in stored_object.model_dump() + + +@pytest.mark.asyncio +async def test_listed_batch_saves_basic_output_without_provider_fetch(): + import litellm.proxy.proxy_server as proxy_server_module + + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma() + response: Final = _batch_response_for_output_listing() + with ( + patch.object(proxy_server_module, "llm_router", None), + patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve, + ): + resolved: Final = await _resolve_batch_for_output_listing( + proxy_managed_files, response + ) + + assert resolved is response + retrieve.assert_not_called() + stored_file: Final = next(iter(managed_file_table.rows.values())) + assert stored_file.file_object is not None + assert stored_file.file_object.bytes == 0 + assert stored_file.file_object.litellm_details_fallback is True + assert stored_file.model_mappings == {"model-123": "provider-output"} + + +@pytest.mark.asyncio +async def test_listed_batch_leaves_existing_marked_output_untouched(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + proxy_managed_files, _ = _managed_files_with_fake_prisma() + unified_file_id: Final = proxy_managed_files.get_unified_output_file_id( + output_file_id="provider-output", + model_id="model-123", + model_name="bedrock/model-x", + ) + row: Final = _marked_fallback_file_row().model_copy( + update={"unified_file_id": unified_file_id} + ) + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(row) + response: Final = _batch_response_for_output_listing() + with ( + patch.object(proxy_server_module, "llm_router", None), + patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve, + ): + resolved: Final = await _resolve_batch_for_output_listing( + proxy_managed_files, response + ) + await proxy_managed_files.store_batch_output_file( + unified_file_id=unified_file_id, + provider_file_id="provider-output", + model_id="model-123", + model_name="bedrock/model-x", + owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"), + litellm_parent_otel_span=None, + fetch_provider_details=False, + ) + + assert resolved is response + retrieve.assert_not_called() + assert managed_file_table.rows[unified_file_id] is row + assert managed_file_table.upsert_calls == [] + + +@pytest.mark.asyncio +async def test_batch_retrieve_registration_fetches_details_by_default(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.openai_files_endpoints.common_utils import ( + ensure_batch_response_managed_file_ids, + ) + + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma() + response: Final = _batch_response_for_output_listing() + router: Final = MagicMock() + router.get_deployment_credentials_with_provider.return_value = { + "api_key": "key", + "custom_llm_provider": "bedrock", + } + provider_file_object: Final = OpenAIFileObject( + id="provider-output", + object="file", + bytes=836, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + status="processed", + ) + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + return_value=provider_file_object, + ) as retrieve, + ): + await ensure_batch_response_managed_file_ids( + response=response, + managed_files_obj=proxy_managed_files, + prisma_client=proxy_managed_files.prisma_client, + verbose_proxy_logger=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="user-123"), + db_batch_object=SimpleNamespace(created_by="user-123", team_id=None), + ) + + retrieve.assert_awaited_once() + output_file_id: Final = response.output_file_id + assert output_file_id is not None + stored_file: Final = managed_file_table.rows[output_file_id] + assert stored_file.file_object is not None + assert stored_file.file_object.bytes == 836 + + +@pytest.mark.parametrize( + ("error_factory", "expected_attempts"), + [ + pytest.param(lambda: HTTPException(status_code=429), 2, id="429"), + pytest.param(lambda: HTTPException(status_code=408), 2, id="408"), + pytest.param(lambda: HTTPException(status_code=503), 2, id="503"), + pytest.param( + lambda: httpx.ConnectError( + "connection failed", + request=httpx.Request("GET", "https://api.openai.com/v1/files/file-1"), + ), + 2, + id="connect-error", + ), + pytest.param( + lambda: APIConnectionError( + message="connection failed", + request=httpx.Request("GET", "https://api.openai.com/v1/files/file-1"), + ), + 2, + id="openai-api-connection-error", + ), + pytest.param(lambda: asyncio.TimeoutError(), 2, id="async-timeout"), + pytest.param(lambda: HTTPException(status_code=400), 1, id="400"), + pytest.param(lambda: HTTPException(status_code=403), 1, id="403"), + pytest.param(lambda: HTTPException(status_code=404), 1, id="404"), + pytest.param(lambda: ValueError("invalid file"), 1, id="value-error"), + ], +) +@pytest.mark.asyncio +async def test_batch_file_detail_retry_classification( + error_factory: Callable[[], Exception], + expected_attempts: int, +) -> None: + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + sleep_calls: Final[list[float]] = [] + + async def record_sleep(delay: float) -> None: + sleep_calls.append(delay) + + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + sleep=record_sleep + ) + router: Final = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + provider_file_object: Final = OpenAIFileObject( + id="provider-output", + object="file", + bytes=836, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + status="processed", + ) + retrieve_side_effect: Final = ( + [error_factory(), provider_file_object] + if expected_attempts == 2 + else error_factory() + ) + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=retrieve_side_effect, + ) as retrieve, + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + ) + + assert retrieve.await_count == expected_attempts + assert [call.kwargs["max_retries"] for call in retrieve.call_args_list] == [ + 0 + ] * expected_attempts + assert sleep_calls == ([0.5] if expected_attempts == 2 else []) + stored_object: Final = managed_file_table.rows["unified-output"].file_object + assert stored_object is not None + assert stored_object.bytes == (836 if expected_attempts == 2 else 0) + + +@pytest.mark.asyncio +async def test_store_batch_output_file_ignores_unmarked_row_with_provider_route(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.llms.openai import OpenAIFileObject + + existing_file_object: Final = OpenAIFileObject( + id="unified-output", + object="file", + bytes=1, + created_at=1, + filename="output.jsonl", + purpose="batch_output", + ) + existing_file_row: Final = LiteLLM_ManagedFileTable( + unified_file_id="unified-output", + file_object=existing_file_object, + model_mappings={"model-123": "provider-output"}, + flat_model_file_ids=["provider-output"], + created_by="user-123", + ) + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + existing_file_row + ) + router: Final = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + with ( + patch.object(proxy_server_module, "llm_router", router), + patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve, + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + ) + + retrieve.assert_not_called() + assert managed_file_table.rows["unified-output"].file_object is existing_file_object + assert managed_file_table.upsert_calls == [] + + +@pytest.mark.asyncio +async def test_store_batch_output_file_rejects_non_provider_model_prefix(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma() + with ( + patch.object(proxy_server_module, "llm_router", None), + patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve, + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + model_name="my-team/gpt-4o", + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + ) + + retrieve.assert_not_called() + stored_object: Final = managed_file_table.rows["unified-output"].file_object + assert stored_object is not None + assert stored_object.litellm_details_fallback is None + + +@pytest.mark.asyncio +async def test_store_batch_output_file_marks_fallback_after_four_transient_failures(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + sleep_calls: list[float] = [] + + async def record_sleep(delay: float) -> None: + sleep_calls.append(delay) + + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + sleep=record_sleep + ) + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=[HTTPException(status_code=500) for _ in range(4)], + ) as retrieve, + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"), + litellm_parent_otel_span=None, + size_bytes=836, + ) + + assert retrieve.await_count == 4 + assert [call.kwargs["max_retries"] for call in retrieve.call_args_list] == [0, 0, 0, 0] + assert sleep_calls == [0.5, 1.0, 2.0] + stored_object = managed_file_table.rows["unified-output"].file_object + assert stored_object is not None + assert stored_object.litellm_details_fallback is True + stored_file_object_json: Final = managed_file_table.upsert_calls[0][1]["create"]["file_object"] + assert isinstance(stored_file_object_json, str) + assert json.loads(stored_file_object_json)["litellm_details_fallback"] is True + + +@pytest.mark.asyncio +async def test_store_batch_output_file_skips_lookup_for_fallback_written_moments_ago(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + _marked_fallback_file_row(created_at=int(time.time())) + ) + + with ( + patch.object(proxy_server_module, "llm_router", router), + patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve, + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"), + litellm_parent_otel_span=None, + size_bytes=836, + ) + + retrieve.assert_not_awaited() + stored_object = managed_file_table.rows["unified-output"].file_object + assert stored_object is not None + assert (stored_object.bytes, stored_object.litellm_details_fallback) == (836, True) + assert managed_file_table.upsert_calls == [] + assert managed_file_table.update_many_calls[0][0] == { + "unified_file_id": "unified-output" + } + assert set(managed_file_table.update_many_calls[0][1]) == {"file_object"} + + +@pytest.mark.asyncio +async def test_store_batch_output_file_does_not_retry_not_found(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + sleep_calls: list[float] = [] + + async def record_sleep(delay: float) -> None: + sleep_calls.append(delay) + + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + sleep=record_sleep + ) + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=HTTPException(status_code=404), + ) as retrieve, + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + ) + + retrieve.assert_awaited_once() + assert sleep_calls == [] + stored_object = managed_file_table.rows["unified-output"].file_object + assert stored_object is not None + assert stored_object.litellm_details_fallback is True + + +@pytest.mark.asyncio +async def test_store_batch_output_file_refreshes_marked_fallback_on_later_write(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.llms.openai import OpenAIFileObject + + row = _marked_fallback_file_row() + provider_object = OpenAIFileObject( + id="provider-output", + object="file", + bytes=836, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + status="processed", + ) + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(row) + with ( + patch.object(proxy_server_module, "llm_router", router), + patch("litellm.afile_retrieve", new_callable=AsyncMock, return_value=provider_object) as retrieve, + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"), + litellm_parent_otel_span=None, + ) + + retrieve.assert_awaited_once() + refreshed = managed_file_table.rows["unified-output"].file_object + assert refreshed is not None + assert refreshed.id == "unified-output" + assert refreshed.bytes == 836 + assert refreshed.litellm_details_fallback is None + assert managed_file_table.update_many_calls == [ + ({"unified_file_id": "unified-output"}, {"file_object": refreshed.model_dump_json()}) + ] + assert managed_file_table.upsert_calls == [] + + +@pytest.mark.asyncio +async def test_store_batch_output_file_does_not_recreate_deleted_row_after_refresh(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + provider_object: Final = OpenAIFileObject( + id="provider-output", + object="file", + bytes=836, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + status="processed", + ) + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + _marked_fallback_file_row() + ) + router: Final = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + + async def delete_then_return_file(**_kwargs: object) -> OpenAIFileObject: + await managed_file_table.delete(where={"unified_file_id": "unified-output"}) + return provider_object + + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=delete_then_return_file, + ), + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123"), + litellm_parent_otel_span=None, + ) + + assert "unified-output" not in managed_file_table.rows + assert managed_file_table.upsert_calls == [] + assert len(managed_file_table.update_many_calls) == 1 + saved_file_object_json: Final = cast( + str, managed_file_table.update_many_calls[0][1]["file_object"] + ) + saved_file_object: Final = json.loads(saved_file_object_json) + assert saved_file_object["id"] == "unified-output" + assert ( + await proxy_managed_files.internal_usage_cache.async_get_cache( + key="unified-output", + litellm_parent_otel_span=None, + ) + is None + ) + + +@pytest.mark.asyncio +async def test_store_batch_output_file_keeps_marked_fallback_when_retryable_details_still_fail(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + _marked_fallback_file_row() + ) + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=HTTPException(status_code=404), + ) as retrieve, + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"), + litellm_parent_otel_span=None, + ) + + retrieve.assert_awaited_once() + assert managed_file_table.upsert_calls == [] + assert managed_file_table.rows["unified-output"].file_object.created_at == 123 + + +@pytest.mark.asyncio +async def test_store_batch_output_file_updates_only_size_for_failed_marked_fallback(): + import litellm.proxy.proxy_server as proxy_server_module + from litellm.proxy._types import UserAPIKeyAuth + + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + _marked_fallback_file_row() + ) + with ( + patch.object(proxy_server_module, "llm_router", router), + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=HTTPException(status_code=404), + ), + ): + await proxy_managed_files.store_batch_output_file( + unified_file_id="unified-output", + provider_file_id="provider-output", + model_id="model-123", + owner=UserAPIKeyAuth(user_id="user-123", team_id="team-123"), + litellm_parent_otel_span=None, + size_bytes=836, + ) + + updated = managed_file_table.rows["unified-output"].file_object + assert updated is not None + assert updated.bytes == 836 + assert updated.created_at == 123 + assert updated.litellm_details_fallback is True + + @pytest.mark.asyncio async def test_afile_delete_returns_provider_response_when_stored_file_object_none(): """ @@ -1585,19 +2766,21 @@ async def test_afile_retrieve_fetches_from_provider_when_file_object_none(): """ from litellm.types.llms.openai import OpenAIFileObject - prisma_client = AsyncMock() - internal_usage_cache = MagicMock() - - proxy_managed_files = PROXY_LiteLLMManagedFiles( - internal_usage_cache=internal_usage_cache, - prisma_client=prisma_client, + stored_file = LiteLLM_ManagedFileTable( + unified_file_id="test-unified-file-id", + file_object=None, + model_mappings={"model-123": "file-provider-xyz"}, + flat_model_file_ids=["file-provider-xyz"], + created_by="user-123", ) + sleep_calls: list[float] = [] - # Mock get_unified_file_id to return a stored object with file_object=None - stored_file = MagicMock() - stored_file.file_object = None - stored_file.model_mappings = {"model-123": "file-provider-xyz"} - proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file) + async def record_sleep(delay: float) -> None: + sleep_calls.append(delay) + + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + stored_file, sleep=record_sleep + ) # Mock the router and provider response provider_file_response = OpenAIFileObject( @@ -1617,7 +2800,11 @@ async def test_afile_retrieve_fetches_from_provider_when_file_object_none(): } ) - with patch("litellm.afile_retrieve", new_callable=AsyncMock) as mock_afile_retrieve: + with patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=[HTTPException(status_code=500), provider_file_response], + ) as mock_afile_retrieve: mock_afile_retrieve.return_value = provider_file_response unified_file_id = "test-unified-file-id" @@ -1630,7 +2817,10 @@ async def test_afile_retrieve_fetches_from_provider_when_file_object_none(): # Should return the provider response with the unified file ID assert result is not None assert result.id == unified_file_id - mock_afile_retrieve.assert_called_once() + assert mock_afile_retrieve.await_count == 2 + assert [call.kwargs["max_retries"] for call in mock_afile_retrieve.call_args_list] == [0, 0] + assert sleep_calls == [0.5] + assert managed_file_table.upsert_calls == [] @pytest.mark.asyncio @@ -1639,19 +2829,14 @@ async def test_afile_retrieve_raises_error_when_no_router_and_file_object_none() Test that afile_retrieve raises an appropriate error when file_object is None and no llm_router is provided to fetch from the provider. """ - prisma_client = AsyncMock() - internal_usage_cache = MagicMock() - - proxy_managed_files = PROXY_LiteLLMManagedFiles( - internal_usage_cache=internal_usage_cache, - prisma_client=prisma_client, + stored_file = LiteLLM_ManagedFileTable( + unified_file_id="test-unified-file-id", + file_object=None, + model_mappings={"model-123": "file-provider-xyz"}, + flat_model_file_ids=["file-provider-xyz"], + created_by="user-123", ) - - # Mock get_unified_file_id to return a stored object with file_object=None - stored_file = MagicMock() - stored_file.file_object = None - stored_file.model_mappings = {"model-123": "file-provider-xyz"} - proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file) + proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file) unified_file_id = "test-unified-file-id" @@ -1662,7 +2847,7 @@ async def test_afile_retrieve_raises_error_when_no_router_and_file_object_none() llm_router=None, ) - assert "llm_router is required" in str(exc_info.value) + assert "no provider route to fetch it" in str(exc_info.value) @pytest.mark.asyncio @@ -1673,15 +2858,6 @@ async def test_afile_retrieve_returns_stored_file_object_when_exists(): """ from litellm.types.llms.openai import OpenAIFileObject - prisma_client = AsyncMock() - internal_usage_cache = MagicMock() - - proxy_managed_files = PROXY_LiteLLMManagedFiles( - internal_usage_cache=internal_usage_cache, - prisma_client=prisma_client, - ) - - # Mock get_unified_file_id to return a stored object WITH file_object stored_file_object = OpenAIFileObject( id="test-unified-file-id", object="file", @@ -1690,18 +2866,371 @@ async def test_afile_retrieve_returns_stored_file_object_when_exists(): filename="input.jsonl", purpose="batch", ) - stored_file = MagicMock() - stored_file.file_object = stored_file_object - proxy_managed_files.get_unified_file_id = AsyncMock(return_value=stored_file) + stored_file = LiteLLM_ManagedFileTable( + unified_file_id="test-unified-file-id", + file_object=stored_file_object, + model_mappings={"model-123": "provider-file-id"}, + flat_model_file_ids=["provider-file-id"], + created_by="user-123", + ) + proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file) + router = MagicMock() - result = await proxy_managed_files.afile_retrieve( - file_id="test-unified-file-id", + with patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve: + result = await proxy_managed_files.afile_retrieve( + file_id="test-unified-file-id", + litellm_parent_otel_span=None, + llm_router=router, + ) + + assert result == stored_file_object + retrieve.assert_not_called() + + +@pytest.mark.asyncio +async def test_afile_retrieve_refreshes_marked_fallback_and_preserves_ownership(): + from litellm.types.llms.openai import OpenAIFileObject + + row = _marked_fallback_file_row() + provider_object = OpenAIFileObject( + id="provider-output", + object="file", + bytes=836, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + status="processed", + ) + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(row) + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + + with patch("litellm.afile_retrieve", new_callable=AsyncMock, return_value=provider_object) as retrieve: + response = await proxy_managed_files.afile_retrieve( + file_id="unified-output", + litellm_parent_otel_span=None, + llm_router=router, + ) + + assert response == provider_object.model_copy(update={"id": "unified-output"}) + assert "litellm_details_fallback" not in response.model_dump() + retrieve.assert_awaited_once() + updated_row = managed_file_table.rows["unified-output"] + assert updated_row.file_object == response + assert updated_row.created_by == "user-123" + assert updated_row.team_id == "team-123" + assert managed_file_table.update_many_calls[0][1] == { + "file_object": response.model_dump_json() + } + assert managed_file_table.upsert_calls == [] + cached_row = await proxy_managed_files.internal_usage_cache.async_get_cache( + key="unified-output", litellm_parent_otel_span=None, - llm_router=None, + ) + assert cached_row["file_object"]["bytes"] == 836 + + +@pytest.mark.asyncio +async def test_afile_retrieve_refreshes_marked_fallback_without_router_from_model_name(): + row: Final = _marked_fallback_file_row().model_copy( + update={"model_mappings": {"bedrock/model-x": "provider-output"}} + ) + provider_object: Final = OpenAIFileObject( + id="provider-output", + object="file", + bytes=836, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + status="processed", + ) + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(row) + + with patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + return_value=provider_object, + ) as retrieve: + response: Final = await proxy_managed_files.afile_retrieve( + file_id="unified-output", + litellm_parent_otel_span=None, + llm_router=None, + ) + + assert response.bytes == 836 + assert response.id == "unified-output" + assert response.litellm_details_fallback is None + retrieve.assert_awaited_once() + assert retrieve.await_args.kwargs["custom_llm_provider"] == "bedrock" + assert retrieve.await_args.kwargs["max_retries"] == 0 + assert "_litellm_internal_model_credentials" not in retrieve.await_args.kwargs + updated_row: Final = managed_file_table.rows["unified-output"] + assert updated_row.file_object == response + assert updated_row.model_mappings == {"bedrock/model-x": "provider-output"} + assert updated_row.created_by == "user-123" + assert updated_row.team_id == "team-123" + + +@pytest.mark.asyncio +async def test_afile_retrieve_uses_model_name_when_router_does_not_know_deployment(): + row: Final = _marked_fallback_file_row().model_copy( + update={"model_mappings": {"bedrock/model-x": "provider-output"}} + ) + provider_object: Final = OpenAIFileObject( + id="provider-output", + object="file", + bytes=836, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + status="processed", + ) + router: Final = MagicMock() + router.get_deployment_credentials_with_provider.return_value = None + proxy_managed_files, _ = _managed_files_with_fake_prisma(row) + + with patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + return_value=provider_object, + ) as retrieve: + response: Final = await proxy_managed_files.afile_retrieve( + file_id="unified-output", + litellm_parent_otel_span=None, + llm_router=router, + ) + + assert response.bytes == 836 + retrieve.assert_awaited_once() + router.get_deployment_credentials_with_provider.assert_called_once_with( + "bedrock/model-x" + ) + assert retrieve.await_args.kwargs["custom_llm_provider"] == "bedrock" + assert retrieve.await_args.kwargs["max_retries"] == 0 + assert "_litellm_internal_model_credentials" not in retrieve.await_args.kwargs + + +@pytest.mark.asyncio +async def test_afile_retrieve_does_not_recreate_deleted_row_after_refresh(): + row: Final = _marked_fallback_file_row() + provider_object: Final = OpenAIFileObject( + id="provider-output", + object="file", + bytes=836, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + status="processed", + ) + router: Final = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma(row) + + async def delete_row_then_return_file( + *, file_id: str, **_options: object + ) -> OpenAIFileObject: + assert file_id == "provider-output" + deleted_row: Final = await managed_file_table.delete( + where={"unified_file_id": "unified-output"} + ) + assert deleted_row is not None + return provider_object + + with patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=delete_row_then_return_file, + ): + response: Final = await proxy_managed_files.afile_retrieve( + file_id="unified-output", + litellm_parent_otel_span=None, + llm_router=router, + ) + + assert response.bytes == 836 + assert "unified-output" not in managed_file_table.rows + assert managed_file_table.upsert_calls == [] + cached_row: Final = await proxy_managed_files.internal_usage_cache.async_get_cache( + key="unified-output", + litellm_parent_otel_span=None, + ) + assert cached_row is None + + +@pytest.mark.asyncio +async def test_afile_retrieve_case3_includes_provider_error_text(): + stored_file: Final = LiteLLM_ManagedFileTable( + unified_file_id="unified-output", + file_object=None, + model_mappings={"model-123": "provider-output"}, + flat_model_file_ids=["provider-output"], + created_by="user-123", + ) + router: Final = MagicMock() + router.get_deployment_credentials_with_provider.return_value = None + proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file) + + with ( + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=HTTPException( + status_code=404, + detail="provider file is missing", + ), + ), + pytest.raises( + Exception, + match="Failed to retrieve file unified-output from provider", + ) as error, + ): + await proxy_managed_files.afile_retrieve( + file_id="unified-output", + litellm_parent_otel_span=None, + llm_router=router, + ) + + assert type(error.value) is Exception + assert str(error.value) == ( + "Failed to retrieve file unified-output from provider: " + "404: provider file is missing" ) - # Should return the stored file object directly - assert result == stored_file_object + +@pytest.mark.asyncio +async def test_afile_retrieve_case3_uses_default_provider_for_unknown_router_deployment(): + stored_file: Final = LiteLLM_ManagedFileTable( + unified_file_id="unified-output", + file_object=None, + model_mappings={"model-123": "provider-output"}, + flat_model_file_ids=["provider-output"], + created_by="user-123", + ) + router: Final = MagicMock() + router.get_deployment_credentials_with_provider.return_value = None + proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file) + provider_file_object: Final = OpenAIFileObject( + id="provider-output", + object="file", + bytes=836, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + status="processed", + ) + + with patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + return_value=provider_file_object, + ) as retrieve: + response: Final = await proxy_managed_files.afile_retrieve( + file_id="unified-output", + litellm_parent_otel_span=None, + llm_router=router, + ) + + assert response.bytes == 836 + retrieve.assert_awaited_once() + assert retrieve.await_args.kwargs["file_id"] == "provider-output" + assert retrieve.await_args.kwargs["max_retries"] == 0 + assert "custom_llm_provider" not in retrieve.await_args.kwargs + assert "_litellm_internal_model_credentials" not in retrieve.await_args.kwargs + + +@pytest.mark.asyncio +async def test_afile_retrieve_case3_times_out_on_default_provider_route(): + stored_file: Final = LiteLLM_ManagedFileTable( + unified_file_id="unified-output", + file_object=None, + model_mappings={"model-123": "provider-output"}, + flat_model_file_ids=["provider-output"], + created_by="user-123", + ) + router: Final = MagicMock() + router.get_deployment_credentials_with_provider.return_value = None + proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file) + + async def never_answers(**_options: object) -> None: + await asyncio.Event().wait() + + with ( + patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=never_answers, + ) as retrieve, + patch( + "litellm_enterprise.proxy.hooks.managed_files.BATCH_OUTPUT_FILE_LOOKUP_TIMEOUT_SECONDS", + 0.01, + ), + pytest.raises(Exception, match="Provider file retrieve timed out"), + ): + await asyncio.wait_for( + proxy_managed_files.afile_retrieve( + file_id="unified-output", + litellm_parent_otel_span=None, + llm_router=router, + ), + timeout=1, + ) + + retrieve.assert_awaited_once() + assert retrieve.await_args.kwargs["file_id"] == "provider-output" + assert retrieve.await_args.kwargs["max_retries"] == 0 + assert "custom_llm_provider" not in retrieve.await_args.kwargs + + +@pytest.mark.asyncio +async def test_afile_retrieve_case3_without_route_raises_accurate_error(): + import litellm.proxy.proxy_server as proxy_server_module + + stored_file: Final = LiteLLM_ManagedFileTable( + unified_file_id="unified-output", + file_object=None, + model_mappings={"my-team/gpt-4o": "provider-output"}, + flat_model_file_ids=["provider-output"], + created_by="user-123", + ) + proxy_managed_files, _ = _managed_files_with_fake_prisma(stored_file) + with ( + patch.object(proxy_server_module, "llm_router", None), + patch("litellm.afile_retrieve", new_callable=AsyncMock) as retrieve, + pytest.raises(Exception, match="no provider route to fetch it"), + ): + await proxy_managed_files.afile_retrieve( + file_id="unified-output", + litellm_parent_otel_span=None, + llm_router=None, + ) + + retrieve.assert_not_called() + + +@pytest.mark.asyncio +async def test_afile_retrieve_returns_marked_fallback_without_marker_when_refresh_fails(): + router = MagicMock() + router.get_deployment_credentials_with_provider.return_value = {"api_key": "key"} + proxy_managed_files, managed_file_table = _managed_files_with_fake_prisma( + _marked_fallback_file_row() + ) + + with patch( + "litellm.afile_retrieve", + new_callable=AsyncMock, + side_effect=HTTPException(status_code=404), + ) as retrieve: + response = await proxy_managed_files.afile_retrieve( + file_id="unified-output", + litellm_parent_otel_span=None, + llm_router=router, + ) + + assert response.bytes == 0 + assert response.id == "unified-output" + assert "litellm_details_fallback" not in response.model_dump() + retrieve.assert_awaited_once() + assert managed_file_table.upsert_calls == [] @pytest.mark.asyncio @@ -1710,16 +3239,7 @@ async def test_afile_retrieve_raises_error_for_non_managed_file(): Test that afile_retrieve raises an error when the file_id is not found in the managed files table. """ - prisma_client = AsyncMock() - internal_usage_cache = MagicMock() - - proxy_managed_files = PROXY_LiteLLMManagedFiles( - internal_usage_cache=internal_usage_cache, - prisma_client=prisma_client, - ) - - # Mock get_unified_file_id to return None (file not found) - proxy_managed_files.get_unified_file_id = AsyncMock(return_value=None) + proxy_managed_files, _ = _managed_files_with_fake_prisma() with pytest.raises(Exception, match='LiteLLM Managed File object with id=non-existent-file-id') as exc_info: await proxy_managed_files.afile_retrieve( @@ -1936,6 +3456,47 @@ async def test_list_batches_registers_and_returns_unified_output_file_ids(): assert c.kwargs["data"]["create"]["team_id"] == "owner-team" +@pytest.mark.asyncio +async def test_afile_list_returns_persisted_batch_output_file(): + from litellm.proxy._types import UserAPIKeyAuth + + proxy_managed_files, _ = _managed_files_with_fake_prisma( + _marked_fallback_file_row() + ) + + result = await proxy_managed_files.afile_list( + purpose="batch_output", + user_api_key_dict=UserAPIKeyAuth(user_id="owner-user"), + litellm_parent_otel_span=None, + limit=10, + ) + + listed = result.data[0] + assert listed.id == "unified-output" + assert listed.filename == "output.jsonl" + assert listed.bytes == 0 + assert listed.purpose == "batch_output" + assert "litellm_details_fallback" not in listed.model_dump() + + +@pytest.mark.asyncio +async def test_get_user_created_file_ids_strips_fallback_marker(): + from litellm.proxy._types import UserAPIKeyAuth + + proxy_managed_files, _ = _managed_files_with_fake_prisma( + _marked_fallback_file_row() + ) + + files = await proxy_managed_files.get_user_created_file_ids( + UserAPIKeyAuth(user_id="user-123"), + ["provider-output"], + ) + + assert len(files) == 1 + assert files[0].id == "unified-output" + assert "litellm_details_fallback" not in files[0].model_dump() + + @pytest.mark.asyncio async def test_list_batches_resolves_existing_managed_rows_without_minting(): """When the raw provider file IDs already have managed file rows, listing must @@ -3453,8 +5014,10 @@ async def test_post_call_batch_sync_does_not_claim_ownership(): prisma_client = AsyncMock() prisma_client.db.litellm_managedobjecttable.update_many.return_value = 0 proxy_managed_files = PROXY_LiteLLMManagedFiles( - MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client + MagicMock(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock()), + prisma_client=prisma_client, ) + prisma_client.db.litellm_managedfiletable.find_first.return_value = None await proxy_managed_files.async_post_call_success_hook( data={"batch_id": MODEL_ENCODED_BATCH_ID}, @@ -3477,22 +5040,16 @@ async def test_post_call_batch_sync_updates_existing_row(): prisma_client.db.litellm_managedobjecttable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) - proxy_managed_files = PROXY_LiteLLMManagedFiles( - MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client - ) + proxy_managed_files = PROXY_LiteLLMManagedFiles(MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client) await proxy_managed_files.async_post_call_success_hook( data={"batch_id": MODEL_ENCODED_BATCH_ID}, - user_api_key_dict=UserAPIKeyAuth( - user_id="user_a", team_id="team_a", parent_otel_span=MagicMock() - ), + user_api_key_dict=UserAPIKeyAuth(user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()), response=_batch_response(MODEL_ENCODED_BATCH_ID), ) update_call = prisma_client.db.litellm_managedobjecttable.update_many.await_args - assert update_call.kwargs["where"] == { - "unified_object_id": MODEL_ENCODED_BATCH_ID - } + assert update_call.kwargs["where"] == {"unified_object_id": MODEL_ENCODED_BATCH_ID} assert update_call.kwargs["data"]["status"] == "completed" prisma_client.db.litellm_managedobjecttable.upsert.assert_not_awaited() @@ -3512,8 +5069,10 @@ async def test_post_call_batch_sync_stores_output_file_ownership_from_batch_row( _owned_record(created_by="user_a", team_id="team_a") ) proxy_managed_files = PROXY_LiteLLMManagedFiles( - MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client + MagicMock(async_get_cache=AsyncMock(return_value=None), async_set_cache=AsyncMock()), + prisma_client=prisma_client, ) + prisma_client.db.litellm_managedfiletable.find_first.return_value = None await proxy_managed_files.async_post_call_success_hook( data={"batch_id": MODEL_ENCODED_BATCH_ID}, diff --git a/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py b/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py index c48af1f6177..c6ab83a0d80 100644 --- a/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py +++ b/tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py @@ -9,12 +9,12 @@ import pytest from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.openai_files_endpoints.common_utils import ( + ManagedBatchOutputFileWriter, ensure_batch_response_managed_file_ids, get_batch_from_database, ) from litellm.types.utils import LiteLLMBatch - UNIFIED_BATCH_ID = "litellm_proxy;model_id:my-model;llm_batch_id:batch-raw-123" ENCODED_UNIFIED_BATCH_ID = ( base64.urlsafe_b64encode(UNIFIED_BATCH_ID.encode()).decode().rstrip("=") @@ -40,9 +40,9 @@ def _build_batch_response( def _build_managed_files_mock(unified_id: str = "file-bWFuYWdlZF9vdXRwdXRfaWQ="): - mock = MagicMock() + mock = MagicMock(spec=ManagedBatchOutputFileWriter) mock.get_unified_output_file_id = MagicMock(return_value=unified_id) - mock.store_unified_file_id = AsyncMock() + mock.store_batch_output_file = AsyncMock() return mock @@ -73,11 +73,12 @@ async def test_ensure_batch_response_derives_model_id_from_unified_batch_id(): ) assert response.output_file_id == unified_output_file_id - mock_managed_files.store_unified_file_id.assert_called_once() - store_kwargs = mock_managed_files.store_unified_file_id.call_args.kwargs - assert store_kwargs["model_mappings"] == {"my-model": raw_output_file_id} - assert store_kwargs["user_api_key_dict"].user_id == "batch-owner" - assert store_kwargs["user_api_key_dict"].team_id == "team-owner" + mock_managed_files.store_batch_output_file.assert_awaited_once() + store_kwargs = mock_managed_files.store_batch_output_file.await_args.kwargs + assert store_kwargs["model_id"] == "my-model" + assert store_kwargs["provider_file_id"] == raw_output_file_id + assert store_kwargs["owner"].user_id == "batch-owner" + assert store_kwargs["owner"].team_id == "team-owner" @pytest.mark.asyncio @@ -100,13 +101,11 @@ async def test_ensure_batch_response_registers_output_and_error_file_ids(): assert response.output_file_id == unified_id assert response.error_file_id == unified_id - assert mock_managed_files.store_unified_file_id.call_count == 2 - mappings = [ - call.kwargs["model_mappings"] - for call in mock_managed_files.store_unified_file_id.call_args_list + assert mock_managed_files.store_batch_output_file.await_count == 2 + stored_provider_file_ids = [ + call.kwargs["provider_file_id"] for call in mock_managed_files.store_batch_output_file.await_args_list ] - assert {"my-model": "file-raw-output"} in mappings - assert {"my-model": "file-raw-error"} in mappings + assert stored_provider_file_ids == ["file-raw-output", "file-raw-error"] @pytest.mark.asyncio @@ -148,10 +147,11 @@ async def test_get_batch_from_database_registers_missing_output_file_id(): assert response is not None assert response.output_file_id == unified_output_file_id - mock_managed_files.store_unified_file_id.assert_called_once() - store_kwargs = mock_managed_files.store_unified_file_id.call_args.kwargs - assert store_kwargs["model_mappings"] == {"my-model": raw_output_file_id} - assert store_kwargs["user_api_key_dict"].user_id == "batch-owner" + mock_managed_files.store_batch_output_file.assert_awaited_once() + store_kwargs = mock_managed_files.store_batch_output_file.await_args.kwargs + assert store_kwargs["model_id"] == "my-model" + assert store_kwargs["provider_file_id"] == raw_output_file_id + assert store_kwargs["owner"].user_id == "batch-owner" @pytest.mark.asyncio @@ -172,9 +172,7 @@ async def test_ensure_batch_response_uses_batch_owner_when_db_batch_object_prese ) # batch owner from db_batch_object wins over the caller auth context - forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[ - "user_api_key_dict" - ] + forwarded_auth = mock_managed_files.store_batch_output_file.await_args.kwargs["owner"] assert forwarded_auth.user_id == "batch-owner" assert forwarded_auth.team_id == "team-owner" @@ -190,7 +188,10 @@ async def test_registered_output_file_row_denies_cross_user_access(): prisma.db.litellm_managedfiletable.upsert = AsyncMock() prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) managed_files = PROXY_LiteLLMManagedFiles( - internal_usage_cache=MagicMock(), + internal_usage_cache=MagicMock( + async_get_cache=AsyncMock(return_value=None), + async_set_cache=AsyncMock(), + ), prisma_client=prisma, ) response = _build_batch_response(output_file_id=raw_output_file_id) diff --git a/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py b/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py index 2049ff58f95..cf305c660cd 100644 --- a/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py +++ b/tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py @@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, MagicMock from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.openai_files_endpoints.common_utils import ( + ManagedBatchOutputFileWriter, ensure_batch_response_managed_file_ids, update_batch_in_database, ) @@ -39,9 +40,9 @@ def _build_batch_response( def _build_managed_files_mock(unified_id: str = "file-bWFuYWdlZF9vdXRwdXRfaWQ="): - mock = MagicMock() + mock = MagicMock(spec=ManagedBatchOutputFileWriter) mock.get_unified_output_file_id = MagicMock(return_value=unified_id) - mock.store_unified_file_id = AsyncMock() + mock.store_batch_output_file = AsyncMock() return mock @@ -114,9 +115,7 @@ async def test_cancel_path_registers_output_file_under_batch_owner(): operation="cancel", ) - forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[ - "user_api_key_dict" - ] + forwarded_auth = mock_managed_files.store_batch_output_file.await_args.kwargs["owner"] assert forwarded_auth.user_id == "batch-owner" assert forwarded_auth.team_id == "batch-team" stored = json.loads( @@ -156,9 +155,7 @@ async def test_update_batch_skips_lookup_when_db_batch_object_supplied(): ) mock_prisma.db.litellm_managedobjecttable.find_first.assert_not_called() - forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[ - "user_api_key_dict" - ] + forwarded_auth = mock_managed_files.store_batch_output_file.await_args.kwargs["owner"] assert forwarded_auth.user_id == "caller-owner" assert forwarded_auth.team_id == "caller-team" @@ -222,11 +219,11 @@ async def test_ensure_batch_response_swallows_conversion_errors(): hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"}, ) - mock_managed_files = MagicMock() + mock_managed_files = MagicMock(spec=ManagedBatchOutputFileWriter) mock_managed_files.get_unified_output_file_id = MagicMock( side_effect=RuntimeError("boom") ) - mock_managed_files.store_unified_file_id = AsyncMock() + mock_managed_files.store_batch_output_file = AsyncMock() mock_logger = MagicMock() await ensure_batch_response_managed_file_ids( @@ -263,9 +260,7 @@ async def test_ensure_batch_response_builds_auth_from_db_batch_object(): db_batch_object=db_batch_object, ) - forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[ - "user_api_key_dict" - ] + forwarded_auth = mock_managed_files.store_batch_output_file.await_args.kwargs["owner"] assert forwarded_auth.user_id == "user-from-db" assert forwarded_auth.team_id == "team-from-db" diff --git a/tests/unit/enterprise/proxy/test_managed_files_hook.py b/tests/unit/enterprise/proxy/test_managed_files_hook.py index 3ea3b97e1fe..7bf53b11cdb 100644 --- a/tests/unit/enterprise/proxy/test_managed_files_hook.py +++ b/tests/unit/enterprise/proxy/test_managed_files_hook.py @@ -183,7 +183,9 @@ def _make_managed_files_instance(): ) mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) mock_prisma = MagicMock() + mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) instance = PROXY_LiteLLMManagedFiles( internal_usage_cache=mock_cache, @@ -1384,9 +1386,11 @@ def _make_real_managed_files_instance(): ) mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_set_cache = AsyncMock() mock_prisma = MagicMock() + mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) mock_prisma.db.litellm_managedfiletable.upsert = AsyncMock() mock_prisma.db.litellm_managedfiletable.create = AsyncMock( side_effect=AssertionError( diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index b06ed9468b0..9a30e1990b1 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -2936,8 +2936,11 @@ async def test_cancellation_delivers_termination_over_tcp( listener: Final = await asyncio.start_server(handle_connection, "127.0.0.1", 0) port: Final = listener.sockets[0].getsockname()[1] + client_timeout: Final = 2 if cancel_mode == "read_timeout" else 30 client: Final = MCPClient( - server_url=f"http://127.0.0.1:{port}/mcp", protocol_version=protocol_version, timeout=2 if cancel_mode == "read_timeout" else 30 + server_url=f"http://127.0.0.1:{port}/mcp", + protocol_version=protocol_version, + timeout=client_timeout, ) async def calls(): diff --git a/tests/unit/files/test_main.py b/tests/unit/files/test_main.py index cb70b39f4d5..e91951fb5ba 100644 --- a/tests/unit/files/test_main.py +++ b/tests/unit/files/test_main.py @@ -1,3 +1,4 @@ +from types import MappingProxyType from typing import Final from urllib.parse import parse_qs, urlparse @@ -121,3 +122,68 @@ async def test_afile_retrieve_rejects_a_provider_file_without_its_size(): assert exc_info.value.title == "OpenAIFileObject" assert [error["loc"] for error in exc_info.value.errors()] == [("bytes",)] + + +_FILE_BODY: Final = b'{"prompt": "Hello", "completion": "Hi"}' +_FINE_TUNE_FILE_JSON: Final = MappingProxyType( + { + "id": "file-abc123", + "object": "file", + "bytes": len(_FILE_BODY), + "created_at": 1699000000, + "filename": "mydata.jsonl", + "purpose": "fine-tune", + } +) + + +@pytest.mark.asyncio +async def test_openai_file_operations_roundtrip(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + files_route: Final = respx_mock.post("https://api.openai.com/v1/files").mock( + return_value=httpx.Response(200, json=dict(_FINE_TUNE_FILE_JSON)) + ) + list_route: Final = respx_mock.get("https://api.openai.com/v1/files").mock( + return_value=httpx.Response(200, json={"object": "list", "data": [dict(_FINE_TUNE_FILE_JSON)]}) + ) + retrieve_route: Final = respx_mock.get("https://api.openai.com/v1/files/file-abc123").mock( + return_value=httpx.Response(200, json=dict(_FINE_TUNE_FILE_JSON)) + ) + content_route: Final = respx_mock.get("https://api.openai.com/v1/files/file-abc123/content").mock( + return_value=httpx.Response(200, content=_FILE_BODY) + ) + delete_route: Final = respx_mock.delete("https://api.openai.com/v1/files/file-abc123").mock( + return_value=httpx.Response(200, json={"id": "file-abc123", "object": "file", "deleted": True}) + ) + + uploaded: Final = await litellm.acreate_file( + file=("mydata.jsonl", _FILE_BODY), purpose="fine-tune", custom_llm_provider="openai", api_key="fake-key" + ) + assert files_route.call_count == 1 + upload_body: Final = files_route.calls.last.request.content + assert b'name="purpose"\r\n\r\nfine-tune' in upload_body + assert b'filename="mydata.jsonl"' in upload_body + assert _FILE_BODY in upload_body + assert uploaded.id == "file-abc123" + + listed: Final = await litellm.afile_list(custom_llm_provider="openai", api_key="fake-key") + assert list_route.call_count == 1 + assert [file.id for file in listed.data] == ["file-abc123"] + + retrieved: Final = await litellm.afile_retrieve( + file_id="file-abc123", custom_llm_provider="openai", api_key="fake-key" + ) + assert retrieve_route.call_count == 1 + assert retrieved.filename == "mydata.jsonl" + assert retrieved.purpose == "fine-tune" + + content: Final = await litellm.afile_content( + file_id="file-abc123", custom_llm_provider="openai", api_key="fake-key" + ) + assert content_route.call_count == 1 + assert content.content == _FILE_BODY + + deleted: Final = await litellm.afile_delete(file_id="file-abc123", custom_llm_provider="openai", api_key="fake-key") + assert delete_route.call_count == 1 + assert deleted.id == "file-abc123" + assert deleted.deleted is True diff --git a/tests/unit/images/test_image_edit.py b/tests/unit/images/test_image_edit.py new file mode 100644 index 00000000000..583474cba90 --- /dev/null +++ b/tests/unit/images/test_image_edit.py @@ -0,0 +1,152 @@ +import asyncio +import io +from collections.abc import Iterator, Mapping +from datetime import datetime +from typing import Final + +import httpx +import pytest +import respx +from pydantic import TypeAdapter +from typing_extensions import ReadOnly, TypedDict, override + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.types.utils import ImageResponse + +_PNG_SIGNATURE: Final = b"\x89PNG\r\n\x1a\n" +_FIRST_IMAGE: Final = _PNG_SIGNATURE + b"first-reference-image" +_SECOND_IMAGE: Final = _PNG_SIGNATURE + b"second-reference-image" +_EDITED_IMAGE_B64: Final = ( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg==" +) +_TEXT_TOKENS: Final = 50 +_IMAGE_TOKENS: Final = 50 +_OUTPUT_TOKENS: Final = 1000 +_EDIT_RESPONSE: Final = { + "created": 1589478378, + "data": [{"b64_json": _EDITED_IMAGE_B64}], + "usage": { + "total_tokens": _TEXT_TOKENS + _IMAGE_TOKENS + _OUTPUT_TOKENS, + "input_tokens": _TEXT_TOKENS + _IMAGE_TOKENS, + "input_tokens_details": {"image_tokens": _IMAGE_TOKENS, "text_tokens": _TEXT_TOKENS}, + "output_tokens": _OUTPUT_TOKENS, + }, +} + + +class _LoggedImageEdit(TypedDict): + model: ReadOnly[str] + custom_llm_provider: ReadOnly[str] + response_cost: ReadOnly[float] + + +_LOGGED_IMAGE_EDIT: Final = TypeAdapter(_LoggedImageEdit) + + +class _SuccessLogger(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.payload: _LoggedImageEdit | None = None + self.logged: Final = asyncio.Event() + + @override + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + self.payload = _LOGGED_IMAGE_EDIT.validate_python(kwargs.get("standard_logging_object")) + self.logged.set() + + +@pytest.fixture +def httpx_transport(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + +def _multipart_image_parts(request: httpx.Request) -> tuple[bytes, ...]: + body: Final = request.read() + boundary: Final = request.headers["content-type"].split("boundary=", 1)[1].encode() + parts: Final = body.split(b"--" + boundary) + return tuple(part.split(b"\r\n\r\n", 1)[1].removesuffix(b"\r\n") for part in parts if b'name="image[]"' in part) + + +@pytest.mark.asyncio +async def test_openai_image_edit_accepts_bytesio_images(respx_mock: respx.MockRouter, httpx_transport: None) -> None: + route: Final = respx_mock.post("https://api.openai.com/v1/images/edits").mock( + return_value=httpx.Response(200, json=_EDIT_RESPONSE) + ) + + result: Final = await litellm.aimage_edit( + prompt="combine the reference images", + model="gpt-image-1", + image=[io.BytesIO(_FIRST_IMAGE), io.BytesIO(_SECOND_IMAGE)], + api_key="fake-key", + ) + + assert isinstance(result, ImageResponse) + assert result.data is not None and result.data[0].b64_json == _EDITED_IMAGE_B64 + assert route.call_count == 1 + assert _multipart_image_parts(route.calls[0].request) == (_FIRST_IMAGE, _SECOND_IMAGE) + + +@pytest.mark.asyncio +async def test_openai_image_edit_accepts_mixed_bytes_and_bytesio( + respx_mock: respx.MockRouter, httpx_transport: None +) -> None: + route: Final = respx_mock.post("https://api.openai.com/v1/images/edits").mock( + return_value=httpx.Response(200, json=_EDIT_RESPONSE) + ) + + result: Final = await litellm.aimage_edit( + prompt="Create a cohesive artistic style across all images", + model="gpt-image-1", + image=[_FIRST_IMAGE, io.BytesIO(_SECOND_IMAGE)], + api_key="fake-key", + ) + + assert isinstance(result, ImageResponse) + assert result.data is not None and len(result.data) == 1 + assert result.data[0].b64_json == _EDITED_IMAGE_B64 + assert route.call_count == 1 + assert _multipart_image_parts(route.calls[0].request) == (_FIRST_IMAGE, _SECOND_IMAGE) + + +@pytest.mark.asyncio +async def test_azure_image_edit_logs_deployment_model_and_positive_cost( + respx_mock: respx.MockRouter, httpx_transport: None, monkeypatch: pytest.MonkeyPatch +) -> None: + logger: Final = _SuccessLogger() + monkeypatch.setattr(litellm, "callbacks", [logger]) + route: Final = respx_mock.post( + url__startswith="https://fake.openai.azure.com/openai/deployments/CUSTOM_AZURE_DEPLOYMENT_NAME/images/edits" + ).mock(return_value=httpx.Response(200, json=_EDIT_RESPONSE)) + + result: Final = await litellm.aimage_edit( + prompt="combine the reference images", + model="azure/CUSTOM_AZURE_DEPLOYMENT_NAME", + base_model="azure/gpt-image-1", + image=[_FIRST_IMAGE, _SECOND_IMAGE], + api_key="fake-key", + api_base="https://fake.openai.azure.com", + api_version="2025-04-01-preview", + ) + await asyncio.wait_for(logger.logged.wait(), timeout=10) + + assert isinstance(result, ImageResponse) + assert route.call_count == 1 + payload: Final = logger.payload + assert payload is not None + assert payload["model"] == "CUSTOM_AZURE_DEPLOYMENT_NAME" + assert payload["custom_llm_provider"] == "azure" + pricing: Final = litellm.model_cost["azure/gpt-image-1"] + expected_cost: Final = ( + _TEXT_TOKENS * pricing["input_cost_per_token"] + + _IMAGE_TOKENS * pricing["input_cost_per_image_token"] + + _OUTPUT_TOKENS * pricing["output_cost_per_image_token"] + ) + assert expected_cost > 0 + assert payload["response_cost"] == pytest.approx(expected_cost) + assert result._hidden_params["response_cost"] == pytest.approx(expected_cost) # pyright: ignore[reportPrivateUsage] # cost is only surfaced on _hidden_params diff --git a/tests/unit/integrations/SlackAlerting/test_slack_alerting_delivery.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting_delivery.py new file mode 100644 index 00000000000..c8fb5d30376 --- /dev/null +++ b/tests/unit/integrations/SlackAlerting/test_slack_alerting_delivery.py @@ -0,0 +1,344 @@ +import asyncio +import datetime +from collections.abc import Sequence +from typing import Final, Literal, TypedDict + +from typing_extensions import ReadOnly + +import httpx +import pytest +import respx +from pydantic import TypeAdapter + +import litellm +import litellm.proxy.proxy_server as proxy_server +from litellm.caching.caching import DualCache +from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.langfuse.langfuse import LangFuseLogger, installed_langfuse_version +from litellm.integrations.langfuse.langfuse_sdk import ( + build_langfuse_client, + build_langfuse_tracing, + resolve_trace_id, +) +from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting +from litellm.integrations.SlackAlerting.utils import add_langfuse_trace_id_to_alert +from litellm.litellm_core_utils import litellm_logging +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.proxy.utils import ProxyLogging +from litellm.types.integrations.slack_alerting import AlertType +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +_WEBHOOK: Final = "https://hooks.slack.example/services/delivery" +_LANGFUSE_HOST: Final = "https://langfuse.alerts.example" +_AZURE_BASE: Final = "https://openai-gpt-4-test-v-1.openai.azure.com/" +_DAILY_BASE: Final = "https://daily-report.openai.example/v1" + + +class _SlackPayload(TypedDict): + text: ReadOnly[str] + + +class _TeamRow(TypedDict): + team_alias: ReadOnly[str] + total_spend: ReadOnly[float] + + +class _TagRow(TypedDict): + individual_request_tag: ReadOnly[str] + total_spend: ReadOnly[float] + + +_PAYLOAD: Final = TypeAdapter(_SlackPayload) + + +@pytest.fixture(autouse=True) +def _slack_webhook(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setenv("SLACK_WEBHOOK_URL", _WEBHOOK) + + +def _webhook(respx_mock: respx.MockRouter) -> respx.Route: + return respx_mock.post(_WEBHOOK).mock(return_value=httpx.Response(200, text="ok")) + + +def _posted_texts(route: respx.Route) -> tuple[str, ...]: + return tuple(_PAYLOAD.validate_json(call.request.content)["text"] for call in route.calls) + + +@pytest.mark.asyncio +async def test_slow_response_alert_names_the_azure_api_base_and_reaches_the_webhook( + respx_mock: respx.MockRouter, +) -> None: + route: Final = _webhook(respx_mock) + proxy_logging: Final = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging.update_values(alerting=["slack"], alerting_threshold=100, redis_cache=None) + start: Final = datetime.datetime(2026, 1, 1, 12, 0, 0) + messages: Final = ({"role": "user", "content": "Hey how's it going?"},) + + helper_result: Final = proxy_logging.slack_alerting_instance._response_taking_too_long_callback_helper( + kwargs={ + "model": "chatgpt-v-3", + "messages": messages, + "litellm_params": {"api_base": _AZURE_BASE, "custom_llm_provider": "azure"}, + }, + start_time=start, + end_time=start + datetime.timedelta(seconds=150), + ) + + assert helper_result == (150.0, "chatgpt-v-3", _AZURE_BASE, str(messages)[:100]) + + slow_message: Final = ( + f"`Responses are slow - 150.0s response time > Alerting threshold: 100s`\nAPI Base: `{_AZURE_BASE}`" + ) + await proxy_logging.alerting_handler(message=slow_message, level="Low", alert_type=AlertType.llm_too_slow) + await proxy_logging.slack_alerting_instance.flush_queue() + + texts: Final = _posted_texts(route) + assert len(texts) == 1 + assert texts[0].startswith("Alert type: `llm_too_slow`\nLevel: `Low`\n") + assert texts[0].endswith(f"Message: {slow_message}") + + +@pytest.mark.asyncio +async def test_send_alert_is_queued_until_flush_then_posted_to_the_webhook_once(respx_mock: respx.MockRouter) -> None: + route: Final = _webhook(respx_mock) + slack_alerting: Final = SlackAlerting(alerting_threshold=1, internal_usage_cache=DualCache(), alerting=["slack"]) + + await slack_alerting.send_alert("Test message", "Low", AlertType.budget_alerts, alerting_metadata={}) + + assert route.call_count == 0 + + await slack_alerting.flush_queue() + await slack_alerting.flush_queue() + + texts: Final = _posted_texts(route) + assert len(texts) == 1 + assert texts[0].startswith("Alert type: `budget_alerts`\nLevel: `Low`\n") + assert texts[0].endswith("Message: Test message") + + +@pytest.mark.asyncio +async def test_a_queued_alert_is_posted_by_the_periodic_flush_without_a_manual_flush( + respx_mock: respx.MockRouter, +) -> None: + delivered: Final = asyncio.Event() + + def deliver(request: httpx.Request) -> httpx.Response: + delivered.set() + return httpx.Response(200, text="ok") + + route: Final = respx_mock.post(_WEBHOOK).mock(side_effect=deliver) + slack_alerting: Final = SlackAlerting(alerting_threshold=1, internal_usage_cache=DualCache(), alerting=["slack"]) + slack_alerting.flush_interval = 0 + slack_alerting.update_values(alerting=["slack"]) + flush_task: Final = slack_alerting._periodic_flush_task + assert flush_task is not None + try: + await slack_alerting.send_alert("Timed message", "Low", AlertType.budget_alerts, alerting_metadata={}) + await asyncio.wait_for(delivered.wait(), timeout=5) + finally: + flush_task.cancel() + + texts: Final = _posted_texts(route) + assert len(texts) == 1 + assert texts[0].endswith("Message: Timed message") + + +class _DeploymentSettled(CustomLogger): + def __init__(self, model_id: str) -> None: + super().__init__() + self.model_id: Final = model_id + self.succeeded: Final = asyncio.Event() + self.failed: Final = asyncio.Event() + + def _is_mine(self, kwargs: dict[str, object]) -> bool: + payload: Final = kwargs.get("standard_logging_object") + return isinstance(payload, dict) and payload.get("model_id") == self.model_id + + async def async_log_success_event( + self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + if self._is_mine(kwargs): + self.succeeded.set() + + async def async_log_failure_event( + self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + if self._is_mine(kwargs): + self.failed.set() + + +@pytest.mark.asyncio +async def test_daily_report_lists_router_latency_after_success_and_failures_after_an_auth_error( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +) -> None: + webhook: Final = _webhook(respx_mock) + respx_mock.post(f"{_DAILY_BASE}/chat/completions").mock( + side_effect=( + httpx.Response( + 200, + json={ + "id": "chatcmpl-daily", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-5-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "fine"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 4, "total_tokens": 9}, + }, + ), + httpx.Response( + 401, + json={ + "error": { + "message": "Incorrect API key provided", + "type": "invalid_request_error", + "code": "invalid_api_key", + } + }, + ), + ) + ) + model_id: Final = "daily-report-deployment" + slack_alerting: Final = SlackAlerting( + alerting=["slack"], internal_usage_cache=DualCache(), alert_types=[AlertType.daily_reports] + ) + settled: Final = _DeploymentSettled(model_id) + monkeypatch.setattr(litellm, "callbacks", [slack_alerting, settled]) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "daily-report-model", + "litellm_params": {"model": "openai/gpt-5-mini", "api_key": "sk-daily", "api_base": _DAILY_BASE}, + "model_info": {"id": model_id}, + } + ] + ) + request: Final = ({"role": "user", "content": "Hey, how's it going?"},) + + await router.acompletion(model="daily-report-model", messages=list(request)) + await asyncio.wait_for(settled.succeeded.wait(), timeout=5) + after_success: Final = await slack_alerting.send_daily_reports(router=router) + await slack_alerting.flush_queue() + + with pytest.raises(litellm.AuthenticationError): + await router.acompletion(model="daily-report-model", messages=list(request)) + await asyncio.wait_for(settled.failed.wait(), timeout=5) + after_failure: Final = await slack_alerting.send_daily_reports(router=router) + await slack_alerting.flush_queue() + + texts: Final = _posted_texts(webhook) + assert (after_success, after_failure) == (True, True) + assert len(texts) == 2 + assert "Most Failed Requests:*\n\n\tNone\n" in texts[0] + assert "1. Deployment: `openai/gpt-5-mini`, Latency per output token: `" in texts[0] + assert f"1. Deployment: `openai/gpt-5-mini`, Failed Requests: `1`, API Base: `{_DAILY_BASE}`" in texts[1] + assert "Top Slowest Deployments:*\n\n\tNone\n" in texts[1] + + +class _CallLogged(CustomLogger): + def __init__(self, call_id: str, loop: asyncio.AbstractEventLoop) -> None: + super().__init__() + self.call_id: Final = call_id + self.loop: Final = loop + self.logged: Final = asyncio.Event() + + def log_success_event( + self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + if kwargs.get("litellm_call_id") == self.call_id: + self.loop.call_soon_threadsafe(self.logged.set) + + +@pytest.mark.asyncio +async def test_langfuse_trace_link_ends_with_the_trace_id_the_logger_emitted(monkeypatch: pytest.MonkeyPatch) -> None: + logger: Final = LangFuseLogger.__new__(LangFuseLogger) + logger.tracing = build_langfuse_tracing( + exporter=InMemorySpanExporter(), environment=None, release=None, sample_rate=1.0, flush_interval_millis=10 + ) + logger.api_client = build_langfuse_client( + public_key="pk-alert-trace", secret_key="sk-alert-trace", base_url=_LANGFUSE_HOST, httpx_client=None + ) + logger.langfuse_sdk_version = installed_langfuse_version() + call_id: Final = "slack-alert-langfuse-trace" + logged: Final = _CallLogged(call_id, asyncio.get_running_loop()) + monkeypatch.setenv("LANGFUSE_HOST", _LANGFUSE_HOST) + monkeypatch.setattr(litellm_logging, "langFuseLogger", logger) + monkeypatch.setattr(litellm, "success_callback", ["langfuse", logged]) + monkeypatch.setattr(litellm, "_async_success_callback", []) + monkeypatch.setattr(litellm, "callbacks", []) + logging_obj: Final = Logging( + model="gpt-5-mini", + messages=[{"role": "user", "content": "hi"}], + stream=False, + call_type="completion", + litellm_call_id=call_id, + start_time=datetime.datetime.now(), + function_id=call_id, + ) + + litellm.completion( + model="gpt-5-mini", + messages=[{"role": "user", "content": "Hey how's it going?"}], + mock_response="Hey!", + litellm_logging_obj=logging_obj, + ) + await asyncio.wait_for(logged.logged.wait(), timeout=5) + trace_url: Final = await add_langfuse_trace_id_to_alert(request_data={"litellm_logging_obj": logging_obj}) + + expected_trace_id: Final = resolve_trace_id(logging_obj.litellm_trace_id) + assert logging_obj.get_trace_id(service_name="langfuse") == expected_trace_id + assert trace_url == f"{_LANGFUSE_HOST}/trace/{expected_trace_id}" + + +class _SpendReportDb: + def __init__(self, teams: Sequence[_TeamRow], tags: Sequence[_TagRow]) -> None: + self.teams: Final = teams + self.tags: Final = tags + + async def query_raw(self, query: str, *args: object) -> Sequence[_TeamRow] | Sequence[_TagRow]: + return self.teams if "team_alias" in query else self.tags + + +class _SpendReportPrisma: + def __init__(self, db: _SpendReportDb) -> None: + self.db: Final = db + + +@pytest.mark.parametrize("report_type", ["weekly", "monthly"]) +@pytest.mark.asyncio +async def test_spend_report_is_sent_once_per_period( + report_type: Literal["weekly", "monthly"], respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +) -> None: + route: Final = _webhook(respx_mock) + monkeypatch.setattr( + proxy_server, + "prisma_client", + _SpendReportPrisma( + _SpendReportDb( + teams=( + _TeamRow(team_alias="team1", total_spend=100.0), + _TeamRow(team_alias="team2", total_spend=200.0), + ), + tags=( + _TagRow(individual_request_tag="tag1", total_spend=150.0), + _TagRow(individual_request_tag="tag2", total_spend=150.0), + ), + ) + ), + ) + slack_alerting: Final = SlackAlerting(alerting=["slack"], internal_usage_cache=DualCache()) + send_report: Final = ( + slack_alerting.send_weekly_spend_report if report_type == "weekly" else slack_alerting.send_monthly_spend_report + ) + + await send_report() + await slack_alerting.flush_queue() + await send_report() + await slack_alerting.flush_queue() + + texts: Final = _posted_texts(route) + assert len(texts) == 1 + assert "Team: `team1` | Spend: `$100.0`\nTeam: `team2` | Spend: `$200.0`\n" in texts[0] + assert "Tag: `tag1` | Spend: `$150.0`\nTag: `tag2` | Spend: `$150.0`\n" in texts[0] diff --git a/tests/unit/integrations/datadog/test_datadog.py b/tests/unit/integrations/datadog/test_datadog.py index e86d83ba467..eec0022c1db 100644 --- a/tests/unit/integrations/datadog/test_datadog.py +++ b/tests/unit/integrations/datadog/test_datadog.py @@ -2,14 +2,20 @@ import gzip import json import os from datetime import datetime -from typing import Coroutine, Final +from pathlib import Path +from typing import Coroutine, Final, TypedDict from unittest.mock import AsyncMock, patch import pytest +import respx from httpx import Request, Response +from pydantic import TypeAdapter +from typing_extensions import ReadOnly import litellm import litellm.integrations.datadog.datadog as datadog_module +from litellm.caching.caching import Cache +from litellm.caching.llm_caching_handler import LLMClientCache from litellm.integrations.datadog.datadog import DataDogLogger from litellm.integrations.datadog.datadog_handler import ( get_datadog_env, @@ -19,6 +25,7 @@ from litellm.integrations.datadog.datadog_handler import ( get_datadog_source, get_datadog_tags, ) +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.types.integrations.datadog import DatadogInitParams, DatadogPayload, DataDogStatus from litellm.types.utils import ( StandardLoggingHiddenParams, @@ -803,3 +810,113 @@ def create_standard_logging_payload() -> StandardLoggingPayload: additional_headers=None, ), ) + + +_INTAKE_URL: Final = "https://http-intake.logs.test.datadoghq.com/api/v2/logs" + + +class _ServiceEventMessage(TypedDict): + service: ReadOnly[str] + call_type: ReadOnly[str] + error: ReadOnly[str] + is_error: ReadOnly[bool] + + +@pytest.fixture +def delivery( + datadog_env: None, monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> tuple[DataDogLogger, respx.Route]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache()) + monkeypatch.delenv("DD_SOURCE", raising=False) + monkeypatch.delenv("DD_SERVICE", raising=False) + with patch("asyncio.create_task", side_effect=_discard_periodic_flush): + logger: Final = DataDogLogger() + intake: Final = respx_mock.post(_INTAKE_URL).mock(return_value=Response(202, text="Accepted")) + return logger, intake + + +def _delivered_logs(intake: respx.Route) -> list[DatadogPayload]: + return TypeAdapter(list[DatadogPayload]).validate_json(gzip.decompress(intake.calls.last.request.content)) + + +@pytest.mark.asyncio +async def test_a_successful_request_is_delivered_as_an_info_log_carrying_the_standard_payload( + delivery: tuple[DataDogLogger, respx.Route], +) -> None: + datadog_logger, intake = delivery + standard_payload: Final = _standard_logging_payload() + + await datadog_logger.async_log_success_event( + kwargs={"standard_logging_object": standard_payload}, + response_obj=None, + start_time=STANDARD_START_TIME, + end_time=STANDARD_END_TIME, + ) + await datadog_logger.async_send_batch() + + assert intake.call_count == 1 + logs: Final = _delivered_logs(intake) + assert len(logs) == 1 + assert logs[0]["ddsource"] == "litellm" + assert logs[0]["service"] == "litellm-server" + assert logs[0]["status"] == DataDogStatus.INFO + assert TypeAdapter(dict[str, object]).validate_json(logs[0]["message"]) == standard_payload + + +@pytest.mark.asyncio +async def test_a_failed_request_is_delivered_as_an_error_log_that_keeps_the_error_string( + delivery: tuple[DataDogLogger, respx.Route], +) -> None: + datadog_logger, intake = delivery + standard_payload: Final = _standard_logging_payload() + standard_payload["status"] = "failure" + standard_payload["error_str"] = "Test error" + + await datadog_logger.async_log_failure_event( + kwargs={"standard_logging_object": standard_payload}, + response_obj=None, + start_time=STANDARD_START_TIME, + end_time=STANDARD_END_TIME, + ) + await datadog_logger.async_send_batch() + + assert intake.call_count == 1 + logs: Final = _delivered_logs(intake) + assert len(logs) == 1 + assert logs[0]["status"] == DataDogStatus.ERROR + message: Final = TypeAdapter(dict[str, object]).validate_json(logs[0]["message"]) + assert message == standard_payload + assert message["error_str"] == "Test error" + + +@pytest.mark.asyncio +async def test_a_failing_redis_cache_is_delivered_to_datadog_as_redis_warnings( + delivery: tuple[DataDogLogger, respx.Route], monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + datadog_logger, intake = delivery + absent_socket: Final = str(tmp_path / "absent.sock") + redis_cache: Final = Cache(type="redis", url=f"unix://{absent_socket}") + monkeypatch.setattr(redis_cache.cache.service_logger_obj, "dd_logger", datadog_logger, raising=False) + monkeypatch.setattr(litellm, "service_callback", ["datadog"]) + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "cache", redis_cache) + + for _ in range(3): + await litellm.acompletion( + model="gpt-4.1-mini", + messages=[{"role": "user", "content": "what llm are u"}], + mock_response="Accepted", + caching=True, + ) + await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10) + await datadog_logger.async_send_batch() + + assert intake.call_count == 1 + logs: Final = _delivered_logs(intake) + assert len(logs) > 0 + assert {log["status"] for log in logs} == {DataDogStatus.WARN} + messages: Final = [TypeAdapter(_ServiceEventMessage).validate_json(log["message"]) for log in logs] + assert {message["service"] for message in messages} == {"redis"} + assert all(message["is_error"] is True for message in messages) + assert all(absent_socket in message["error"] for message in messages) diff --git a/tests/unit/integrations/focus/test_s3_destination.py b/tests/unit/integrations/focus/test_s3_destination.py index 8e54b561f82..12ddf412ac3 100644 --- a/tests/unit/integrations/focus/test_s3_destination.py +++ b/tests/unit/integrations/focus/test_s3_destination.py @@ -6,6 +6,7 @@ from datetime import datetime, timezone from types import SimpleNamespace from typing import Any, Dict +import boto3 import pytest import litellm.integrations.focus.destinations.s3_destination as s3_module @@ -81,7 +82,7 @@ def test_should_upload_with_configured_client(monkeypatch: pytest.MonkeyPatch): return SimpleNamespace(put_object=put_object) - monkeypatch.setattr(s3_module.boto3, "client", fake_client) + monkeypatch.setattr(boto3, "client", fake_client) dest._upload(content=b"payload", object_key="path/file.bin") diff --git a/tests/unit/integrations/langfuse/test_langfuse_sdk.py b/tests/unit/integrations/langfuse/test_langfuse_sdk.py index 5f01b35a338..cf13667608b 100644 --- a/tests/unit/integrations/langfuse/test_langfuse_sdk.py +++ b/tests/unit/integrations/langfuse/test_langfuse_sdk.py @@ -8,17 +8,19 @@ otherwise record its own duration instead of the call's. import json import logging import threading +import time import uuid from base64 import b64encode from datetime import datetime, timedelta, timezone from time import monotonic, sleep -from types import MappingProxyType +from types import MappingProxyType, SimpleNamespace from typing import Final import httpx import opentelemetry.trace as otel_trace import pytest from langfuse import LangfuseOtelSpanAttributes as A +from langfuse.api.core import http_client as langfuse_http_client from langfuse.api.core.api_error import ApiError from langfuse.api.core.request_options import RequestOptions from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest @@ -1042,28 +1044,35 @@ def test_auth_check_fails_when_the_keys_reach_no_project(): @pytest.mark.parametrize("status", [500, 503, 429], ids=["http-500", "http-503", "http-429"]) -def test_auth_check_and_project_id_make_one_round_trip_when_langfuse_is_down(status): +def test_auth_check_and_project_id_make_one_round_trip_when_langfuse_is_down( + status: int, monkeypatch: pytest.MonkeyPatch +) -> None: """Both run on the event loop; the generated client's default retries sleep for seconds, or for Retry-After.""" - requests: list[httpx.Request] = [] + requests: Final[list[httpx.Request]] = [] + sleeps: Final[list[float]] = [] + monkeypatch.setattr( + langfuse_http_client, + "time", + SimpleNamespace(sleep=sleeps.append, time=time.time), + ) def fail(request: httpx.Request) -> httpx.Response: requests.append(request) return httpx.Response(status, request=request, headers={"retry-after": "20"}, json={"message": "down"}) - client = build_langfuse_client( + client: Final = build_langfuse_client( public_key="pk", secret_key="sk", base_url="http://127.0.0.1:1", httpx_client=httpx.Client(transport=httpx.MockTransport(fail)), ) - started = monotonic() - failure = client.auth_check() + failure: Final = client.auth_check() with pytest.raises(ApiError): client.project_id() assert failure is not None and f"status_code: {status}" in failure.reason assert len(requests) == 2 - assert monotonic() - started < 0.5 + assert sleeps == [], "failed auth checks should not sleep for REST retries" @pytest.mark.parametrize( diff --git a/tests/unit/integrations/otel/test_otel_v2_baggage.py b/tests/unit/integrations/otel/test_otel_v2_baggage.py index 930c01e524e..5855cb60aca 100644 --- a/tests/unit/integrations/otel/test_otel_v2_baggage.py +++ b/tests/unit/integrations/otel/test_otel_v2_baggage.py @@ -7,22 +7,22 @@ import pytest pytest.importorskip("opentelemetry") from litellm.integrations.otel import ( # noqa: E402 - GenAI, HTTP, + GenAI, LiteLLM, OpenTelemetryV2Config, promoted_baggage, ) -from litellm.integrations.otel.plumbing import context as ctx_mod # noqa: E402 -from litellm.integrations.otel.plumbing import providers # noqa: E402 from litellm.integrations.otel.emitter import SpanEmitter # noqa: E402 +from litellm.integrations.otel.model.baggage import BAGGAGE_PROMOTED_KEYS # noqa: E402 from litellm.integrations.otel.model.payloads import ( # noqa: E402 GuardrailSpanData, LLMCallSpanData, ServiceSpanData, ) -from litellm.integrations.otel.model.baggage import BAGGAGE_PROMOTED_KEYS # noqa: E402 from litellm.integrations.otel.model.spans import SpanRole # noqa: E402 +from litellm.integrations.otel.plumbing import context as ctx_mod # noqa: E402 +from litellm.integrations.otel.plumbing import providers # noqa: E402 def _payload(): @@ -63,9 +63,7 @@ def test_identity_promoted_onto_every_span(): root = engine.start_span(SpanRole.PROXY_REQUEST, "POST /chat/completions", ctx) root_ctx = ctx_mod.context_from_span(root, ctx) engine.emit(SpanRole.LLM_CALL, data, parent_context=root_ctx) - engine.emit( - SpanRole.GUARDRAIL, GuardrailSpanData("presidio", status="success"), root_ctx - ) + engine.emit(SpanRole.GUARDRAIL, GuardrailSpanData("presidio", status="success"), root_ctx) engine.emit(SpanRole.SERVICE, ServiceSpanData("redis", call_type="set"), root_ctx) root.end() @@ -161,9 +159,7 @@ def test_allowlisted_metadata_subkey_promoted_blob_excluded(): engine.emit(SpanRole.SERVICE, ServiceSpanData("redis", call_type="set"), ctx) (span,) = exporter.get_finished_spans() # allowlisted metadata sub-key is promoted - assert ( - span.attributes.get(f"{LiteLLM.METADATA_PREFIX}user_api_key_org_id") == "org1" - ) + assert span.attributes.get(f"{LiteLLM.METADATA_PREFIX}user_api_key_org_id") == "org1" # non-allowlisted metadata is NOT promoted (no full-blob dumping) assert all("private_note" not in k for k in span.attributes) @@ -208,6 +204,17 @@ def test_nested_metadata_key_promoted_under_caller_path(): assert not any(k.startswith(f"{LiteLLM.METADATA_PREFIX}requester_metadata") for k in span.attributes) +def test_llm_call_promoted_metadata_strips_requester_prefix_and_uses_allowlist(): + payload = _payload() + payload["metadata"]["requester_metadata"] = {"trace_id": "trace-123"} + data = LLMCallSpanData.from_standard_logging_payload( + payload, + metadata_keys=("requester_metadata.trace_id", "user_api_key_org_id", "missing"), + ) + assert data.promoted_metadata == {"trace_id": "trace-123", "user_api_key_org_id": "org1"} + assert LLMCallSpanData.from_standard_logging_payload(payload).promoted_metadata == {} + + def test_http_attributes_never_promoted(): """Even if http.* is present in baggage, the processor must not stamp it on child spans (it belongs on the SERVER span only).""" @@ -228,9 +235,7 @@ def test_http_attributes_never_promoted(): def test_arbitrary_upstream_baggage_not_promoted(): engine, exporter = _engine_and_exporter() - ctx = ctx_mod.set_request_baggage( - {LiteLLM.TEAM_ID: "t1", "some.upstream.key": "leak"} - ) + ctx = ctx_mod.set_request_baggage({LiteLLM.TEAM_ID: "t1", "some.upstream.key": "leak"}) engine.emit(SpanRole.SERVICE, ServiceSpanData("redis", call_type="set"), ctx) (span,) = exporter.get_finished_spans() assert span.attributes.get(LiteLLM.TEAM_ID) == "t1" diff --git a/tests/unit/integrations/otel/test_otel_v2_emitter.py b/tests/unit/integrations/otel/test_otel_v2_emitter.py index 11b2aa5fd67..949eb97bccd 100644 --- a/tests/unit/integrations/otel/test_otel_v2_emitter.py +++ b/tests/unit/integrations/otel/test_otel_v2_emitter.py @@ -20,7 +20,7 @@ from litellm.integrations.otel import ( # noqa: E402 ) from litellm.integrations.otel.plumbing import context as ctx_mod # noqa: E402 from litellm.integrations.otel.plumbing import providers # noqa: E402 -from litellm.integrations.otel.emitter import SpanEmitter, span_attribute_limit # noqa: E402 +from litellm.integrations.otel.emitter import SpanEmitter, attribute_budget, span_attribute_limit # noqa: E402 from litellm.integrations.otel.emitter import stamp_error # noqa: E402 from litellm.integrations.otel.mappers.utils import MAX_TOOL_DEFINITION_ATTRS_PER_SPAN # noqa: E402 from litellm.integrations.otel.model.payloads import ( # noqa: E402 @@ -558,6 +558,113 @@ def test_prompt_turns_are_shed_before_response_choices(): ) +def test_metadata_blob_keeps_the_indexed_messages_it_competes_with(): + """The promoted ``metadata`` blob rides alongside the indexed messages on a boundary-opened span. + + The live regression shape: the LLM-call span is opened at the request + boundary and already carries a few stamped attributes, most of which the + mappers re-emit under the same keys. Charging those overwritten keys + against the attribute budget reserved slots the fit could never spend, so + the new ``metadata`` key displaced a whole middle message group — an + eviction invisible to the SDK's dropped-attributes counter because the fit + sheds before ``span.set_attribute``. With the budget counting only keys + the fit does not overwrite, the metadata key rides and the span keeps the + indexed-message count of the metadata-free span (at shedding granularity). + """ + cfg = OpenTelemetryV2Config( + exporter="in_memory", + legacy_compat=False, + mapper_names=["genai", "openinference"], + capture_message_content="span_only", + ) + provider, exporter = providers.in_memory_provider(cfg) + engine = SpanEmitter(providers.get_tracer(provider, "litellm-test"), cfg) + + def boundary_span(payload): + span = engine.start_span(SpanRole.LLM_CALL, "chat gpt-4o") + # what the boundary opener stamps before the typed payload exists + span.set_attribute(GenAI.REQUEST_MODEL, "gpt-4o") + span.set_attribute(LiteLLM.PROVIDER_MODEL, "gpt-4o-2024") + span.set_attribute("litellm.metadata.user_api_key_alias", "edge-key") + engine.finish_span( + SpanRole.LLM_CALL, + span, + LLMCallSpanData.from_standard_logging_payload( + payload, capture_content=True, metadata_keys=("user_api_key_alias",) + ), + ) + (finished,) = exporter.get_finished_spans() + exporter.clear() + return finished + + with_metadata = boundary_span(_conversation_payload(47, metadata={"user_api_key_alias": "edge-key"})) + without_metadata = boundary_span(_conversation_payload(47)) + _assert_core_intact(with_metadata) + a = with_metadata.attributes + + assert json.loads(a["metadata"]) == {"user_api_key_alias": "edge-key"} + assert "metadata" not in without_metadata.attributes + kept = _indexed_messages(a, "llm.input_messages") + baseline = _indexed_messages(without_metadata.attributes, "llm.input_messages") + assert kept == baseline or kept == baseline[:-1] + assert kept[0] == 0 and kept[-1] == 46 + assert len(a) <= SpanLimits().max_span_attributes + + +def test_preset_message_key_the_fit_sheds_still_never_overflows_the_span(): + """A pre-set indexed-message key the fit later sheds cannot push the span over its limit. + + The budget treats every mapped key already on the span as an overwrite + (free). If the fit then sheds that key, the value stamped earlier simply + stays in its slot, so the span holds one entry for it either way and the + total never exceeds the limit — the SDK's dropped-attributes counter stays + at zero. + """ + cfg = OpenTelemetryV2Config( + exporter="in_memory", + legacy_compat=False, + mapper_names=["genai", "openinference"], + capture_message_content="span_only", + ) + provider, exporter = providers.in_memory_provider(cfg) + engine = SpanEmitter(providers.get_tracer(provider, "litellm-test"), cfg) + span = engine.start_span(SpanRole.LLM_CALL, "chat gpt-4o") + span.set_attribute("llm.input_messages.1.message.role", "stale-role") + engine.finish_span( + SpanRole.LLM_CALL, + span, + LLMCallSpanData.from_standard_logging_payload( + _conversation_payload(60), capture_content=True, metadata_keys=("user_api_key_alias",) + ), + ) + (finished,) = exporter.get_finished_spans() + _assert_core_intact(finished) + assert len(finished.attributes) <= SpanLimits().max_span_attributes + + +def test_attribute_budget_counts_only_keys_the_fit_does_not_overwrite(): + """Pre-set attributes the mapped set overwrites consume no slot against the span limit. + + A boundary-opened LLM-call span already carries a few stamped attributes, + most of which the mappers re-emit under the same keys. Charging those + against the budget reserves slots the fit can never spend and sheds + indexed message attributes for nothing. + """ + cfg = OpenTelemetryV2Config(exporter="in_memory") + provider, _exporter = providers.in_memory_provider(cfg) + span = providers.get_tracer(provider, "litellm-test").start_span("s") + span.set_attribute("gen_ai.request.model", "gpt-4o") + span.set_attribute("litellm.metadata.user_api_key_alias", "edge-key") + try: + assert attribute_budget(span, 0) == SpanLimits().max_span_attributes - 2 + assert ( + attribute_budget(span, 0, frozenset({"gen_ai.request.model"})) + == SpanLimits().max_span_attributes - 1 + ) + finally: + span.end() + + def test_indexed_messages_respect_a_lower_span_attribute_count_limit(monkeypatch): """The budget follows the SDK's configured limit, not a hardcoded default.""" monkeypatch.setenv("OTEL_SPAN_ATTRIBUTE_COUNT_LIMIT", "48") diff --git a/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py b/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py index cdff9c960f3..ca7a9b424ad 100644 --- a/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py +++ b/tests/unit/integrations/otel/test_otel_v2_vendor_mappers.py @@ -7,6 +7,7 @@ backends, so one trace lights up every configured destination. import json from collections.abc import Mapping +from itertools import chain from typing import Final import pytest @@ -19,6 +20,7 @@ from litellm.integrations.otel.mappers import ( WeaveMapper, resolve_mappers, ) +from litellm.integrations.otel.mappers.openinference import fit_indexed_messages from litellm.integrations.otel.model.payloads import ( EmbeddingOutput, LLMCallSpanData, @@ -29,6 +31,7 @@ from litellm.integrations.otel.model.payloads import ( ToolDefinition, ) from litellm.integrations.otel.model.trace_controls import TraceControls +from tests.unit.integrations.otel.test_otel_v2_sources_of_truth import _responses_payload def _llm_call(**overrides): @@ -118,6 +121,378 @@ def test_openinference_multimodal_content_text_only(): assert attrs["llm.input_messages.0.message.content"] == "hi there" +def test_openinference_output_tool_calls_preserve_calls_in_attributes_and_value(): + tool_calls: Final = [ + { + "id": "call_paris", + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + }, + { + "id": "call_search", + "function": {"name": "search", "arguments": {"q": 1}}, + "index": 0, + }, + "ignored", + ] + data: Final = _llm_call( + choices_out=( + { + "finish_reason": "tool_calls", + "message": {"role": "assistant", "content": None, "tool_calls": tool_calls}, + }, + ) + ) + attrs: Final = OpenInferenceMapper().map(data) + assert {key: value for key, value in attrs.items() if ".tool_calls." in key} == { + "llm.output_messages.0.message.tool_calls.0.tool_call.id": "call_paris", + "llm.output_messages.0.message.tool_calls.0.tool_call.function.name": "lookup_weather", + "llm.output_messages.0.message.tool_calls.0.tool_call.function.arguments": '{"city": "Paris"}', + "llm.output_messages.0.message.tool_calls.1.tool_call.id": "call_search", + "llm.output_messages.0.message.tool_calls.1.tool_call.function.name": "search", + "llm.output_messages.0.message.tool_calls.1.tool_call.function.arguments": '{"q": 1}', + } + assert json.loads(attrs["output.value"]) == [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_paris", + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + }, + { + "id": "call_search", + "type": "function", + "function": {"name": "search", "arguments": '{"q": 1}'}, + }, + ], + } + ] + + +def test_openinference_responses_tool_calls_are_emitted_as_output_attributes(): + data: Final = LLMCallSpanData.from_standard_logging_payload( + _responses_payload( + [ + { + "type": "function_call", + "call_id": "call_resp", + "name": "lookup_weather", + "arguments": '{"city": "Paris"}', + } + ] + ), + capture_content=True, + ) + attrs: Final = OpenInferenceMapper().map(data) + tool_call: Final = "llm.output_messages.0.message.tool_calls.0.tool_call." + + assert {key: value for key, value in attrs.items() if ".tool_calls." in key} == { + tool_call + "id": "call_resp", + tool_call + "function.name": "lookup_weather", + tool_call + "function.arguments": '{"city": "Paris"}', + } + assert json.loads(attrs["output.value"]) == [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_resp", + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + } + ], + } + ] + + +def test_openinference_input_tool_calls_stay_in_value_only(): + data: Final = _llm_call( + messages_in=( + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_weather", + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + } + ], + }, + ) + ) + attrs: Final = OpenInferenceMapper().map(data) + assert all(".tool_calls." not in key for key in attrs) + assert json.loads(attrs["input.value"]) == [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_weather", + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + } + ], + } + ] + + +def test_openinference_output_tool_calls_do_not_shed_input_roles_under_budget(): + message_groups: Final = tuple( + ( + {"role": "user", "content": f"Question {index}"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": f"call_{index}", + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + } + ], + }, + {"role": "tool", "tool_call_id": f"call_{index}", "content": f"Result {index}"}, + ) + for index in range(13) + ) + messages_in: Final = tuple(chain.from_iterable(message_groups)) + ({"role": "user", "content": "Final request"},) + tools: Final = tuple( + ToolDefinition(name=name, description="Tool", parameters_json='{"type":"object"}') + for name in ("lookup_weather", "search", "get_location", "convert_units") + ) + data: Final = _llm_call( + messages_in=messages_in, + tools=tools, + choices_out=( + { + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_output", + "type": "function", + "function": {"name": "lookup_weather", "arguments": '{"city": "Paris"}'}, + } + ], + }, + }, + ), + ) + attrs: Final = fit_indexed_messages(OpenInferenceMapper().map(data), 128) + tool_call: Final = "llm.output_messages.0.message.tool_calls.0.tool_call." + + assert { + f"llm.input_messages.{index}.message.role": attrs.get(f"llm.input_messages.{index}.message.role") + for index in range(40) + } == {f"llm.input_messages.{index}.message.role": message["role"] for index, message in enumerate(messages_in)} + assert {key: value for key, value in attrs.items() if ".tool_calls." in key} == { + tool_call + "id": "call_output", + tool_call + "function.name": "lookup_weather", + tool_call + "function.arguments": '{"city": "Paris"}', + } + + +def test_openinference_budget_sheds_trailing_output_tool_calls_before_the_message(): + tool_calls: Final = tuple( + { + "id": f"call_{index}", + "type": "function", + "function": { + "name": "lookup_weather", + "arguments": f'{{"city": "C{index}"}}', + }, + } + for index in range(60) + ) + data: Final = _llm_call( + choices_out=( + { + "finish_reason": "tool_calls", + "message": {"role": "assistant", "content": None, "tool_calls": tool_calls}, + }, + ) + ) + mapped: Final = OpenInferenceMapper().map(data) + attrs: Final = fit_indexed_messages(mapped, len(mapped) - 34) + retained_tool_call_keys: Final = tuple(key for key in attrs if ".tool_calls." in key) + retained_tool_call_indices: Final = frozenset(int(key.split(".")[5]) for key in retained_tool_call_keys) + + assert retained_tool_call_indices == frozenset(range(50)) + assert len(retained_tool_call_keys) == 150 + assert attrs["llm.output_messages.0.message.role"] == "assistant" + assert not any(key.startswith("llm.input_messages.0.") for key in attrs) + assert len(attrs) == len(mapped) - 34 + assert json.loads(attrs["output.value"]) == [{"role": "assistant", "content": None, "tool_calls": list(tool_calls)}] + + +def test_openinference_plain_output_messages_keep_the_existing_value_shape(): + attrs: Final = OpenInferenceMapper().map(_llm_call()) + assert all(".tool_calls." not in key for key in attrs) + assert json.loads(attrs["output.value"]) == [{"role": "assistant", "content": "Sunny."}] + + +def test_openinference_metadata_contains_only_promoted_metadata(): + attrs: Final = OpenInferenceMapper().map( + _llm_call(promoted_metadata={"trace_marker": "m", "user_api_key_alias": "k"}) + ) + assert json.loads(attrs["metadata"]) == {"trace_marker": "m", "user_api_key_alias": "k"} + assert "metadata" not in OpenInferenceMapper().map(_llm_call()) + + +def _long_prompt_messages(count: int) -> tuple[dict[str, str], ...]: + return tuple({"role": "user" if i == 0 else "assistant", "content": f"turn {i}"} for i in range(count)) + + +def _input_message_keys(attrs: Mapping[str, object]) -> frozenset[str]: + return frozenset(key for key in attrs if key.startswith("llm.input_messages.")) + + +def test_openinference_metadata_does_not_displace_indexed_messages_under_budget(): + """The ``metadata`` blob rides in reclaimed slots instead of evicting messages. + + ``metadata`` competes for the OTel 128-attribute span budget, and shedding + happens before ``span.set_attribute``, so the SDK's dropped-attributes + counter stays at zero — the eviction is invisible. The fit pins the + metadata key behind every message group: a squeezed span keeps the message + attributes it would keep without the metadata blob (at whole-group + granularity, at most one group of headroom difference), and the metadata + key survives alongside them. + """ + messages: Final = _long_prompt_messages(62) + + def fitted(promoted: Mapping[str, str], budget: int | None = None): + mapped: Final = OpenInferenceMapper().map(_llm_call(messages_in=messages, promoted_metadata=promoted)) + assert len(mapped) > 128, "fixture must squeeze the span attribute budget" + return fit_indexed_messages(mapped, budget if budget is not None else len(mapped) - 5) + + without_metadata: Final = fitted({}) + with_metadata: Final = fitted({"user_api_key_alias": "edge-key"}) + assert "metadata" not in without_metadata + assert json.loads(with_metadata["metadata"]) == {"user_api_key_alias": "edge-key"} + assert _input_message_keys(with_metadata) == _input_message_keys(without_metadata) + + # Against the absolute span limit the displacement is bounded by the + # whole-group shedding granularity: at most one message group. + without_at_limit: Final = fitted({}, 128) + with_at_limit: Final = fitted({"user_api_key_alias": "edge-key"}, 128) + assert len(_input_message_keys(with_at_limit)) >= len(_input_message_keys(without_at_limit)) - 2 + assert "metadata" in with_at_limit + + +def test_openinference_metadata_sheds_only_after_every_indexed_message(): + """Metadata sheds last: only once every indexed message attribute is gone.""" + mapped: Final = OpenInferenceMapper().map( + _llm_call(messages_in=_long_prompt_messages(6), promoted_metadata={"user_api_key_alias": "edge-key"}) + ) + message_keys: Final = frozenset(key for key in mapped if ".message." in key) + # Budget too small for the message family alone: everything indexed goes, + # and the metadata blob absorbs the residual shortfall with it. + starved: Final = fit_indexed_messages(mapped, len(mapped) - len(message_keys) - 1) + assert not any(key in starved for key in message_keys) + assert "metadata" not in starved + # One slot more and metadata survives alongside zero indexed messages. + last_standing: Final = fit_indexed_messages(mapped, len(mapped) - len(message_keys)) + assert not any(key in last_standing for key in message_keys) + assert "metadata" in last_standing + + +def test_openinference_raw_tool_arguments_fall_back_to_repr_instead_of_raising(): + """Malformed Python tool arguments must not lose the span. + + A plain ``Function()`` constructor JSON-serializes arguments, but provider + adapters and ``model_construct`` responses hand over raw Python objects — + tuple-keyed dicts and cycles that ``json.dumps`` raises on. The mapper + serializes them with a ``repr`` fallback instead of letting the exception + escape before the span is exported. + """ + # rebind-ok: a self-referencing dict cannot be built in one shot — the cycle + # only exists once the finished dict is inserted into itself. + circular: dict[str, object] = {} + circular["self"] = circular + for label, raw_arguments in (("tuple-key", {(1, 2): "v"}), ("circular", circular)): + data: Final = _llm_call( + choices_out=( + { + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "tc1", "type": "function", "function": {"name": "f", "arguments": raw_arguments}} + ], + }, + }, + ) + ) + attrs: Final = OpenInferenceMapper().map(data) + assert attrs["llm.output_messages.0.message.tool_calls.0.tool_call.function.arguments"] == repr( + raw_arguments + ), label + assert attrs["llm.output_messages.0.message.role"] == "assistant" + + +def test_openinference_payload_tool_arguments_with_raw_objects_map_without_raising(): + """The standard-logging payload path (provider adapter responses) survives raw argument objects too.""" + payload: Final = { + "call_type": "acompletion", + "custom_llm_provider": "openai", + "model": "gpt-4o", + "prompt_tokens": 3, + "completion_tokens": 2, + "total_tokens": 5, + "stream": False, + "model_parameters": {}, + "response": { + "id": "resp_bad", + "model": "gpt-4o-2024", + "choices": [ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": "hi", + "tool_calls": [ + {"id": "tc1", "type": "function", "function": {"name": "f", "arguments": {(1, 2): "v"}}} + ], + }, + } + ], + }, + "metadata": {"user_api_key_alias": "edge-key"}, + "status": "success", + "litellm_call_id": "call_raw_args", + "hidden_params": {}, + } + data: Final = LLMCallSpanData.from_standard_logging_payload( + payload, capture_content=True, metadata_keys=("user_api_key_alias",) + ) + attrs: Final = OpenInferenceMapper().map(data) + arguments: Final = attrs["llm.output_messages.0.message.tool_calls.0.tool_call.function.arguments"] + assert arguments == repr({(1, 2): "v"}) + assert json.loads(attrs["metadata"]) == {"user_api_key_alias": "edge-key"} + + +def test_openinference_attribute_fit_keeps_pinned_messages_on_long_prompts(): + messages: Final = _long_prompt_messages(4000) + mapped: Final = OpenInferenceMapper().map( + _llm_call(messages_in=messages, promoted_metadata={"user_api_key_alias": "k"}) + ) + fitted: Final = fit_indexed_messages(mapped, 128) + + assert len(fitted) <= 128 + assert fitted["llm.input_messages.0.message.role"] == "user" + assert fitted["llm.input_messages.3999.message.role"] == "assistant" + + # --------------------------------------------------------------------------- # # Langfuse # --------------------------------------------------------------------------- # diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py index 8392cef8a7b..d96f00cc2b1 100644 --- a/tests/unit/integrations/test_anthropic_cache_control_hook.py +++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py @@ -14,6 +14,7 @@ from pydantic import BaseModel, ConfigDict import litellm from litellm.integrations.anthropic_cache_control_hook import ( AnthropicCacheControlHook, + configured_injection_points, supports_openai_prompt_cache_breakpoint, ) from litellm.litellm_core_utils.prompt_templates.factory import ( @@ -3738,6 +3739,14 @@ class TestPromptCacheBreakpointCapability: ) assert supports_openai_prompt_cache_breakpoint("gpt-5.6") is False + @pytest.mark.parametrize("model", ["gpt-4.1", "gpt-5.6"]) + @pytest.mark.parametrize("flag", ["true", 1, "false", 0]) + def test_listed_model_with_an_odd_typed_flag_is_not_eligible(self, monkeypatch, model, flag): + monkeypatch.setitem( + litellm.model_cost, model, {**litellm.model_cost[model], "supports_prompt_cache_breakpoint": flag} + ) + assert supports_openai_prompt_cache_breakpoint(model) is False + def test_published_map_without_the_flag_still_injects_on_gpt_5_6(self, monkeypatch): unflagged = {k: v for k, v in litellm.model_cost["gpt-5.6"].items() if k != "supports_prompt_cache_breakpoint"} @@ -3769,6 +3778,277 @@ class TestPromptCacheBreakpointCapability: assert model not in litellm.model_cost assert supports_openai_prompt_cache_breakpoint(model) is expected + def test_a_null_prompt_cache_options_takes_the_implicit_default_on_both_paths(self): + points = [{"location": "message", "role": "system"}] + + _, _, chat_params = AnthropicCacheControlHook().get_chat_completion_prompt( + model="openai/gpt-5.6", + messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], + non_default_params={"cache_control_injection_points": copy.deepcopy(points), "prompt_cache_options": None}, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert chat_params["prompt_cache_options"] == {"mode": "implicit"} + + kwargs = {"cache_control_injection_points": copy.deepcopy(points), "prompt_cache_options": None} + AnthropicCacheControlHook.maybe_inject_cache_control( + [{"role": "user", "content": "hi"}], "sys", kwargs, model="gpt-5.6", custom_llm_provider="openai" + ) + assert kwargs["prompt_cache_options"] == {"mode": "implicit"} + + +class TestHostedOpenAIDialectFlag: + """#38666: an OpenAI-shaped model served by another provider can opt in through its own + model-map entry, instead of being excluded by the openai-only provider check.""" + + MANTLE_MODEL = "bedrock_mantle/openai.gpt-5.6-sol" + + def _register(self, monkeypatch, key, provider, flag=True, **extra): + entry = {"litellm_provider": provider, "mode": "chat", **extra} + if flag is not None: + entry["supports_prompt_cache_breakpoint"] = flag + monkeypatch.setitem(litellm.model_cost, key, entry) + + def test_flagged_non_openai_deployment_is_eligible(self, monkeypatch): + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "bedrock_mantle") + is True + ) + + def test_bedrock_api_base_does_not_veto_the_explicit_flag(self, monkeypatch): + """The api_base check exists to sniff for api.openai.com, which a Bedrock host never is.""" + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint( + self.MANTLE_MODEL, + "bedrock_mantle", + api_base="https://bedrock-runtime.us-east-1.amazonaws.com", + ) + is True + ) + + def test_flag_set_false_keeps_the_deployment_ineligible(self, monkeypatch): + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle", flag=False) + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "bedrock_mantle") + is False + ) + + @pytest.mark.parametrize("flag", ["true", 1, "false", 0]) + def test_an_odd_typed_flag_keeps_the_deployment_ineligible(self, monkeypatch, flag): + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle", flag=flag) + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "bedrock_mantle") + is False + ) + + def test_unflagged_non_openai_deployment_stays_ineligible(self, monkeypatch): + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle", flag=None) + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "bedrock_mantle") + is False + ) + + def test_entry_provider_must_match_the_request_provider(self, monkeypatch): + """A flagged entry does not license a different provider serving the same model string.""" + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "azure") is False + ) + + def test_openai_entries_still_go_through_the_api_base_check(self, monkeypatch): + """gpt-5.6 is flagged and openai-provided, so it must not bypass the host gate.""" + assert litellm.model_cost["gpt-5.6"]["supports_prompt_cache_breakpoint"] is True + assert litellm.model_cost["gpt-5.6"]["litellm_provider"] == "openai" + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint( + "gpt-5.6", "openai", api_base="https://some-compatible-host.example.com" + ) + is False + ) + + def test_azure_hosted_gpt_5_6_remains_ineligible(self): + """Regression guard: the openai entry's flag must not leak to another provider.""" + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("gpt-5.6", "azure") is False + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("azure/gpt-5.6", None) is False + + REGIONAL_MODEL = "bedrock_mantle/us-east-1/openai.gpt-5.6-sol" + + def test_region_prefixed_deployment_reads_its_region_free_entry(self, monkeypatch): + """``bedrock_mantle//`` is a documented routing form the map keys without the region.""" + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.REGIONAL_MODEL, "bedrock_mantle") + is True + ) + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.REGIONAL_MODEL, None) is True + + def test_region_prefixed_deployment_honors_a_flag_set_false(self, monkeypatch): + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle", flag=False) + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.REGIONAL_MODEL, "bedrock_mantle") + is False + ) + + def test_region_prefixed_entry_outranks_the_region_free_one(self, monkeypatch): + """A row keyed with the region states that deployment's own dialect; GovCloud rows carry no flag.""" + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + gov_model = "bedrock_mantle/us-gov-west-1/openai.gpt-5.6-sol" + self._register(monkeypatch, gov_model, "bedrock_mantle", flag=None) + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(gov_model, "bedrock_mantle") is False + + def test_region_free_entry_does_not_license_another_provider(self, monkeypatch): + """The candidate keys are built for the request's provider, so a flagged Mantle row stays Mantle's.""" + self._register(monkeypatch, self.MANTLE_MODEL, "bedrock_mantle") + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("us-east-1/openai.gpt-5.6-sol", "azure") + is False + ) + + def test_unmapped_deployment_name_stays_ineligible(self): + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("azure/my-deployment", None) is False + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("my-deployment", "azure") is False + + def test_bare_name_colliding_with_an_openai_row_reads_its_own_provider_entry(self, monkeypatch): + """The Responses layer hands the hook a bare deployment name plus its provider. The openai row keyed by + that bare name neither answers for the deployment nor stops the lookup of the provider's own entry.""" + self._register(monkeypatch, "gpt-collide", "openai") + self._register(monkeypatch, "azure_ai/gpt-collide", "azure_ai") + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("gpt-collide", "azure_ai") is True + + def test_bare_name_colliding_with_an_openai_row_stays_ineligible_without_its_own_flag(self, monkeypatch): + self._register(monkeypatch, "gpt-collide", "openai") + self._register(monkeypatch, "azure_ai/gpt-collide", "azure_ai", flag=None) + assert AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint("gpt-collide", "azure_ai") is False + + +class TestBedrockMantleGptShipsTheOpenAIDialect: + """The shipped cost map flags Bedrock Mantle's GPT-5.6 and newer OpenAI rows, so a configured injection point + on one of them reaches the wire as prompt_cache_breakpoint instead of an Anthropic cache_control the Mantle + bridge strips (verified live against bedrock-mantle.us-east-1 on 2026-10-07: cache_write_tokens then + cached_tokens on the repeat call).""" + + MANTLE_MODEL = "bedrock_mantle/openai.gpt-5.6-sol" + + @pytest.fixture(autouse=True) + def _bundled_model_map(self, monkeypatch): + bundled = os.path.join(os.path.dirname(litellm.__file__), "model_prices_and_context_window_backup.json") + with open(bundled) as handle: + monkeypatch.setattr(litellm, "model_cost", json.load(handle)) + litellm.utils.cached_get_model_info_helper.cache_clear() + yield + litellm.utils.cached_get_model_info_helper.cache_clear() + + def test_shipped_entry_makes_the_deployment_eligible(self): + assert supports_openai_prompt_cache_breakpoint(self.MANTLE_MODEL) is True + assert ( + AnthropicCacheControlHook._targets_openai_prompt_cache_breakpoint(self.MANTLE_MODEL, "bedrock_mantle") + is True + ) + + def test_every_flagged_mantle_row_is_an_openai_gpt_5_6_or_newer_model(self): + flagged = { + key for key, entry in litellm.model_cost.items() + if key.startswith("bedrock_mantle/") and entry.get("supports_prompt_cache_breakpoint") is True + } + assert self.MANTLE_MODEL in flagged + for key in flagged: + bare = key.rsplit("/", 1)[-1].removeprefix("openai.") + assert supports_openai_prompt_cache_breakpoint(bare) is True, key + + def test_seeding_stamps_the_openai_dialect_for_a_configured_point(self): + non_default_params = {"cache_control_injection_points": [{"location": "message", "role": "system"}]} + AnthropicCacheControlHook.maybe_seed_default_injection_points( + non_default_params=non_default_params, + messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], + model=self.MANTLE_MODEL, + custom_llm_provider=None, + ) + assert non_default_params["cache_control_injection_points"][0]["_litellm_openai_dialect"] is True + + def test_configured_point_emits_the_openai_marker_and_default_options(self): + _, messages, params = AnthropicCacheControlHook().get_chat_completion_prompt( + model=self.MANTLE_MODEL, + messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], + non_default_params={"cache_control_injection_points": [{"location": "message", "role": "system"}]}, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert messages[0]["content"] == [{"type": "text", "text": "sys", "prompt_cache_breakpoint": {"mode": "explicit"}}] + assert params["prompt_cache_options"] == {"mode": "implicit"} + assert AnthropicCacheControlHook.count_request_cache_breakpoints(messages) == 1 + + def test_region_prefixed_deployment_emits_the_openai_marker(self): + """The region-prefixed routing form documented for Mantle lands on the same shipped row.""" + _, messages, params = AnthropicCacheControlHook().get_chat_completion_prompt( + model="bedrock_mantle/us-east-1/openai.gpt-5.6-sol", + messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], + non_default_params={"cache_control_injection_points": [{"location": "message", "role": "system"}]}, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert messages[0]["content"] == [{"type": "text", "text": "sys", "prompt_cache_breakpoint": {"mode": "explicit"}}] + assert params["prompt_cache_options"] == {"mode": "implicit"} + + BARE_MODEL = "openai.gpt-5.6-sol" + POINTS = [{"location": "message", "role": "system"}] + MESSAGES = [{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}] + MARKED_SYSTEM = [{"type": "text", "text": "sys", "prompt_cache_breakpoint": {"mode": "explicit"}}] + + def test_a_bare_deployment_name_with_its_provider_is_stamped_on_the_openai_dialect(self): + """A deployment written as ``model: openai.gpt-5.6-sol`` plus ``custom_llm_provider: bedrock_mantle`` + has no row of its own and no openai row of the same name, so only the provider-keyed row can + answer; the stamp must read it the way the dialect resolution does.""" + stamped = AnthropicCacheControlHook._stamped_with_dialect( + copy.deepcopy(self.POINTS), self.BARE_MODEL, "bedrock_mantle", None, None + ) + assert stamped[0]["_litellm_openai_dialect"] is True + + def test_a_bare_deployment_name_without_its_provider_keeps_its_points_and_costs_no_lookup(self): + points = copy.deepcopy(self.POINTS) + with patch.object(AnthropicCacheControlHook, "_resolve_provider") as resolve: + assert AnthropicCacheControlHook._stamped_with_dialect(points, self.BARE_MODEL, None, None, None) is points + resolve.assert_not_called() + + def test_the_chat_seed_carries_the_resolved_provider_for_a_bare_deployment_name(self): + params: dict = {"cache_control_injection_points": copy.deepcopy(self.POINTS)} + AnthropicCacheControlHook.maybe_seed_default_injection_points( + non_default_params=params, + messages=copy.deepcopy(self.MESSAGES), + model=self.BARE_MODEL, + custom_llm_provider="bedrock_mantle", + ) + _, messages, out = AnthropicCacheControlHook().get_chat_completion_prompt( + model=self.BARE_MODEL, + messages=copy.deepcopy(self.MESSAGES), + non_default_params=params, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert messages[0]["content"] == self.MARKED_SYSTEM + assert out["prompt_cache_options"] == {"mode": "implicit"} + + def test_the_responses_stamp_carries_the_resolved_provider_for_a_bare_deployment_name(self): + from litellm.responses.main import _stamp_injection_points_with_dialect + + kwargs: dict = {"cache_control_injection_points": copy.deepcopy(self.POINTS)} + _stamp_injection_points_with_dialect(kwargs, self.BARE_MODEL, "bedrock_mantle") + _, messages, out = AnthropicCacheControlHook().get_chat_completion_prompt( + model=self.BARE_MODEL, + messages=copy.deepcopy(self.MESSAGES), + non_default_params=kwargs, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert messages[0]["content"] == self.MARKED_SYSTEM + assert out["prompt_cache_options"] == {"mode": "implicit"} + class TestRecordGatewayInjection: """The injection marker spend accounting gates prompt-caching savings on.""" @@ -3901,3 +4181,83 @@ class TestRecordGatewayInjection: custom_llm_provider="anthropic", ) assert self.KEY not in kwargs["litellm_metadata"] + + +class TestMalformedInjectionPointsAreIgnored: + """A ``cache_control_injection_points`` value that is not a list of points (a string, an int, a bare + dict, a list of strings) raised inside the hook and turned every request to that deployment into a 500. + Every entry point now reads it as no configured points, the way ``null`` already read.""" + + SHAPES = ("system", 5, {"location": "message", "role": "system"}, ["system"], None) + MIXED = ["system", {"location": "message", "role": "system"}, 3] + MESSAGES = [{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}] + + @pytest.mark.parametrize("value", SHAPES) + def test_reads_as_no_points(self, value): + assert configured_injection_points(value) == () + + def test_keeps_the_point_entries_of_a_mixed_list(self): + assert configured_injection_points(self.MIXED) == ({"location": "message", "role": "system"},) + + def test_the_point_beside_junk_entries_survives_the_chat_seed_on_an_unstamped_deployment(self): + params: dict = {"cache_control_injection_points": copy.deepcopy(self.MIXED)} + AnthropicCacheControlHook.maybe_seed_default_injection_points( + non_default_params=params, + messages=copy.deepcopy(self.MESSAGES), + model="claude-sonnet-4-5", + custom_llm_provider="anthropic", + ) + _, processed, _ = AnthropicCacheControlHook().get_chat_completion_prompt( + model="claude-sonnet-4-5", + messages=copy.deepcopy(self.MESSAGES), + non_default_params=params, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert processed[0] == {"role": "system", "content": "sys", "cache_control": {"type": "ephemeral"}} + + @pytest.mark.parametrize("value", SHAPES) + def test_chat_prompt_hook_leaves_the_request_untouched(self, value): + _, processed, params = AnthropicCacheControlHook().get_chat_completion_prompt( + model="openai/gpt-5.6", + messages=copy.deepcopy(self.MESSAGES), + non_default_params={"cache_control_injection_points": copy.deepcopy(value)}, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + assert processed == self.MESSAGES + assert "cache_control_injection_points" not in params + assert "prompt_cache_options" not in params + + @pytest.mark.parametrize("value", SHAPES) + def test_chat_seeding_falls_through_to_the_defaults(self, monkeypatch, value): + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + params: dict = {"cache_control_injection_points": copy.deepcopy(value)} + AnthropicCacheControlHook.maybe_seed_default_injection_points( + non_default_params=params, + messages=copy.deepcopy(self.MESSAGES), + model="claude-sonnet-4-5", + custom_llm_provider="anthropic", + ) + seeded = params["cache_control_injection_points"] + assert seeded and all(isinstance(point, dict) and "location" in point for point in seeded) + + @pytest.mark.parametrize("value", SHAPES) + def test_messages_path_leaves_the_request_untouched(self, value): + kwargs: dict = {"cache_control_injection_points": copy.deepcopy(value)} + messages, system = AnthropicCacheControlHook.maybe_inject_cache_control( + [{"role": "user", "content": "hi"}], "sys", kwargs, model="gpt-5.6", custom_llm_provider="openai" + ) + assert (messages, system) == ([{"role": "user", "content": "hi"}], "sys") + assert "cache_control_injection_points" not in kwargs + assert "prompt_cache_options" not in kwargs + + @pytest.mark.parametrize("value", SHAPES) + def test_responses_dialect_stamp_leaves_the_request_untouched(self, value): + from litellm.responses.main import _stamp_injection_points_with_dialect + + kwargs: dict = {"cache_control_injection_points": copy.deepcopy(value)} + _stamp_injection_points_with_dialect(kwargs, "gpt-5.6", "openai") + assert kwargs == {"cache_control_injection_points": value} diff --git a/tests/unit/integrations/test_custom_guardrail.py b/tests/unit/integrations/test_custom_guardrail.py index 72c362425bb..e04f8af5a66 100644 --- a/tests/unit/integrations/test_custom_guardrail.py +++ b/tests/unit/integrations/test_custom_guardrail.py @@ -1,5 +1,9 @@ import asyncio +import copy import datetime as dt +import json +import pickle +import threading from typing import TYPE_CHECKING, ClassVar, Final, Literal, Optional from unittest.mock import AsyncMock @@ -8,7 +12,10 @@ import pytest from litellm.integrations.custom_guardrail import ( DEFAULT_ADVISORY_MESSAGE, CustomGuardrail, + _request_is_streaming, + guardrail_request_data_with_streaming, log_guardrail_information, + without_server_streaming_classification, ) from litellm.litellm_core_utils.litellm_logging import Logging from litellm.proxy._types import CallTypes, UserAPIKeyAuth @@ -531,6 +538,233 @@ class TestCustomGuardrailShouldRunGuardrail: assert always_on.should_run_guardrail(data=forged, event_type=GuardrailEventHooks.pre_call) is True +_STREAM_SCOPE_HOOKS: Final = ( + GuardrailEventHooks.pre_call, + GuardrailEventHooks.during_call, + GuardrailEventHooks.post_call, +) + + +class TestCustomGuardrailStreamScope: + @pytest.mark.parametrize("event_type", _STREAM_SCOPE_HOOKS) + @pytest.mark.parametrize("stream", [True, False]) + @pytest.mark.parametrize("stream_scope", [None, "both"]) + def test_both_and_omitted_run_on_streaming_and_non_streaming( + self, + event_type: GuardrailEventHooks, + stream: bool, + stream_scope: str | None, + ): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=event_type, + stream_scope=stream_scope, + ) + assert guardrail.should_run_guardrail({"stream": stream}, event_type) is True + + @pytest.mark.parametrize("event_type", _STREAM_SCOPE_HOOKS) + def test_scalar_streaming_skips_non_streaming(self, event_type: GuardrailEventHooks): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=event_type, + stream_scope="streaming", + ) + assert guardrail.should_run_guardrail({"stream": True}, event_type) is True + assert guardrail.should_run_guardrail({"stream": False}, event_type) is False + assert guardrail.should_run_guardrail({}, event_type) is False + + @pytest.mark.parametrize("event_type", _STREAM_SCOPE_HOOKS) + def test_scalar_non_streaming_skips_streaming(self, event_type: GuardrailEventHooks): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=event_type, + stream_scope="non_streaming", + ) + assert guardrail.should_run_guardrail({"stream": False}, event_type) is True + assert guardrail.should_run_guardrail({}, event_type) is True + assert guardrail.should_run_guardrail({"stream": True}, event_type) is False + + def test_per_mode_map_applies_to_named_hooks_only(self): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call], + stream_scope={"pre_call": "both", "post_call": "streaming"}, + ) + assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.pre_call) is True + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.pre_call) is True + assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False + + def test_default_on_early_return_still_honors_stream_scope(self): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + stream_scope="non_streaming", + ) + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is False + + def test_apply_stream_scope_overwrites_constructor_default(self): + guardrail = CustomGuardrail( + guardrail_name="scoped", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is True + guardrail.apply_stream_scope("streaming") + assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False + + def test_direct_constructor_normalizes_mixed_case_map_keys(self): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope={"Pre_Call": "streaming"}, + ) + assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.pre_call) is True + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.pre_call) is False + + def test_realtime_transcription_counts_as_streaming(self): + streaming_only = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.realtime_input_transcription, + stream_scope="streaming", + ) + assert ( + streaming_only.should_run_guardrail( + {"litellm_metadata": {}}, GuardrailEventHooks.realtime_input_transcription + ) + is True + ) + non_streaming_only = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.realtime_input_transcription, + stream_scope="non_streaming", + ) + assert ( + non_streaming_only.should_run_guardrail( + {"litellm_metadata": {}}, GuardrailEventHooks.realtime_input_transcription + ) + is False + ) + + def test_path_defined_streaming_classification_cannot_be_spoofed(self): + generate_content_body: Final = {"contents": [{"parts": [{"text": "hi"}]}]} + streaming_only = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="streaming", + ) + assert streaming_only.should_run_guardrail(generate_content_body, GuardrailEventHooks.pre_call) is False + assert ( + streaming_only.should_run_guardrail( + {**generate_content_body, "is_streaming_request": True}, + GuardrailEventHooks.pre_call, + ) + is False + ) + assert ( + streaming_only.should_run_guardrail( + {**generate_content_body, "is_streaming_request": "litellm-server-streaming"}, + GuardrailEventHooks.pre_call, + ) + is False + ) + assert ( + streaming_only.should_run_guardrail( + {**generate_content_body, "litellm_server_streaming_classification": True}, + GuardrailEventHooks.pre_call, + ) + is False + ) + server_streaming_data: Final = guardrail_request_data_with_streaming( + generate_content_body, + is_streaming=True, + ) + assert ( + streaming_only.should_run_guardrail( + server_streaming_data, + GuardrailEventHooks.pre_call, + ) + is True + ) + non_streaming_only = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="non_streaming", + ) + assert ( + non_streaming_only.should_run_guardrail( + server_streaming_data, + GuardrailEventHooks.pre_call, + ) + is False + ) + + def test_streaming_classification_is_json_serializable_without_spoofing(self): + d: Final = guardrail_request_data_with_streaming({}, is_streaming=True) + serialized: Final = json.dumps(d) + assert _request_is_streaming(d) is True + + round_tripped: Final = json.loads(serialized) + assert _request_is_streaming(round_tripped) is False + assert round_tripped["litellm_server_streaming_classification"] == "litellm-server-streaming" + assert isinstance(round_tripped["litellm_server_streaming_classification"], str) + assert "litellm_server_streaming_classification" not in without_server_streaming_classification(round_tripped) + + def test_streaming_classification_preserves_caller_fields_and_removes_only_server_marker(self): + caller_data: Final = { + "contents": [{"parts": [{"text": "hi"}]}], + "is_streaming_request": "caller-value", + "litellm_server_streaming_classification": True, + } + non_streaming_data: Final = guardrail_request_data_with_streaming(caller_data, is_streaming=False) + server_streaming_data: Final = guardrail_request_data_with_streaming(caller_data, is_streaming=True) + + assert non_streaming_data is not caller_data + assert server_streaming_data is not caller_data + assert non_streaming_data == caller_data + assert server_streaming_data["is_streaming_request"] == "caller-value" + assert server_streaming_data["litellm_server_streaming_classification"] is not True + assert without_server_streaming_classification(caller_data) == caller_data + assert "litellm_server_streaming_classification" not in without_server_streaming_classification( + server_streaming_data + ) + + def test_server_streaming_classification_survives_scan_raw_request_snapshot(self): + from litellm.litellm_core_utils.core_helpers import independent_snapshot + + generate_content_body: Final = {"contents": [{"parts": [{"text": "hi"}]}]} + snapshot: Final = independent_snapshot( + guardrail_request_data_with_streaming(generate_content_body, is_streaming=True) + ) + streaming_only = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="streaming", + scan_raw_request=True, + ) + assert streaming_only.should_run_guardrail(snapshot, GuardrailEventHooks.pre_call) is True + assert ( + streaming_only.should_run_guardrail( + independent_snapshot({**generate_content_body, "is_streaming_request": True}), + GuardrailEventHooks.pre_call, + ) + is False + ) + + class TestApplyGuardrailCheck: def test_apply_guardrail_check_only_on_direct_implementation(self): """ @@ -3307,3 +3541,113 @@ class TestCustomGuardrailTimeout: ) assert guardrail.timeout == 7.0 + + +@pytest.mark.parametrize("stream_scope", [None, "both", {"post_call": "streaming"}]) +def test_guardrail_survives_deepcopy_and_pickle_with_its_stream_scope(stream_scope): + guardrail: Final = CustomGuardrail( + guardrail_name="copyable", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + stream_scope=stream_scope, + ) + expected: Final = guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) + for clone in (copy.deepcopy(guardrail), pickle.loads(pickle.dumps(guardrail))): + assert dict(clone.stream_scope_by_hook) == dict(guardrail.stream_scope_by_hook) + assert clone.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is expected + + +class _LockHoldingGuardrail(CustomGuardrail): + def __init__(self, **kwargs): + self.lock = threading.Lock() + super().__init__(**kwargs) + + def __getstate__(self): + state: Final = dict(self.__dict__) + state.pop("lock") + return state + + def __setstate__(self, state): + self.__dict__.update(state) + self.lock = threading.Lock() + + +class _SetstateOnlyGuardrail(CustomGuardrail): + def __setstate__(self, state): + self.__dict__.update(state) + self.restored = True + + +class _SlotsGuardrail(CustomGuardrail): + __slots__ = ("vendor_client",) + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.vendor_client = "vendor-client-object" + + +def _subclass_guardrails(): + kwargs: Final = { + "guardrail_name": "vendor", + "default_on": True, + "event_hook": GuardrailEventHooks.post_call, + "stream_scope": {"post_call": "streaming"}, + } + return ( + pytest.param(_LockHoldingGuardrail(**kwargs), id="dict-getstate"), + pytest.param(_SetstateOnlyGuardrail(**kwargs), id="setstate-only"), + pytest.param(_SlotsGuardrail(**kwargs), id="slots"), + ) + + +@pytest.mark.parametrize("guardrail", _subclass_guardrails()) +@pytest.mark.parametrize( + "cloner", + [copy.copy, copy.deepcopy, lambda g: pickle.loads(pickle.dumps(g))], + ids=["copy", "deepcopy", "pickle"], +) +def test_out_of_tree_guardrail_subclasses_survive_copy_and_pickle(guardrail, cloner): + clone = cloner(guardrail) + + if isinstance(clone, _LockHoldingGuardrail): + assert isinstance(clone.lock, type(threading.Lock())) + elif isinstance(clone, _SetstateOnlyGuardrail): + assert clone.restored is True + else: + assert clone.vendor_client == "vendor-client-object" + assert clone.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False + assert clone.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True + + +def test_router_constructs_with_a_dict_getstate_guardrail_in_deployment_callbacks(): + import litellm + + litellm.Router( + model_list=[ + { + "model_name": "m", + "litellm_params": { + "model": "openai/gpt-5.4-mini", + "api_key": "sk-test", + "callbacks": [ + _LockHoldingGuardrail( + guardrail_name="vendor", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + stream_scope={"post_call": "streaming"}, + ) + ], + }, + } + ] + ) + + +def test_subclass_that_skips_super_init_still_runs_with_default_scope(): + class NoSuperInit(CustomGuardrail): + def __init__(self) -> None: + self.guardrail_name = "no-super" + self.event_hook = None + self.default_on = True + + assert NoSuperInit().should_run_guardrail({"stream": True}, GuardrailEventHooks.pre_call) is True diff --git a/tests/unit/integrations/test_langfuse.py b/tests/unit/integrations/test_langfuse.py index bd55216e4bf..3c0238372a5 100644 --- a/tests/unit/integrations/test_langfuse.py +++ b/tests/unit/integrations/test_langfuse.py @@ -2197,6 +2197,58 @@ def test_log_event_returns_the_v2_dict_shape_for_the_alerting_trace_id_cache(): assert returned["generation_id"] +@pytest.mark.parametrize( + ("metadata", "expected_source"), + [ + ({}, None), + ({"trace_id": "my-unique-trace-id"}, "my-unique-trace-id"), + ({"existing_trace_id": "my-unique-existing-trace-id"}, "my-unique-existing-trace-id"), + ( + {"trace_id": "my-unique-trace-id", "existing_trace_id": "my-unique-existing-trace-id"}, + "my-unique-existing-trace-id", + ), + ], +) +def test_logging_get_trace_id_reports_the_langfuse_trace_that_won_precedence( + monkeypatch: pytest.MonkeyPatch, + metadata: dict[str, str], + expected_source: str | None, +) -> None: + from litellm.litellm_core_utils import litellm_logging + from litellm.litellm_core_utils.litellm_logging import Logging + + logger, exporter = _steering_logger() + monkeypatch.setattr(litellm_logging, "langFuseLogger", logger) + monkeypatch.setattr(litellm, "success_callback", ["langfuse"]) + monkeypatch.setattr(litellm, "_async_success_callback", []) + monkeypatch.setattr(litellm, "callbacks", []) + call_id: Final = f"trace-precedence-{len(metadata)}-{expected_source}" + logging_obj: Final = Logging( + model="gpt-5-mini", + messages=[{"role": "user", "content": "hi"}], + stream=False, + call_type="completion", + litellm_call_id=call_id, + start_time=datetime.datetime.now(), + function_id="trace-precedence", + ) + + litellm.completion( + model="gpt-5-mini", + messages=[{"role": "user", "content": "trace precedence"}], + mock_response="ok", + litellm_logging_obj=logging_obj, + metadata=dict(metadata), + ) + deadline: Final = time.monotonic() + 5 + while logging_obj.get_trace_id(service_name="langfuse") is None and time.monotonic() < deadline: + time.sleep(0.01) + + expected_trace_id: Final = resolve_trace_id(expected_source or logging_obj.litellm_trace_id) + assert logging_obj.get_trace_id(service_name="langfuse") == expected_trace_id + assert _span_trace_id(_exported_span(logger, exporter)) == expected_trace_id + + def test_parse_langfuse_debug_only_enables_on_true_strings(): """v4 treats any truthy value as debug=on, so the raw env string "false" would enable debug.""" assert langfuse_module.parse_langfuse_debug("true") is True diff --git a/tests/unit/integrations/test_langfuse_otel.py b/tests/unit/integrations/test_langfuse_otel.py index 0a9ce55fe16..13a49f40d7e 100644 --- a/tests/unit/integrations/test_langfuse_otel.py +++ b/tests/unit/integrations/test_langfuse_otel.py @@ -1,12 +1,17 @@ +import base64 import json import os +from datetime import datetime, timezone +from typing import Final from unittest.mock import MagicMock, patch import pytest +from opentelemetry.sdk.trace import TracerProvider from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger -from litellm.integrations.opentelemetry import OpenTelemetryConfig +from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.utils import StandardCallbackDynamicParams class TestLangfuseOtelIntegration: @@ -1006,5 +1011,58 @@ class TestLangfuseOtelResponsesAPI: assert output_data[0]["arguments"] == {} -if __name__ == "__main__": - pytest.main([__file__]) +LANGFUSE_ENV_ONLY: Final = { + key: value + for key, value in os.environ.items() + if not key.startswith(("LANGFUSE_", "OTEL_")) +} + + +def _decoded_basic_auth(header: str) -> str: + assert header.startswith("Basic ") + return base64.b64decode(header.removeprefix("Basic ")).decode() + + +def test_langfuse_otel_env_config_headers_carry_v4_ingestion_and_basic_auth() -> None: + with patch.dict( + os.environ, {**LANGFUSE_ENV_ONLY, "LANGFUSE_PUBLIC_KEY": "pk-lf-123", "LANGFUSE_SECRET_KEY": "sk-lf-123"}, clear=True + ): + logger: Final = LangfuseOtelLogger() + headers: Final = OpenTelemetry._get_headers_dictionary(logger.config.headers) + assert headers["x-langfuse-ingestion-version"] == "4" + assert _decoded_basic_auth(headers["Authorization"]) == "pk-lf-123:sk-lf-123" + + +def test_langfuse_otel_dynamic_headers_carry_v4_ingestion_and_basic_auth() -> None: + with patch.dict(os.environ, LANGFUSE_ENV_ONLY, clear=True): + logger: Final = LangfuseOtelLogger() + headers: Final = logger.construct_dynamic_otel_headers( + StandardCallbackDynamicParams(langfuse_public_key="pk-lf-dynamic", langfuse_secret_key="sk-lf-dynamic") + ) + assert headers is not None + assert headers["x-langfuse-ingestion-version"] == "4" + assert _decoded_basic_auth(headers["Authorization"]) == "pk-lf-dynamic:sk-lf-dynamic" + + +def test_langfuse_otel_does_not_start_proxy_request_span() -> None: + langfuse_provider: Final = TracerProvider() + generic_provider: Final = TracerProvider() + with patch.dict(os.environ, LANGFUSE_ENV_ONLY, clear=True): + langfuse_logger: Final = LangfuseOtelLogger(tracer_provider=langfuse_provider) + generic_logger: Final = OpenTelemetry( + config=OpenTelemetryConfig(exporter="console", skip_set_global=True), tracer_provider=generic_provider + ) + started_at: Final = datetime(2026, 1, 1, tzinfo=timezone.utc) + request_headers: Final = {"Authorization": "Bearer test"} + try: + assert ( + langfuse_logger.create_litellm_proxy_request_started_span(start_time=started_at, headers=request_headers) + is None + ) + assert ( + generic_logger.create_litellm_proxy_request_started_span(start_time=started_at, headers=request_headers) + is not None + ) + finally: + langfuse_provider.shutdown() + generic_provider.shutdown() diff --git a/tests/unit/integrations/test_opentelemetry_request_spans.py b/tests/unit/integrations/test_opentelemetry_request_spans.py new file mode 100644 index 00000000000..9ec40c00cb4 --- /dev/null +++ b/tests/unit/integrations/test_opentelemetry_request_spans.py @@ -0,0 +1,157 @@ +import asyncio +import json +from collections.abc import Sequence +from typing import Final + +import httpx +import pytest +import respx +from opentelemetry.sdk.trace import ReadableSpan, TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor, SpanExportResult +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +import litellm +from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig + +_OPENAI_URL: Final = "https://api.openai.com/v1/chat/completions" +_EXPECTED_SPAN_NAMES: Final = ("litellm_request", "raw_gen_ai_request") +_USER: Final = "OTEL_USER" +_USAGE: Final = {"prompt_tokens": 8, "completion_tokens": 2, "total_tokens": 10} +_COMPLETION: Final = { + "id": "chatcmpl-otel", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4.1-mini-2025-04-14", + "service_tier": "default", + "system_fingerprint": "fp_otel", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hello"}, "finish_reason": "stop"}], + "usage": _USAGE, +} +_STREAM: Final = ( + "".join( + f"data: {json.dumps(chunk)}\n\n" + for chunk in ( + { + "id": "chatcmpl-otel", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": "gpt-4.1-mini-2025-04-14", + "choices": [{"index": 0, "delta": {"role": "assistant", "content": "hello"}, "finish_reason": None}], + }, + { + "id": "chatcmpl-otel", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": "gpt-4.1-mini-2025-04-14", + "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], + "usage": _USAGE, + }, + ) + ) + + "data: [DONE]\n\n" +) +_LITELLM_REQUEST_ATTRIBUTES: Final = ( + "gen_ai.request.model", + "gen_ai.system", + "gen_ai.request.temperature", + "llm.is_streaming", + "llm.user", + "gen_ai.response.id", + "gen_ai.response.model", + "gen_ai.usage.total_tokens", + "gen_ai.usage.output_tokens", + "gen_ai.usage.input_tokens", +) +_RAW_STREAMING_ATTRIBUTES: Final = ( + "llm.openai.messages", + "llm.openai.temperature", + "llm.openai.user", + "llm.openai.extra_body", + "llm.openai.model", +) +_RAW_NON_STREAMING_ATTRIBUTES: Final = ( + *_RAW_STREAMING_ATTRIBUTES, + "llm.openai.id", + "llm.openai.choices", + "llm.openai.created", + "llm.openai.object", + "llm.openai.service_tier", + "llm.openai.system_fingerprint", + "llm.openai.usage", +) + + +def _is_our_request(span: ReadableSpan) -> bool: + return span.name == "litellm_request" and (span.attributes or {}).get("llm.user") == _USER + + +def _trace_id(span: ReadableSpan) -> int: + assert span.context is not None + return span.context.trace_id + + +class _SignallingExporter(InMemorySpanExporter): + def __init__(self, loop: asyncio.AbstractEventLoop) -> None: + super().__init__() + self.loop: Final = loop + self.request_span_exported: Final = asyncio.Event() + + def export(self, spans: Sequence[ReadableSpan]) -> SpanExportResult: + result: Final = super().export(spans) + if any(_is_our_request(span) for span in spans): + self.loop.call_soon_threadsafe(self.request_span_exported.set) + return result + + +@pytest.mark.asyncio +@pytest.mark.parametrize("streaming", [True, False]) +async def test_otel_callback_emits_the_request_and_raw_provider_spans( + streaming: bool, monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.delenv("OTEL_SEMCONV_STABILITY_OPT_IN", raising=False) + exporter: Final = _SignallingExporter(asyncio.get_running_loop()) + tracer_provider: Final = TracerProvider() + tracer_provider.add_span_processor(SimpleSpanProcessor(exporter)) + monkeypatch.setattr( + litellm, + "callbacks", + [OpenTelemetry(config=OpenTelemetryConfig(exporter=exporter), tracer_provider=tracer_provider)], + ) + respx_mock.post(_OPENAI_URL).mock( + return_value=httpx.Response(200, text=_STREAM, headers={"content-type": "text/event-stream"}) + if streaming + else httpx.Response(200, json=_COMPLETION) + ) + + response: Final = await litellm.acompletion( + model="gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + temperature=0.1, + user=_USER, + stream=streaming, + api_key="sk-unit-test", + ) + if streaming: + assert [chunk async for chunk in response] + await asyncio.wait_for(exporter.request_span_exported.wait(), timeout=10) + + finished: Final = exporter.get_finished_spans() + request_span: Final = next(span for span in finished if _is_our_request(span)) + ours: Final = tuple(span for span in finished if _trace_id(span) == _trace_id(request_span)) + assert tuple(sorted(span.name for span in ours)) == _EXPECTED_SPAN_NAMES + spans: Final = {span.name: span for span in ours} + request_attributes: Final = spans["litellm_request"].attributes or {} + assert all(request_attributes.get(name) is not None for name in _LITELLM_REQUEST_ATTRIBUTES) + assert request_attributes["gen_ai.request.model"] == "gpt-4.1-mini" + assert request_attributes["gen_ai.system"] == "openai" + assert request_attributes["gen_ai.request.temperature"] == 0.1 + assert request_attributes["llm.is_streaming"] == str(streaming) + assert request_attributes["llm.user"] == _USER + assert request_attributes["gen_ai.response.id"] == "chatcmpl-otel" + assert request_attributes["gen_ai.usage.input_tokens"] == _USAGE["prompt_tokens"] + assert request_attributes["gen_ai.usage.output_tokens"] == _USAGE["completion_tokens"] + assert request_attributes["gen_ai.usage.total_tokens"] == _USAGE["total_tokens"] + raw_attributes: Final = spans["raw_gen_ai_request"].attributes or {} + expected_raw: Final = _RAW_STREAMING_ATTRIBUTES if streaming else _RAW_NON_STREAMING_ATTRIBUTES + assert all(raw_attributes.get(name) is not None for name in expected_raw) diff --git a/tests/unit/integrations/test_prometheus_client_ip_user_agent.py b/tests/unit/integrations/test_prometheus_client_ip_user_agent.py index 004ac4dbffb..8f7923428e7 100644 --- a/tests/unit/integrations/test_prometheus_client_ip_user_agent.py +++ b/tests/unit/integrations/test_prometheus_client_ip_user_agent.py @@ -94,7 +94,7 @@ async def test_async_post_call_success_hook_includes_client_ip_user_agent(): logger._increment_token_metrics = MagicMock() logger._increment_remaining_budget_metrics = AsyncMock() logger._set_virtual_key_rate_limit_metrics = MagicMock() - logger._set_key_and_team_rate_limit_metrics = MagicMock() + logger._set_v3_rate_limit_allowed_and_used_metrics = MagicMock() logger._set_latency_metrics = MagicMock() logger.set_llm_deployment_success_metrics = MagicMock() logger._increment_cache_metrics = MagicMock() diff --git a/tests/unit/integrations/test_prometheus_rate_limit_labels.py b/tests/unit/integrations/test_prometheus_rate_limit_labels.py index bf1d68c7714..e4c71ff7f77 100644 --- a/tests/unit/integrations/test_prometheus_rate_limit_labels.py +++ b/tests/unit/integrations/test_prometheus_rate_limit_labels.py @@ -14,8 +14,10 @@ Covers two follow-up gaps to the unified rate-limit error work: """ from collections.abc import Mapping +from typing import Final from unittest.mock import MagicMock, patch +import litellm import pytest from litellm.exceptions import ( @@ -474,11 +476,13 @@ def test_should_ignore_non_int_v3_header_values(bad_value): ) -KEY_AND_TEAM_RATE_LIMIT_METRICS = ( +RATE_LIMIT_METRICS = ( "litellm_api_key_rate_limit_allowed_metric", "litellm_api_key_rate_limit_used_metric", "litellm_team_rate_limit_allowed_metric", "litellm_team_rate_limit_used_metric", + "litellm_project_model_rate_limit_allowed_metric", + "litellm_project_model_rate_limit_used_metric", ) @@ -534,6 +538,8 @@ def _success_kwargs_with_rate_limit_headers(additional_headers: Mapping[str, obj "user_api_key_alias": "key-alias", "user_api_key_team_id": "team-id", "user_api_key_team_alias": "team-alias", + "user_api_key_project_id": "project-id", + "user_api_key_project_alias": "project-alias", "user_api_key_user_id": "u", "user_api_key_user_email": "e@x.com", "user_api_key_org_id": None, @@ -627,6 +633,68 @@ async def test_should_emit_key_and_team_rate_limit_allowed_and_used_from_v3_head _clear_prometheus_registry() +@pytest.mark.asyncio +async def test_should_emit_project_model_rate_limit_allowed_and_used_from_v3_headers() -> None: + _clear_prometheus_registry() + try: + await _run_success_event( + { + "x-ratelimit-model_per_project-limit-requests": 100, + "x-ratelimit-model_per_project-remaining-requests": 99, + "x-ratelimit-model_per_project-limit-tokens": 10000, + "x-ratelimit-model_per_project-remaining-tokens": 9950, + "x-ratelimit-model_per_project_itpm-limit-tokens": 2000, + "x-ratelimit-model_per_project_itpm-remaining-tokens": 1900, + "x-ratelimit-model_per_project_otpm-limit-tokens": 3000, + "x-ratelimit-model_per_project_otpm-remaining-tokens": 2750, + } + ) + + project_requests: Final = ( + ("project_alias", "project-alias"), + ("project_id", "project-id"), + ("rate_limit_type", "requests"), + ("requested_model", "anthropic-haiku-4-5"), + ) + project_tokens: Final = ( + ("project_alias", "project-alias"), + ("project_id", "project-id"), + ("rate_limit_type", "tokens"), + ("requested_model", "anthropic-haiku-4-5"), + ) + project_input_tokens: Final = ( + ("project_alias", "project-alias"), + ("project_id", "project-id"), + ("rate_limit_type", "input_tokens"), + ("requested_model", "anthropic-haiku-4-5"), + ) + project_output_tokens: Final = ( + ("project_alias", "project-alias"), + ("project_id", "project-id"), + ("rate_limit_type", "output_tokens"), + ("requested_model", "anthropic-haiku-4-5"), + ) + + assert _collected_samples("litellm_project_model_rate_limit_allowed_metric") == { + project_requests: 100, + project_tokens: 10000, + project_input_tokens: 2000, + project_output_tokens: 3000, + } + assert _collected_samples("litellm_project_model_rate_limit_used_metric") == { + project_requests: 1, + project_tokens: 50, + project_input_tokens: 100, + project_output_tokens: 250, + } + assert _collected_samples("litellm_api_key_rate_limit_allowed_metric") == {} + assert _collected_samples("litellm_api_key_rate_limit_used_metric") == {} + assert _collected_samples("litellm_team_rate_limit_allowed_metric") == {} + assert _collected_samples("litellm_team_rate_limit_used_metric") == {} + finally: + _clear_prometheus_registry() + + @pytest.mark.asyncio async def test_should_emit_only_the_dimensions_the_limiter_enforced(): """ @@ -701,6 +769,79 @@ async def test_should_drop_key_and_team_series_once_the_limiter_stops_reporting_ _clear_prometheus_registry() +@pytest.mark.asyncio +async def test_should_drop_project_model_series_once_the_limiter_stops_reporting_a_limit() -> None: + _clear_prometheus_registry() + try: + logger: Final = PrometheusLogger() + await _run_success_event( + { + "x-ratelimit-model_per_project-limit-requests": 100, + "x-ratelimit-model_per_project-remaining-requests": 99, + "x-ratelimit-model_per_project-limit-tokens": 10000, + "x-ratelimit-model_per_project-remaining-tokens": 9950, + "x-ratelimit-model_per_project_itpm-limit-tokens": 2000, + "x-ratelimit-model_per_project_itpm-remaining-tokens": 1900, + "x-ratelimit-model_per_project_otpm-limit-tokens": 3000, + "x-ratelimit-model_per_project_otpm-remaining-tokens": 2750, + }, + logger=logger, + ) + await _run_success_event( + { + "x-ratelimit-model_per_project-limit-requests": 100, + "x-ratelimit-model_per_project-remaining-requests": 96, + }, + logger=logger, + ) + + project_requests: Final = ( + ("project_alias", "project-alias"), + ("project_id", "project-id"), + ("rate_limit_type", "requests"), + ("requested_model", "anthropic-haiku-4-5"), + ) + assert _collected_samples("litellm_project_model_rate_limit_allowed_metric") == { + project_requests: 100, + } + assert _collected_samples("litellm_project_model_rate_limit_used_metric") == { + project_requests: 4, + } + finally: + _clear_prometheus_registry() + + +@pytest.mark.asyncio +async def test_project_model_rate_limit_allowed_uses_same_custom_project_alias_label_as_requests() -> None: + original_custom_labels: Final = litellm.custom_prometheus_metadata_labels + litellm.custom_prometheus_metadata_labels = ["metadata.user_api_key_project_alias"] + _clear_prometheus_registry() + try: + logger: Final = PrometheusLogger() + await _run_success_event( + { + "x-ratelimit-model_per_project-limit-requests": 100, + "x-ratelimit-model_per_project-remaining-requests": 99, + }, + logger=logger, + ) + + allowed_samples: Final = _collected_samples("litellm_project_model_rate_limit_allowed_metric") + request_samples: Final = _collected_samples("litellm_proxy_total_requests_metric_total") + assert len(allowed_samples) == 1 + assert len(request_samples) == 1 + allowed_labels: Final = dict(next(iter(allowed_samples))) + request_labels: Final = dict(next(iter(request_samples))) + assert allowed_labels["metadata_user_api_key_project_alias"] == "project-alias" + assert ( + allowed_labels["metadata_user_api_key_project_alias"] + == request_labels["metadata_user_api_key_project_alias"] + ) + finally: + litellm.custom_prometheus_metadata_labels = original_custom_labels + _clear_prometheus_registry() + + @pytest.mark.asyncio @pytest.mark.parametrize( "additional_headers", @@ -710,16 +851,25 @@ async def test_should_drop_key_and_team_series_once_the_limiter_stops_reporting_ {"x-ratelimit-api_key-limit-requests": 10}, {"x-ratelimit-api_key-limit-requests": "10", "x-ratelimit-api_key-remaining-requests": "7"}, {"x-ratelimit-team-limit-tokens": True, "x-ratelimit-team-remaining-tokens": 5}, + {"x-ratelimit-model_per_project-limit-requests": 10}, + { + "x-ratelimit-model_per_project-limit-requests": "10", + "x-ratelimit-model_per_project-remaining-requests": "7", + }, + { + "x-ratelimit-model_per_project-limit-tokens": True, + "x-ratelimit-model_per_project-remaining-tokens": 5, + }, ], ) -async def test_should_emit_no_key_or_team_rate_limit_series_without_a_complete_int_pair( - additional_headers, -): +async def test_should_emit_no_rate_limit_series_without_a_complete_int_pair( + additional_headers: Mapping[str, object] | None, +) -> None: _clear_prometheus_registry() try: await _run_success_event(additional_headers) - for metric_name in KEY_AND_TEAM_RATE_LIMIT_METRICS: + for metric_name in RATE_LIMIT_METRICS: assert _collected_samples(metric_name) == {}, metric_name finally: _clear_prometheus_registry() diff --git a/tests/unit/integrations/vector_store_integrations/test_bedrock_kb_context_offline.py b/tests/unit/integrations/vector_store_integrations/test_bedrock_kb_context_offline.py new file mode 100644 index 00000000000..20d3a310bdb --- /dev/null +++ b/tests/unit/integrations/vector_store_integrations/test_bedrock_kb_context_offline.py @@ -0,0 +1,289 @@ +import itertools +import json +from typing import Final, TypedDict + +from typing_extensions import ReadOnly + +import httpx +import pytest +import respx +from pydantic import TypeAdapter + +import litellm +import litellm.proxy.proxy_server as proxy_server +from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import VectorStorePreCallHook +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.vector_stores.vector_store_registry import LiteLLM_ManagedVectorStore, VectorStoreRegistry + +_KB_ID: Final = "T37J8R4WTM" +_KB_URL: Final = f"https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases/{_KB_ID}/retrieve" +_ANTHROPIC_URL: Final = "https://api.anthropic.com/v1/messages" +_OPENAI_URL: Final = "https://api.openai.com/v1/chat/completions" +_KB_TEXT: Final = "LiteLLM is a library that simplifies LLM API access" +_PREFIX: Final = VectorStorePreCallHook.CONTENT_PREFIX_STRING +_BODY: Final = TypeAdapter(dict[str, object]) + +_ANTHROPIC_MESSAGE: Final = { + "id": "msg_kb", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "LiteLLM simplifies LLM access."}], + "model": "claude-haiku-4-5-20251001", + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 100, "output_tokens": 50}, +} +_OPENAI_COMPLETION: Final = { + "id": "chatcmpl-kb", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-5-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12}, +} +_ANTHROPIC_STREAM: Final = ( + 'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_kb_stream","type":"message",' + '"role":"assistant","content":[],"model":"claude-haiku-4-5-20251001","stop_reason":null,"stop_sequence":null,' + '"usage":{"input_tokens":10,"output_tokens":1}}}\n\n' + 'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}\n\n' + 'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"LiteLLM"}}\n\n' + 'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n' + 'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},' + '"usage":{"output_tokens":2}}\n\n' + 'event: message_stop\ndata: {"type":"message_stop"}\n\n' +) + + +class _ChatMessage(TypedDict): + role: ReadOnly[str] + content: ReadOnly[str] + + +@pytest.fixture(autouse=True) +def _knowledge_base(monkeypatch: pytest.MonkeyPatch, fake_provider_credentials: None) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setenv("AWS_REGION", "us-west-2") + monkeypatch.setenv("AWS_REGION_NAME", "us-west-2") + monkeypatch.setattr( + litellm, + "vector_store_registry", + VectorStoreRegistry( + vector_stores=[LiteLLM_ManagedVectorStore(vector_store_id=_KB_ID, custom_llm_provider="bedrock")] + ), + raising=False, + ) + + +def _kb_route(respx_mock: respx.MockRouter) -> respx.Route: + return respx_mock.post(_KB_URL).mock( + return_value=httpx.Response( + 200, + json={"retrievalResults": [{"content": {"text": _KB_TEXT, "type": "TEXT"}, "score": 0.9, "metadata": {}}]}, + ) + ) + + +def _sent_body(route: respx.Route) -> dict[str, object]: + return _BODY.validate_json(route.calls.last.request.content) + + +@pytest.mark.asyncio +async def test_completion_with_vector_store_ids_prepends_the_kb_context_block(respx_mock: respx.MockRouter) -> None: + _kb_route(respx_mock) + anthropic: Final = respx_mock.post(_ANTHROPIC_URL).mock(return_value=httpx.Response(200, json=_ANTHROPIC_MESSAGE)) + + await litellm.acompletion( + model="anthropic/claude-haiku-4-5-20251001", + messages=[{"role": "user", "content": "what is litellm?"}], + vector_store_ids=[_KB_ID], + ) + + messages: Final = TypeAdapter(tuple[dict[str, object], ...]).validate_python(_sent_body(anthropic)["messages"]) + content: Final = TypeAdapter(tuple[dict[str, str], ...]).validate_python(messages[0]["content"]) + assert anthropic.call_count == 1 + assert [block["type"] for block in content] == ["text", "text"] + assert content[0]["text"] == f"{_PREFIX}{_KB_TEXT}\n\n" + assert content[1]["text"] == "what is litellm?" + + +@pytest.mark.asyncio +async def test_streaming_completion_carries_the_search_results_on_a_chunk_delta(respx_mock: respx.MockRouter) -> None: + _kb_route(respx_mock) + respx_mock.post(_ANTHROPIC_URL).mock( + return_value=httpx.Response(200, text=_ANTHROPIC_STREAM, headers={"content-type": "text/event-stream"}) + ) + + response: Final = await litellm.acompletion( + model="anthropic/claude-haiku-4-5-20251001", + messages=[{"role": "user", "content": "what is litellm?"}], + vector_store_ids=[_KB_ID], + stream=True, + ) + chunks: Final = tuple([chunk async for chunk in response]) + choices: Final = tuple(itertools.chain.from_iterable(chunk.choices for chunk in chunks)) + annotated: Final = tuple( + choice.delta.provider_specific_fields["search_results"] + for choice in choices + if choice.delta.provider_specific_fields and "search_results" in choice.delta.provider_specific_fields + ) + + assert len(chunks) > 0 + assert len(annotated) >= 1 + assert annotated[0][0]["object"] == "vector_store.search_results.page" + assert annotated[0][0]["data"][0]["content"][0]["text"] == _KB_TEXT + + +@pytest.mark.asyncio +async def test_file_search_filters_reach_the_kb_as_a_bedrock_equals_filter(respx_mock: respx.MockRouter) -> None: + kb: Final = _kb_route(respx_mock) + respx_mock.post(_ANTHROPIC_URL).mock(return_value=httpx.Response(200, json=_ANTHROPIC_MESSAGE)) + + response: Final = await litellm.acompletion( + model="anthropic/claude-haiku-4-5-20251001", + messages=[{"role": "user", "content": "what is litellm?"}], + max_tokens=10, + tools=[ + { + "type": "file_search", + "vector_store_ids": [_KB_ID], + "filters": {"key": "user_id", "value": "fake-user-id", "operator": "eq"}, + } + ], + ) + + retrieval: Final = TypeAdapter(dict[str, dict[str, dict[str, object]]]).validate_python( + _sent_body(kb)["retrievalConfiguration"] + ) + assert retrieval["vectorSearchConfiguration"]["filter"] == {"equals": {"key": "user_id", "value": "fake-user-id"}} + assert response.choices[0].message.content == "LiteLLM simplifies LLM access." + + +def _openai_messages(route: respx.Route) -> tuple[_ChatMessage, ...]: + return TypeAdapter(tuple[_ChatMessage, ...]).validate_python(_sent_body(route)["messages"]) + + +@pytest.mark.asyncio +async def test_openai_request_with_vector_store_ids_leads_with_a_kb_context_user_message( + respx_mock: respx.MockRouter, +) -> None: + _kb_route(respx_mock) + openai: Final = respx_mock.post(_OPENAI_URL).mock(return_value=httpx.Response(200, json=_OPENAI_COMPLETION)) + + await litellm.acompletion( + model="gpt-5-mini", + messages=[{"role": "user", "content": "what is litellm?"}], + vector_store_ids=[_KB_ID], + ) + + assert _openai_messages(openai) == ( + _ChatMessage(role="user", content=f"{_PREFIX}{_KB_TEXT}\n\n"), + _ChatMessage(role="user", content="what is litellm?"), + ) + + +@pytest.mark.asyncio +async def test_a_managed_file_search_tool_is_resolved_locally_and_not_sent_upstream( + respx_mock: respx.MockRouter, +) -> None: + _kb_route(respx_mock) + openai: Final = respx_mock.post(_OPENAI_URL).mock(return_value=httpx.Response(200, json=_OPENAI_COMPLETION)) + + await litellm.acompletion( + model="gpt-5-mini", + messages=[{"role": "user", "content": "what is litellm?"}], + tools=[{"type": "file_search", "vector_store_ids": [_KB_ID]}], + ) + + assert _openai_messages(openai)[0] == _ChatMessage(role="user", content=f"{_PREFIX}{_KB_TEXT}\n\n") + assert "tools" not in _sent_body(openai) + + +@pytest.mark.asyncio +async def test_an_unknown_vector_store_tool_is_forwarded_while_the_known_one_is_resolved( + respx_mock: respx.MockRouter, +) -> None: + _kb_route(respx_mock) + openai: Final = respx_mock.post(_OPENAI_URL).mock(return_value=httpx.Response(200, json=_OPENAI_COMPLETION)) + + await litellm.acompletion( + model="gpt-5-mini", + messages=[{"role": "user", "content": "what is litellm?"}], + tools=[ + {"type": "file_search", "vector_store_ids": [_KB_ID]}, + {"type": "file_search", "vector_store_ids": ["unknownVS"]}, + ], + ) + + assert _openai_messages(openai)[0] == _ChatMessage(role="user", content=f"{_PREFIX}{_KB_TEXT}\n\n") + assert _sent_body(openai)["tools"] == [{"type": "file_search", "vector_store_ids": ["unknownVS"]}] + + +def _authorized_key() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-kb-proxy", user_id="kb-proxy-user") + + +class _SearchResultsPage(TypedDict): + object: ReadOnly[str] + search_query: ReadOnly[str] + data: ReadOnly[list[dict[str, object]]] + + +class _ProviderFields(TypedDict): + search_results: ReadOnly[list[_SearchResultsPage]] + + +class _ProxyMessage(TypedDict): + role: ReadOnly[str] + content: ReadOnly[str] + provider_specific_fields: ReadOnly[_ProviderFields] + + +class _ProxyChoice(TypedDict): + message: ReadOnly[_ProxyMessage] + + +class _ProxyCompletion(TypedDict): + choices: ReadOnly[list[_ProxyChoice]] + + +@pytest.mark.asyncio +async def test_proxy_http_response_keeps_the_kb_search_results_in_provider_specific_fields( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + kb: Final = _kb_route(respx_mock) + upstream: Final = respx_mock.post(_OPENAI_URL).mock(return_value=httpx.Response(200, json=_OPENAI_COMPLETION)) + monkeypatch.setattr( + proxy_server, + "llm_router", + litellm.Router( + model_list=[ + {"model_name": "gpt-5-mini", "litellm_params": {"model": "openai/gpt-5-mini", "api_key": "sk-fixture"}} + ], + num_retries=0, + ), + ) + monkeypatch.setitem(proxy_server.app.dependency_overrides, user_api_key_auth, _authorized_key) + + async with httpx.AsyncClient( + transport=httpx.ASGITransport(proxy_server.app), base_url="http://kb-proxy.test" + ) as client: + result: Final = await client.post( + "/v1/chat/completions", + json={ + "model": "gpt-5-mini", + "messages": [{"role": "user", "content": "what is litellm?"}], + "vector_store_ids": [_KB_ID], + }, + ) + + assert result.status_code == 200, result.text + assert kb.call_count == 1 + assert upstream.call_count == 1 + message: Final = TypeAdapter(_ProxyCompletion).validate_json(result.content)["choices"][0]["message"] + assert message["content"] == "ok" + pages: Final = message["provider_specific_fields"]["search_results"] + assert [page["object"] for page in pages] == ["vector_store.search_results.page"] + assert pages[0]["search_query"] == "what is litellm?" + assert pages[0]["data"] + assert _KB_TEXT in json.dumps(pages[0]["data"]) diff --git a/tests/unit/interactions/test_main.py b/tests/unit/interactions/test_main.py index 80cce7f1b6f..626d53d3384 100644 --- a/tests/unit/interactions/test_main.py +++ b/tests/unit/interactions/test_main.py @@ -1,4 +1,10 @@ +import base64 +import json +from typing import Final + +import httpx import pytest +from respx import MockRouter import litellm import litellm.interactions as interactions @@ -18,3 +24,71 @@ class TestGoogleInteractionsCreate: input="Hello", api_key=api_key, ) + + +class TestInteractionsAcreateOffline: + @pytest.fixture(autouse=True) + def _httpx_only_transport(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + + @pytest.mark.usefixtures("fake_provider_credentials") + @pytest.mark.asyncio + async def test_acreate_simple_gemini(self, respx_mock: MockRouter) -> None: + route: Final = respx_mock.post("https://generativelanguage.googleapis.com/v1beta/interactions").mock( + return_value=httpx.Response( + 200, + json={ + "id": "interaction-offline", + "object": "interaction", + "model": "gemini-2.5-flash", + "status": "completed", + "steps": [{"type": "model_output", "content": [{"type": "text", "text": "299792458"}]}], + "usage": {"input_tokens": 6, "output_tokens": 3}, + }, + ) + ) + response: Final = await interactions.acreate( + model="gemini/gemini-2.5-flash", + input="What is the speed of light?", + api_key="gemini-offline", + ) + body: Final = json.loads(route.calls.last.request.content) + assert body["model"] == "gemini-2.5-flash" + assert body["input"] == "What is the speed of light?" + assert response.id == "interaction-offline" + assert response.status == "completed" + + @pytest.mark.usefixtures("fake_provider_credentials") + @pytest.mark.asyncio + async def test_acreate_simple_litellm_responses_bridge(self, respx_mock: MockRouter) -> None: + route: Final = respx_mock.post("https://api.openai.com/v1/responses").mock( + return_value=httpx.Response( + 200, + json={ + "id": "resp-offline", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o", + "output": [ + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "299792458 m/s"}], + } + ], + "usage": {"input_tokens": 6, "output_tokens": 3, "total_tokens": 9}, + }, + ) + ) + response: Final = await interactions.acreate( + model="gpt-4o", + input="What is the speed of light?", + api_key="sk-offline", + ) + body: Final = json.loads(route.calls.last.request.content) + assert body["model"] == "gpt-4o" + serialized: Final = json.dumps(body) + assert "What is the speed of light?" in serialized + assert "response_id:resp-offline" in base64.b64decode(response.id.removeprefix("resp_")).decode() + assert response.status == "completed" diff --git a/tests/unit/litellm_core_utils/llm_cost_calc/test_web_search_logged_cost.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_web_search_logged_cost.py new file mode 100644 index 00000000000..1448318b07d --- /dev/null +++ b/tests/unit/litellm_core_utils/llm_cost_calc/test_web_search_logged_cost.py @@ -0,0 +1,210 @@ +import asyncio +import json +from typing import Final, Literal, TypedDict + +from typing_extensions import ReadOnly + +import httpx +import pytest +import respx +from pydantic import TypeAdapter + +import litellm +from litellm.integrations.custom_logger import CustomLogger + +_MODEL: Final = "gpt-sized-search-unit" +_INPUT_COST: Final = 1e-06 +_OUTPUT_COST: Final = 4e-06 +_PER_QUERY: Final = { + "search_context_size_low": 0.011, + "search_context_size_medium": 0.022, + "search_context_size_high": 0.033, +} +_PROMPT_TOKENS: Final = 100 +_COMPLETION_TOKENS: Final = 20 +_USAGE: Final = { + "prompt_tokens": _PROMPT_TOKENS, + "completion_tokens": _COMPLETION_TOKENS, + "total_tokens": _PROMPT_TOKENS + _COMPLETION_TOKENS, +} +_CHAT_RESPONSE: Final = { + "id": "chatcmpl-search", + "object": "chat.completion", + "created": 1700000000, + "model": _MODEL, + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "A positive story", + "annotations": [ + { + "type": "url_citation", + "url_citation": { + "start_index": 0, + "end_index": 5, + "title": "news", + "url": "https://news.example/a", + }, + } + ], + }, + "finish_reason": "stop", + } + ], + "usage": _USAGE, +} +_RESPONSES_BODY: Final = { + "id": "resp_search", + "object": "response", + "created_at": 1700000000, + "status": "completed", + "model": _MODEL, + "output": [ + {"type": "web_search_call", "id": "ws_search", "status": "completed"}, + { + "type": "message", + "id": "msg_search", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "A positive story", "annotations": []}], + }, + ], + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + "usage": { + "input_tokens": _PROMPT_TOKENS, + "output_tokens": _COMPLETION_TOKENS, + "total_tokens": _PROMPT_TOKENS + _COMPLETION_TOKENS, + }, +} +_RESPONSES_STREAM: Final = "".join( + f"event: {event['type']}\ndata: {json.dumps(event)}\n\n" + for event in ( + {"type": "response.created", "response": {**_RESPONSES_BODY, "status": "in_progress", "output": []}}, + {"type": "response.completed", "response": _RESPONSES_BODY}, + ) +) + +_ContextSize = Literal["search_context_size_low", "search_context_size_medium", "search_context_size_high"] + + +class _LoggedCost(TypedDict): + response_cost: ReadOnly[float] + prompt_tokens: ReadOnly[int] + completion_tokens: ReadOnly[int] + + +class _CostRecorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.payloads: Final[list[_LoggedCost]] = [] + self.logged: Final = asyncio.Event() + + async def async_log_success_event( + self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + self.payloads.append(TypeAdapter(_LoggedCost).validate_python(kwargs["standard_logging_object"])) + self.logged.set() + + +@pytest.fixture +def recorder(monkeypatch: pytest.MonkeyPatch) -> _CostRecorder: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + entry: Final = { + "input_cost_per_token": _INPUT_COST, + "output_cost_per_token": _OUTPUT_COST, + "litellm_provider": "openai", + "mode": "chat", + "max_tokens": 4096, + "max_input_tokens": 4096, + "max_output_tokens": 4096, + "supports_web_search": True, + "search_context_cost_per_query": _PER_QUERY, + } + monkeypatch.setitem(litellm.model_cost, _MODEL, entry) + monkeypatch.setitem(litellm.model_cost, f"openai/{_MODEL}", entry) + cost_recorder: Final = _CostRecorder() + monkeypatch.setattr(litellm, "callbacks", [cost_recorder]) + return cost_recorder + + +async def _logged_cost(recorder: _CostRecorder) -> _LoggedCost: + await asyncio.wait_for(recorder.logged.wait(), timeout=10) + return recorder.payloads[-1] + + +def _expected_cost(payload: _LoggedCost, size: _ContextSize) -> float: + return payload["prompt_tokens"] * _INPUT_COST + payload["completion_tokens"] * _OUTPUT_COST + _PER_QUERY[size] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("web_search_options", "size"), + [ + (None, "search_context_size_medium"), + ({"search_context_size": "low"}, "search_context_size_low"), + ({"search_context_size": "high"}, "search_context_size_high"), + ], +) +async def test_chat_web_search_logged_cost_adds_the_per_query_cost_for_the_context_size( + web_search_options: dict[str, str] | None, + size: _ContextSize, + recorder: _CostRecorder, + respx_mock: respx.MockRouter, +) -> None: + respx_mock.post("https://api.openai.com/v1/chat/completions").mock( + return_value=httpx.Response(200, json=_CHAT_RESPONSE) + ) + options: Final = {"web_search_options": web_search_options} if web_search_options is not None else {} + + await litellm.acompletion( + model=f"openai/{_MODEL}", + messages=[{"role": "user", "content": "What was a positive news story from today?"}], + api_key="sk-unit-test", + **options, + ) + payload: Final = await _logged_cost(recorder) + + assert (payload["prompt_tokens"], payload["completion_tokens"]) == (_PROMPT_TOKENS, _COMPLETION_TOKENS) + assert payload["response_cost"] == pytest.approx(_expected_cost(payload, size), abs=1e-12) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("tools", "size", "stream"), + [ + ([{"type": "web_search_preview", "search_context_size": "low"}], "search_context_size_low", True), + ([{"type": "web_search_preview", "search_context_size": "low"}], "search_context_size_low", False), + ([{"type": "web_search_preview"}], "search_context_size_medium", True), + ([{"type": "web_search_preview"}], "search_context_size_medium", False), + ], +) +async def test_responses_web_search_logged_cost_adds_the_per_query_cost_for_the_context_size( + tools: list[dict[str, str]], + size: _ContextSize, + stream: bool, + recorder: _CostRecorder, + respx_mock: respx.MockRouter, +) -> None: + respx_mock.post("https://api.openai.com/v1/responses").mock( + return_value=httpx.Response(200, text=_RESPONSES_STREAM, headers={"content-type": "text/event-stream"}) + if stream + else httpx.Response(200, json=_RESPONSES_BODY) + ) + + response: Final = await litellm.aresponses( + model=f"openai/{_MODEL}", + input=[{"role": "user", "content": "What was a positive news story from today?"}], + tools=tools, + stream=stream, + api_key="sk-unit-test", + ) + if stream: + assert [event async for event in response] + payload: Final = await _logged_cost(recorder) + + assert (payload["prompt_tokens"], payload["completion_tokens"]) == (_PROMPT_TOKENS, _COMPLETION_TOKENS) + assert payload["response_cost"] == pytest.approx(_expected_cost(payload, size), abs=1e-12) diff --git a/tests/unit/litellm_core_utils/test_default_encoding.py b/tests/unit/litellm_core_utils/test_default_encoding.py new file mode 100644 index 00000000000..fd00a4b013a --- /dev/null +++ b/tests/unit/litellm_core_utils/test_default_encoding.py @@ -0,0 +1,40 @@ +import importlib +import os +from collections.abc import Iterator +from pathlib import Path +from typing import Final + +import pytest + +import litellm.litellm_core_utils.default_encoding as default_encoding + +BUNDLED_TOKENIZERS: Final = Path(default_encoding.filename) + + +@pytest.fixture +def reload_default_encoding(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.delenv("TIKTOKEN_CACHE_DIR", raising=False) + monkeypatch.delenv("CUSTOM_TIKTOKEN_CACHE_DIR", raising=False) + yield + monkeypatch.delenv("TIKTOKEN_CACHE_DIR", raising=False) + monkeypatch.delenv("CUSTOM_TIKTOKEN_CACHE_DIR", raising=False) + importlib.reload(default_encoding) + + +@pytest.mark.usefixtures("reload_default_encoding") +def test_tiktoken_cache_dir_defaults_to_bundled_tokenizers_for_non_root(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_NON_ROOT", "true") + importlib.reload(default_encoding) + assert Path(os.environ["TIKTOKEN_CACHE_DIR"]) == BUNDLED_TOKENIZERS + assert BUNDLED_TOKENIZERS.name == "tokenizers" + assert default_encoding.encoding.name == "cl100k_base" + assert default_encoding.encoding.decode(default_encoding.encoding.encode("hello world")) == "hello world" + + +@pytest.mark.usefixtures("reload_default_encoding") +def test_custom_tiktoken_cache_dir_overrides_and_is_created(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + custom_dir: Final = tmp_path / "tiktoken_cache" + monkeypatch.setenv("CUSTOM_TIKTOKEN_CACHE_DIR", str(custom_dir)) + importlib.reload(default_encoding) + assert os.environ["TIKTOKEN_CACHE_DIR"] == str(custom_dir) + assert custom_dir.is_dir() diff --git a/tests/unit/litellm_core_utils/test_duration_parser.py b/tests/unit/litellm_core_utils/test_duration_parser.py index cb9f273a0a7..f5e50d67e71 100644 --- a/tests/unit/litellm_core_utils/test_duration_parser.py +++ b/tests/unit/litellm_core_utils/test_duration_parser.py @@ -1,11 +1,13 @@ import unittest from datetime import datetime, time, timezone +from typing import Final from unittest.mock import patch from zoneinfo import ZoneInfo import litellm.litellm_core_utils.duration_parser as duration_parser from litellm.litellm_core_utils.duration_parser import ( duration_in_seconds, + get_budget_window_start, get_next_standardized_reset_time, ) @@ -23,9 +25,7 @@ class TestStandardizedResetTime(unittest.TestCase): # Weekly reset (7d) - should reset on next Monday wednesday = datetime(2023, 5, 17, 15, 45, 0, tzinfo=timezone.utc) # A Wednesday - weekly_expected = datetime( - 2023, 5, 22, 0, 0, 0, tzinfo=timezone.utc - ) # Next Monday + weekly_expected = datetime(2023, 5, 22, 0, 0, 0, tzinfo=timezone.utc) # Next Monday weekly_result = get_next_standardized_reset_time("7d", wednesday, "UTC") self.assertEqual(weekly_result, weekly_expected) @@ -105,17 +105,13 @@ class TestStandardizedResetTime(unittest.TestCase): # Europe/London (UTC+1): 11:30 PM, so next 15m reset is 11:45 PM london = ZoneInfo("Europe/London") london_expected = datetime(2023, 5, 15, 23, 45, 0, tzinfo=london) - london_result = get_next_standardized_reset_time( - "15m", base_time, "Europe/London" - ) + london_result = get_next_standardized_reset_time("15m", base_time, "Europe/London") self.assertEqual(london_result, london_expected) # Test Bangkok timezone (UTC+7): 5:30 AM next day, so next reset is midnight the day after bangkok = ZoneInfo("Asia/Bangkok") bangkok_expected = datetime(2023, 5, 17, 0, 0, 0, tzinfo=bangkok) - bangkok_result = get_next_standardized_reset_time( - "1d", base_time, "Asia/Bangkok" - ) + bangkok_result = get_next_standardized_reset_time("1d", base_time, "Asia/Bangkok") self.assertEqual(bangkok_result, bangkok_expected) def test_edge_cases(self): @@ -137,16 +133,12 @@ class TestStandardizedResetTime(unittest.TestCase): # 30m near midnight - should roll over to next day midnight_minute_expected = datetime(2023, 5, 16, 0, 0, 0, tzinfo=timezone.utc) - midnight_minute_result = get_next_standardized_reset_time( - "30m", near_midnight, "UTC" - ) + midnight_minute_result = get_next_standardized_reset_time("30m", near_midnight, "UTC") self.assertEqual(midnight_minute_result, midnight_minute_expected) # Invalid timezone - should fall back to UTC invalid_tz_expected = datetime(2023, 5, 16, 0, 0, 0, tzinfo=timezone.utc) - invalid_tz_result = get_next_standardized_reset_time( - "1d", on_hour, "NonExistentTimeZone" - ) + invalid_tz_result = get_next_standardized_reset_time("1d", on_hour, "NonExistentTimeZone") self.assertEqual(invalid_tz_result, invalid_tz_expected) def test_iana_timezones_previously_unsupported(self): @@ -164,17 +156,13 @@ class TestStandardizedResetTime(unittest.TestCase): sydney = ZoneInfo("Australia/Sydney") # At 15:00 UTC it's 01:00 AEST May 16 → next midnight is May 17 00:00 AEST sydney_expected = datetime(2023, 5, 17, 0, 0, 0, tzinfo=sydney) - sydney_result = get_next_standardized_reset_time( - "1d", base_time, "Australia/Sydney" - ) + sydney_result = get_next_standardized_reset_time("1d", base_time, "Australia/Sydney") self.assertEqual(sydney_result, sydney_expected) # America/Chicago (UTC-5): at 15:00 UTC it's 10:00 CDT → next midnight is May 16 00:00 CDT chicago = ZoneInfo("America/Chicago") chicago_expected = datetime(2023, 5, 16, 0, 0, 0, tzinfo=chicago) - chicago_result = get_next_standardized_reset_time( - "1d", base_time, "America/Chicago" - ) + chicago_result = get_next_standardized_reset_time("1d", base_time, "America/Chicago") self.assertEqual(chicago_result, chicago_expected) def test_dst_fall_back(self): @@ -209,107 +197,77 @@ class TestResetTimeOfDay(unittest.TestCase): def test_daily_reset_before_offset_is_today(self): now = datetime(2023, 5, 15, 8, 0, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "1d", now, "UTC", reset_time_of_day=time(12, 0) - ) + result = get_next_standardized_reset_time("1d", now, "UTC", reset_time_of_day=time(12, 0)) self.assertEqual(result, datetime(2023, 5, 15, 12, 0, 0, tzinfo=timezone.utc)) def test_daily_reset_after_offset_is_tomorrow(self): now = datetime(2023, 5, 15, 14, 0, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "1d", now, "UTC", reset_time_of_day=time(12, 0) - ) + result = get_next_standardized_reset_time("1d", now, "UTC", reset_time_of_day=time(12, 0)) self.assertEqual(result, datetime(2023, 5, 16, 12, 0, 0, tzinfo=timezone.utc)) def test_daily_reset_exactly_at_offset_rolls_forward(self): now = datetime(2023, 5, 15, 12, 0, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "1d", now, "UTC", reset_time_of_day=time(12, 0) - ) + result = get_next_standardized_reset_time("1d", now, "UTC", reset_time_of_day=time(12, 0)) self.assertEqual(result, datetime(2023, 5, 16, 12, 0, 0, tzinfo=timezone.utc)) def test_daily_reset_with_seconds_offset(self): now = datetime(2023, 5, 15, 8, 0, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "1d", now, "UTC", reset_time_of_day=time(9, 30, 15) - ) + result = get_next_standardized_reset_time("1d", now, "UTC", reset_time_of_day=time(9, 30, 15)) self.assertEqual(result, datetime(2023, 5, 15, 9, 30, 15, tzinfo=timezone.utc)) def test_offset_applies_in_configured_timezone(self): # 2023-05-15 22:30 UTC == 2023-05-16 01:30 in Jerusalem (IDT, UTC+3), # so the next noon-Jerusalem reset is 2023-05-16 12:00 IDT. now = datetime(2023, 5, 15, 22, 30, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "1d", now, "Asia/Jerusalem", reset_time_of_day=time(12, 0) - ) + result = get_next_standardized_reset_time("1d", now, "Asia/Jerusalem", reset_time_of_day=time(12, 0)) jerusalem = result.astimezone(ZoneInfo("Asia/Jerusalem")) - self.assertEqual( - (jerusalem.year, jerusalem.month, jerusalem.day), (2023, 5, 16) - ) + self.assertEqual((jerusalem.year, jerusalem.month, jerusalem.day), (2023, 5, 16)) self.assertEqual(jerusalem.hour, 12) self.assertEqual(jerusalem.minute, 0) def test_weekly_reset_lands_on_monday_at_offset(self): wednesday = datetime(2023, 5, 17, 15, 45, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "7d", wednesday, "UTC", reset_time_of_day=time(12, 0) - ) + result = get_next_standardized_reset_time("7d", wednesday, "UTC", reset_time_of_day=time(12, 0)) self.assertEqual(result, datetime(2023, 5, 22, 12, 0, 0, tzinfo=timezone.utc)) def test_weekly_reset_today_is_monday_before_offset_is_today(self): monday_morning = datetime(2023, 5, 22, 9, 0, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "7d", monday_morning, "UTC", reset_time_of_day=time(12, 0) - ) + result = get_next_standardized_reset_time("7d", monday_morning, "UTC", reset_time_of_day=time(12, 0)) self.assertEqual(result, datetime(2023, 5, 22, 12, 0, 0, tzinfo=timezone.utc)) def test_weekly_reset_today_is_monday_after_offset_is_next_week(self): monday_afternoon = datetime(2023, 5, 22, 15, 0, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "7d", monday_afternoon, "UTC", reset_time_of_day=time(12, 0) - ) + result = get_next_standardized_reset_time("7d", monday_afternoon, "UTC", reset_time_of_day=time(12, 0)) self.assertEqual(result, datetime(2023, 5, 29, 12, 0, 0, tzinfo=timezone.utc)) def test_monthly_30d_lands_on_first_at_offset(self): now = datetime(2023, 5, 15, 10, 30, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "30d", now, "UTC", reset_time_of_day=time(12, 0) - ) + result = get_next_standardized_reset_time("30d", now, "UTC", reset_time_of_day=time(12, 0)) self.assertEqual(result, datetime(2023, 6, 1, 12, 0, 0, tzinfo=timezone.utc)) def test_monthly_1mo_today_is_first_before_offset_is_today(self): now = datetime(2023, 5, 1, 9, 0, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "1mo", now, "UTC", reset_time_of_day=time(12, 0) - ) + result = get_next_standardized_reset_time("1mo", now, "UTC", reset_time_of_day=time(12, 0)) self.assertEqual(result, datetime(2023, 5, 1, 12, 0, 0, tzinfo=timezone.utc)) def test_monthly_year_rollover_at_offset(self): now = datetime(2023, 12, 15, 9, 0, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "1mo", now, "UTC", reset_time_of_day=time(12, 0) - ) + result = get_next_standardized_reset_time("1mo", now, "UTC", reset_time_of_day=time(12, 0)) self.assertEqual(result, datetime(2024, 1, 1, 12, 0, 0, tzinfo=timezone.utc)) def test_custom_day_reset_applies_offset(self): now = datetime(2023, 5, 15, 10, 30, 0, tzinfo=timezone.utc) - result = get_next_standardized_reset_time( - "3d", now, "UTC", reset_time_of_day=time(12, 0) - ) + result = get_next_standardized_reset_time("3d", now, "UTC", reset_time_of_day=time(12, 0)) self.assertEqual(result, datetime(2023, 5, 18, 12, 0, 0, tzinfo=timezone.utc)) def test_sub_day_durations_ignore_offset(self): base = datetime(2023, 5, 15, 15, 20, 30, tzinfo=timezone.utc) self.assertEqual( - get_next_standardized_reset_time( - "2h", base, "UTC", reset_time_of_day=time(12, 0) - ), + get_next_standardized_reset_time("2h", base, "UTC", reset_time_of_day=time(12, 0)), datetime(2023, 5, 15, 16, 0, 0, tzinfo=timezone.utc), ) self.assertEqual( - get_next_standardized_reset_time( - "30m", base, "UTC", reset_time_of_day=time(12, 0) - ), + get_next_standardized_reset_time("30m", base, "UTC", reset_time_of_day=time(12, 0)), datetime(2023, 5, 15, 15, 30, 0, tzinfo=timezone.utc), ) @@ -385,5 +343,35 @@ class TestWordFormBudgetDurations(unittest.TestCase): self.assertIn("garbage", mock_warning.call_args.args) +class TestGetBudgetWindowStart(unittest.TestCase): + def test_window_start_is_the_previous_reset_on_the_same_schedule(self): + created_at: Final = datetime(2024, 10, 1, 0, 30, tzinfo=timezone.utc) + durations: Final = "1d 24h daily 7d weekly 2w 10d 30d monthly 1mo 4h 5h 30m 45s 1hr fortnightly".split() + for duration in durations: + with self.subTest(duration=duration): + reset_at = get_next_standardized_reset_time(duration, created_at, "UTC") + window_start = get_budget_window_start(duration, reset_at) + self.assertLessEqual(window_start, created_at) + self.assertEqual(get_next_standardized_reset_time(duration, window_start, "UTC"), reset_at) + + def test_thirty_days_spans_the_calendar_month_it_resets_on(self): + self.assertEqual( + get_budget_window_start("30d", datetime(2024, 3, 1, tzinfo=timezone.utc)), + datetime(2024, 2, 1, tzinfo=timezone.utc), + ) + + def test_months_clamp_to_the_shorter_month(self): + self.assertEqual( + get_budget_window_start("1mo", datetime(2024, 3, 31, tzinfo=timezone.utc)), + datetime(2024, 2, 29, tzinfo=timezone.utc), + ) + + def test_unrecognized_duration_is_the_one_day_window_the_scheduler_falls_back_to(self): + self.assertEqual( + get_budget_window_start("1hr", datetime(2024, 1, 16, tzinfo=timezone.utc)), + datetime(2024, 1, 15, tzinfo=timezone.utc), + ) + + if __name__ == "__main__": unittest.main() diff --git a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py index 7e7f1b536f8..1bae82014cf 100644 --- a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py +++ b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py @@ -1624,6 +1624,16 @@ def test_guardrail_block_raised_inside_an_llm_call_is_returned_unmapped(block: E assert returned is block +@pytest.mark.parametrize("provider", ["bedrock", "bedrock_mantle"]) +@pytest.mark.parametrize( + "failure", [ImportError("Run 'pip install boto3'."), ModuleNotFoundError(name="unrelated_dependency")] +) +def test_bedrock_import_errors_preserve_the_original_exception(provider, failure): + assert exception_type( + model="test-model", original_exception=failure, custom_llm_provider=provider + ) is failure + + def test_guardrail_provider_failure_status_is_still_mapped(): upstream_failure = HTTPException(status_code=401, detail={"error": "guardrail provider rejected the key"}) @@ -2374,3 +2384,18 @@ def _pre_call_utils_httpx( original_function = litellm.atext_completion return data, original_function, mapped_target + + +@pytest.mark.parametrize("provider", ["sagemaker", "sagemaker_chat", "aws_polly", "openai"]) +@pytest.mark.parametrize("dependency", ["boto3", "botocore"]) +def test_missing_aws_dependency_is_not_mapped_to_provider_failure(provider, dependency): + failure = ModuleNotFoundError(f"No module named '{dependency}'", name=dependency) + assert exception_type( + model="test-model", original_exception=failure, custom_llm_provider=provider + ) is failure + + +@pytest.mark.parametrize("failure", [ImportError("broken import"), ModuleNotFoundError(name="unrelated_dependency")]) +def test_non_aws_import_failure_keeps_provider_mapping(failure): + with pytest.raises(litellm.APIConnectionError): + exception_type(model="test-model", original_exception=failure, custom_llm_provider="openai") diff --git a/tests/unit/litellm_core_utils/test_health_check_helpers.py b/tests/unit/litellm_core_utils/test_health_check_helpers.py index ffe50c137d2..c0298bc1c80 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -1,7 +1,6 @@ """Test health check helper functions""" import json -import os import socket import struct import zlib @@ -1235,3 +1234,193 @@ async def test_health_check_with_custom_llm_provider( assert "error" not in response, response assert upstream.called assert json.loads(upstream.calls[0].request.content)["model"] == "deepseek-r1-distill-qwen-1.5B-q4" + + +@pytest.mark.asyncio +async def test_azure_chat_health_check_surfaces_provider_rate_limit_headers( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post( + url__regex=r"https://resource\.example/openai/deployments/gpt-4\.1-mini/chat/completions.*" + ).respond( + json={ + "id": "chatcmpl-health", + "object": "chat.completion", + "created": 1, + "model": "gpt-4.1-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }, + headers={"x-ratelimit-remaining-tokens": "42"}, + ) + + response: Final = await ahealth_check( + { + "model": "azure/gpt-4.1-mini", + "api_key": "fake-key", + "api_base": "https://resource.example", + "api_version": "2024-06-01", + }, + mode="chat", + ) + + assert response["x-ratelimit-remaining-tokens"] == "42" + assert upstream.called + + +@pytest.mark.asyncio +async def test_azure_embedding_health_check_surfaces_provider_rate_limit_headers( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post( + url__regex=r"https://resource\.example/openai/deployments/text-embedding-ada-002/embeddings.*" + ).respond( + json={ + "object": "list", + "data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}], + "model": "text-embedding-ada-002", + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + }, + headers={"x-ratelimit-remaining-tokens": "84"}, + ) + + response: Final = await ahealth_check( + { + "model": "azure/text-embedding-ada-002", + "api_key": "fake-key", + "api_base": "https://resource.example", + "api_version": "2024-06-01", + }, + input=["health check"], + mode="embedding", + ) + + assert response["x-ratelimit-remaining-tokens"] == "84" + assert upstream.called + + +@pytest.mark.asyncio +async def test_image_generation_health_check_returns_a_successful_response( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post("https://api.openai.com/v1/images/generations").respond( + json={"created": 1, "data": [{"b64_json": "AA=="}]} + ) + + response: Final = await ahealth_check( + {"model": "gpt-image-1", "api_key": "fake-key"}, + mode="image_generation", + prompt="health check", + ) + + assert "error" not in response + assert upstream.called + + +@pytest.mark.asyncio +async def test_groq_wildcard_health_check_uses_a_concrete_model( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr( + litellm, + "models_by_provider", + {"groq": ["groq/openai/gpt-oss-20b"]}, + ) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post("https://api.groq.com/openai/v1/chat/completions").respond( + json={ + "id": "chatcmpl-health", + "object": "chat.completion", + "created": 1, + "model": "groq/openai/gpt-oss-20b", + "service_tier": "on_demand", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "2"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } + ) + + response: Final = await ahealth_check( + { + "model": "groq/*", + "api_key": "fake-key", + "messages": [{"role": "user", "content": "What is 1 + 1?"}], + } + ) + + assert upstream.called + assert json.loads(upstream.calls.last.request.content)["model"] == "openai/gpt-oss-20b" + assert response == {} + + +@pytest.mark.asyncio +async def test_cohere_rerank_health_check_returns_a_successful_response( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post("https://api.cohere.com/v2/rerank").respond( + json={ + "id": "rerank-health", + "results": [{"index": 0, "relevance_score": 0.7}], + "meta": {"billed_units": {"search_units": 1}}, + } + ) + + response: Final = await ahealth_check( + {"model": "cohere/rerank-english-v3.0", "api_key": "fake-key"}, + mode="rerank", + prompt="health check", + ) + + assert "error" not in response + assert upstream.called + assert json.loads(upstream.calls.last.request.content)["query"] == "health check" + + +@pytest.mark.asyncio +async def test_audio_speech_health_check_returns_audio( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post("https://api.openai.com/v1/audio/speech").respond( + content=b"audio", + headers={"content-type": "audio/mpeg"}, + ) + + response: Final = await ahealth_check( + {"model": "openai/tts-1", "api_key": "fake-key"}, + mode="audio_speech", + prompt="health check", + ) + + assert "error" not in response + assert upstream.called + assert json.loads(upstream.calls.last.request.content)["input"] == "health check" + + +@pytest.mark.asyncio +async def test_audio_transcription_health_check_returns_transcribed_text( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post("https://api.openai.com/v1/audio/transcriptions").respond( + json={"text": "health check audio"} + ) + + response: Final = await ahealth_check( + {"model": "openai/whisper-1", "api_key": "fake-key"}, + mode="audio_transcription", + ) + + assert "error" not in response + assert upstream.called + assert b'name="file"' in upstream.calls.last.request.content diff --git a/tests/unit/litellm_core_utils/test_moderation_standard_logging.py b/tests/unit/litellm_core_utils/test_moderation_standard_logging.py new file mode 100644 index 00000000000..1f8c722785f --- /dev/null +++ b/tests/unit/litellm_core_utils/test_moderation_standard_logging.py @@ -0,0 +1,88 @@ +import asyncio +from typing import Final, Literal, TypedDict + +from typing_extensions import ReadOnly + +import httpx +import pytest +import respx +from pydantic import TypeAdapter + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.router import Router + +_MODERATIONS_URL: Final = "https://api.openai.com/v1/moderations" +_MODEL_GROUP: Final = "internal-moderation-model" +_INPUT: Final = "Hello, how are you?" +_CATEGORIES: Final = ("harassment", "hate", "self-harm", "sexual", "violence") +_MODERATION_RESPONSE: Final = { + "id": "modr-logging", + "model": "omni-moderation-latest", + "results": [ + { + "flagged": False, + "categories": {name: False for name in _CATEGORIES}, + "category_scores": {name: 0.001 for name in _CATEGORIES}, + } + ], +} + + +class _LoggedModeration(TypedDict): + call_type: ReadOnly[str] + status: ReadOnly[str] + custom_llm_provider: ReadOnly[str | None] + messages: ReadOnly[object] + response: ReadOnly[object] + model_group: ReadOnly[str | None] + + +class _ModerationRecorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.payloads: Final[list[_LoggedModeration]] = [] + self.logged: Final = asyncio.Event() + + async def async_log_success_event( + self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + self.payloads.append(TypeAdapter(_LoggedModeration).validate_python(kwargs["standard_logging_object"])) + self.logged.set() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("caller", ["default-model", "named-model", "router-group"]) +async def test_moderation_call_is_logged_as_an_amoderation_standard_payload( + caller: Literal["default-model", "named-model", "router-group"], + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setenv("OPENAI_API_KEY", "sk-unit-test") + recorder: Final = _ModerationRecorder() + monkeypatch.setattr(litellm, "callbacks", [recorder]) + respx_mock.post(_MODERATIONS_URL).mock(return_value=httpx.Response(200, json=_MODERATION_RESPONSE)) + router: Final = Router( + model_list=[{"model_name": _MODEL_GROUP, "litellm_params": {"model": "openai/omni-moderation-latest"}}] + ) + + response: Final = ( + await router.amoderation(input=_INPUT, model=_MODEL_GROUP) + if caller == "router-group" + else await litellm.amoderation( + input=_INPUT, model=None if caller == "default-model" else "omni-moderation-latest" + ) + ) + await asyncio.wait_for(recorder.logged.wait(), timeout=10) + + payload: Final = recorder.payloads[-1] + assert payload["call_type"] == litellm.utils.CallTypes.amoderation.value + assert payload["status"] == "success" + assert payload["custom_llm_provider"] == litellm.LlmProviders.OPENAI.value + assert TypeAdapter(tuple[dict[str, str], ...]).validate_python(payload["messages"])[0]["content"] == _INPUT + assert dict(TypeAdapter(dict[str, object]).validate_python(payload["response"])) == response.model_dump() + if caller == "router-group": + assert payload["model_group"] == _MODEL_GROUP + else: + assert not payload["model_group"] diff --git a/tests/unit/litellm_core_utils/test_optional_imports.py b/tests/unit/litellm_core_utils/test_optional_imports.py new file mode 100644 index 00000000000..8b6cc539885 --- /dev/null +++ b/tests/unit/litellm_core_utils/test_optional_imports.py @@ -0,0 +1,39 @@ +import builtins +from typing import Final +from unittest.mock import patch + +import pytest + +from litellm.litellm_core_utils.optional_imports import ensure_optional_import + + +@pytest.mark.parametrize("module", ["boto3", "botocore", "tokenizers"]) +def test_missing_optional_dependency_names_the_installable_package(module: str) -> None: + with patch.dict("sys.modules", {module: None}): + with pytest.raises(ModuleNotFoundError) as caught: + ensure_optional_import(module) + package: Final = "boto3" if module == "botocore" else module + assert str(caught.value) == f"Missing optional dependency '{module}'. Run 'pip install {package}'." + assert caught.value.name == module + assert isinstance(caught.value.__cause__, ModuleNotFoundError) + + +@pytest.mark.parametrize( + "failure", [ModuleNotFoundError(name="unrelated_dependency"), ImportError("broken installation")] +) +def test_optional_import_preserves_unrelated_failure(failure: ImportError) -> None: + original_import: Final = builtins.__import__ + + def import_dependency(name, *args, **kwargs): + if name == "botocore": + raise failure + return original_import(name, *args, **kwargs) + + with patch("builtins.__import__", side_effect=import_dependency): + with pytest.raises(ImportError) as caught: + ensure_optional_import("botocore") + assert caught.value is failure + + +def test_available_optional_dependency_does_not_raise() -> None: + assert ensure_optional_import("json") is None diff --git a/tests/unit/litellm_core_utils/test_stream_usage_logging.py b/tests/unit/litellm_core_utils/test_stream_usage_logging.py new file mode 100644 index 00000000000..a108af09dd1 --- /dev/null +++ b/tests/unit/litellm_core_utils/test_stream_usage_logging.py @@ -0,0 +1,138 @@ +import asyncio +import json +from typing import Final, TypedDict + +from typing_extensions import ReadOnly + +import httpx +import pytest +import respx +from pydantic import TypeAdapter + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.redact_messages import REDACTED_BY_LITELLM +from litellm.types.utils import Usage + +_OPENAI_URL: Final = "https://api.openai.com/v1/chat/completions" +_BODY: Final = TypeAdapter(dict[str, object]) +_PROMPT_TOKENS: Final = 607 +_COMPLETION_TOKENS: Final = 23 + + +def _sse(chunks: tuple[dict[str, object], ...]) -> str: + return "".join(f"data: {json.dumps(chunk)}\n\n" for chunk in chunks) + "data: [DONE]\n\n" + + +def _chunk(delta: dict[str, str], finish_reason: str | None) -> dict[str, object]: + return { + "id": "chatcmpl-usage", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": "gpt-5.5", + "choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}], + } + + +_STREAM: Final = _sse( + ( + _chunk({"role": "assistant", "content": "I am"}, None), + _chunk({"content": " well"}, "stop"), + { + "id": "chatcmpl-usage", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": "gpt-5.5", + "choices": [], + "usage": { + "prompt_tokens": _PROMPT_TOKENS, + "completion_tokens": _COMPLETION_TOKENS, + "total_tokens": _PROMPT_TOKENS + _COMPLETION_TOKENS, + }, + }, + ) +) + + +class _LoggedUsage(TypedDict): + prompt_tokens: ReadOnly[int] + completion_tokens: ReadOnly[int] + total_tokens: ReadOnly[int] + messages: ReadOnly[object] + + +class _UsageRecorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.payloads: Final[list[_LoggedUsage]] = [] + self.logged: Final = asyncio.Event() + + async def async_log_success_event( + self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + self.payloads.append(TypeAdapter(_LoggedUsage).validate_python(kwargs["standard_logging_object"])) + self.logged.set() + + +async def _stream_and_record( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter, include_usage: bool +) -> tuple[Usage, _LoggedUsage, dict[str, object]]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + recorder: Final = _UsageRecorder() + monkeypatch.setattr(litellm, "callbacks", [recorder]) + route: Final = respx_mock.post(_OPENAI_URL).mock( + return_value=httpx.Response(200, text=_STREAM, headers={"content-type": "text/event-stream"}) + ) + stream_options: Final = {"stream_options": {"include_usage": True}} if include_usage else {} + response: Final = await litellm.acompletion( + model="gpt-5.5", + messages=[{"role": "user", "content": "Hello, how are you?" * 100}], + stream=True, + api_key="sk-unit-test", + **stream_options, + ) + usages: Final = tuple([chunk.usage async for chunk in response if getattr(chunk, "usage", None) is not None]) + await asyncio.wait_for(recorder.logged.wait(), timeout=10) + return usages[-1], recorder.payloads[-1], _BODY.validate_json(route.calls.last.request.content) + + +def _assert_logged_usage_matches(client_usage: Usage, payload: _LoggedUsage) -> None: + assert client_usage.prompt_tokens == _PROMPT_TOKENS + assert client_usage.completion_tokens == _COMPLETION_TOKENS + assert (payload["prompt_tokens"], payload["completion_tokens"], payload["total_tokens"]) == ( + client_usage.prompt_tokens, + client_usage.completion_tokens, + client_usage.total_tokens, + ) + + +@pytest.mark.asyncio +async def test_logged_stream_usage_equals_the_final_chunk_usage_with_include_usage( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + client_usage, payload, body = await _stream_and_record(monkeypatch, respx_mock, include_usage=True) + + assert body["stream_options"] == {"include_usage": True} + _assert_logged_usage_matches(client_usage, payload) + + +@pytest.mark.asyncio +async def test_logged_stream_usage_equals_the_usage_chunk_without_stream_options( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + client_usage, payload, body = await _stream_and_record(monkeypatch, respx_mock, include_usage=False) + + assert body["stream_options"] == {"include_usage": True} + _assert_logged_usage_matches(client_usage, payload) + + +@pytest.mark.asyncio +async def test_logged_stream_usage_survives_message_redaction( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + + client_usage, payload, _ = await _stream_and_record(monkeypatch, respx_mock, include_usage=False) + + _assert_logged_usage_matches(client_usage, payload) + assert payload["messages"] == [{"role": "user", "content": REDACTED_BY_LITELLM}] diff --git a/tests/unit/litellm_core_utils/test_tokenizer.py b/tests/unit/litellm_core_utils/test_tokenizer.py index 9d08442b164..354c91920fc 100644 --- a/tests/unit/litellm_core_utils/test_tokenizer.py +++ b/tests/unit/litellm_core_utils/test_tokenizer.py @@ -401,3 +401,138 @@ def test_huggingface_encoding_exposes_the_tokenizers_lookup_and_mutation_surface assert merged.offsets == type(expected).merge([expected, reference.encode("more")]).offsets with pytest.raises(ValueError, match="direction"): actual.pad(8, direction="sideways") + + +@pytest.mark.parametrize("log_level", ["WARNING", "ERROR"]) +def test_missing_python_tokenizer_warns_before_approximate_count(caplog, monkeypatch, log_level): + from unittest.mock import patch + from litellm.utils import _load_huggingface_tokenizer, _select_tokenizer_helper + + monkeypatch.setenv("LITELLM_RUST", "false") + _load_huggingface_tokenizer.cache_clear() + with caplog.at_level(log_level, logger="LiteLLM"), patch.dict(sys.modules, {"tokenizers": None}): + result = _select_tokenizer_helper("llama-2") + assert result["type"] == "openai_tokenizer" + assert result["tokenizer"].encode("hello") + assert ("token counts may be approximate" in caplog.text) is (log_level == "WARNING") + assert ("install tokenizers" in caplog.text) is (log_level == "WARNING") + + +@pytest.mark.parametrize("python_installed", [False, True]) +def test_runtime_aliases_accept_available_tokenizer_instances(python_installed): + from contextlib import nullcontext + from unittest.mock import patch + from litellm.litellm_core_utils import tokenizer as types + + native = HuggingFaceTokenizer.from_str(TOKENIZER_JSON) + python = ReferenceTokenizer.from_str(TOKENIZER_JSON) + with nullcontext() if python_installed else patch.dict(sys.modules, {"tokenizers": None}): + assert isinstance(native, types.HuggingFace) + assert isinstance(native, types.Tokenizer) + assert isinstance(tiktoken.get_encoding("cl100k_base"), types.Tokenizer) + assert isinstance(OpenAIEncoding.from_tiktoken("cl100k_base"), types.Tokenizer) + assert isinstance(python, types.HuggingFace) is python_installed + assert isinstance(python, types.Tokenizer) is python_installed + + +def test_runtime_alias_does_not_hide_broken_tokenizer_installation(): + from unittest.mock import patch + from litellm.litellm_core_utils import tokenizer as types + + failure = ModuleNotFoundError("broken installation", name="tokenizer_dependency") + with patch("builtins.__import__", side_effect=failure): + with pytest.raises(ModuleNotFoundError) as error: + getattr(types, "Tokenizer") + assert error.value is failure + + +def test_unknown_tokenizer_export_raises_attribute_error(): + from litellm.litellm_core_utils import tokenizer as types + + with pytest.raises(AttributeError, match="unknown_tokenizer"): + getattr(types, "unknown_tokenizer") + + +@pytest.mark.parametrize("python_installed", [False, True]) +def test_added_token_return_annotation_resolves_without_optional_import(python_installed): + from contextlib import nullcontext + from typing import get_type_hints + from unittest.mock import patch + + with nullcontext() if python_installed else patch.dict(sys.modules, {"tokenizers": None}): + hints = get_type_hints(HuggingFaceTokenizer.get_added_tokens_decoder) + assert "return" in hints + decoder = HuggingFaceTokenizer.from_str(TOKENIZER_JSON).get_added_tokens_decoder() + for token in decoder.values(): + assert isinstance(token.content, str) + assert isinstance(token.special, bool) + + +def test_tokenizer_fallback_logs_safe_diagnostic_context(caplog, monkeypatch): + from types import SimpleNamespace + from unittest.mock import patch + from litellm.utils import _load_huggingface_tokenizer, _select_tokenizer_helper + + monkeypatch.setenv("LITELLM_RUST", "false") + _load_huggingface_tokenizer.cache_clear() + model = "llama-2\r\nforged-model\x1b[31m\u2028\u2029" + secret = "sk-" + "x" * 48 + failure = OSError("download failed\r\nforged-error\x1b[31m api_key=" + secret) + + def fail_download(*args, **kwargs): + raise failure + + dependency = SimpleNamespace(Tokenizer=SimpleNamespace(from_pretrained=fail_download)) + with caplog.at_level("WARNING", logger="LiteLLM"), patch.dict(sys.modules, {"tokenizers": dependency}): + result = _select_tokenizer_helper(model) + assert result["type"] == "openai_tokenizer" + assert result["tokenizer"].encode("hello") + message = next(record.getMessage() for record in caplog.records if "token counts may be approximate" in record.getMessage()) + assert "llama-2" in message + assert "download failed" in message + assert "forged-model" in message and "forged-error" in message + assert message.isascii() and message.isprintable() + assert secret not in message + assert "REDACTED" in message + assert "install tokenizers and huggingface-hub" in message + + +@pytest.mark.parametrize("python_installed", [False, True]) +def test_native_added_token_decoder_preserves_fields_without_python_dependency(monkeypatch, python_installed): + from tokenizers import AddedToken + + native = HuggingFaceTokenizer.from_str(TOKENIZER_JSON) + expected = ReferenceTokenizer.from_str(TOKENIZER_JSON).get_added_tokens_decoder() + if not python_installed: + monkeypatch.setitem(sys.modules, "tokenizers", None) + actual = native.get_added_tokens_decoder() + attributes = ("content", "single_word", "lstrip", "rstrip", "normalized", "special") + assert { + token_id: tuple(getattr(token, name) for name in attributes) for token_id, token in actual.items() + } == { + token_id: tuple(getattr(token, name) for name in attributes) for token_id, token in expected.items() + } + assert {token_id: str(token) for token_id, token in actual.items()} == { + token_id: str(token) for token_id, token in expected.items() + } + if python_installed: + assert all(isinstance(token, AddedToken) for token in actual.values()) + + +def test_native_added_token_decoder_preserves_unrelated_import_failure(): + import builtins + from unittest.mock import patch + + native = HuggingFaceTokenizer.from_str(TOKENIZER_JSON) + failure = ModuleNotFoundError(name="broken_tokenizer_dependency") + original_import = builtins.__import__ + + def import_dependency(name, *args, **kwargs): + if name == "tokenizers": + raise failure + return original_import(name, *args, **kwargs) + + with patch("builtins.__import__", side_effect=import_dependency): + with pytest.raises(ModuleNotFoundError) as caught: + native.get_added_tokens_decoder() + assert caught.value is failure diff --git a/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py index 70c6ec44790..7edb531a63b 100644 --- a/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py +++ b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py @@ -655,6 +655,49 @@ class _FakePsycopgConn: return _FakeCursor() +class TestLensRenamePendingCheck: + _DATABASE_URL: Final = "postgresql://litellm:hunter2@localhost:5432/litellm" + + def test_database_failure_text_reaches_the_raised_message(self, monkeypatch: pytest.MonkeyPatch) -> None: + import psycopg + + monkeypatch.setenv("DATABASE_URL", self._DATABASE_URL) + + def refuse(conninfo: str, *, connect_timeout: int, autocommit: bool) -> NoReturn: + raise psycopg.OperationalError("FATAL: sorry, too many clients already") + + with pytest.raises(RuntimeError) as err: + ProxyExtrasDBManager.raise_if_lens_rename_pending(connect=refuse) + assert "FATAL: sorry, too many clients already" in str(err.value) + + def test_password_libpq_echoes_is_redacted_from_the_message(self, monkeypatch: pytest.MonkeyPatch) -> None: + import psycopg + + password: Final = "p%zzword" + database_url: Final = f"postgresql://litellm:{password}@localhost:5432/litellm" + monkeypatch.setenv("DATABASE_URL", database_url) + monkeypatch.delenv("DIRECT_URL", raising=False) + with pytest.raises(psycopg.Error) as libpq: + psycopg.connect(database_url, connect_timeout=10, autocommit=True) + + with pytest.raises(RuntimeError) as err: + ProxyExtrasDBManager.raise_if_lens_rename_pending() + assert password not in str(err.value) + assert str(err.value).endswith(str(libpq.value).strip().replace(password, "REDACTED")) + + def test_legacy_tables_keep_their_own_message(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DATABASE_URL", self._DATABASE_URL) + executed: Final[list[tuple[str, tuple[str]]]] = [] + + def legacy_tables_present(conninfo: str, *, connect_timeout: int, autocommit: bool) -> _FakePsycopgConn: + return _FakePsycopgConn(executed) + + with pytest.raises(RuntimeError) as err: + ProxyExtrasDBManager.raise_if_lens_rename_pending(connect=legacy_tables_present) + assert str(err.value).startswith("Legacy Lens tables exist.") + assert executed[0][1] == ("public",) + + class TestSpendLogsPartitionDetectionSchemaScope: """A same-named LiteLLM_SpendLogs in another schema must not trip the detector: the catalog lookup has to be scoped to Prisma's target schema.""" diff --git a/tests/unit/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/unit/llms/anthropic/chat/test_anthropic_chat_transformation.py index fe0c4c84923..74d68d76813 100644 --- a/tests/unit/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/unit/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -7917,3 +7917,54 @@ def test_is_prompt_caching_enabled(anthropic_messages): custom_llm_provider="anthropic", model="anthropic/claude-sonnet-4-5-20250929", ) + + +def test_calculate_usage_sums_compaction_and_message_iterations(): + usage: Final = AnthropicConfig().calculate_usage( + usage_object={ + "input_tokens": 100, + "output_tokens": 50, + "iterations": [ + {"iteration": 1, "type": "compaction", "input_tokens": 1000, "output_tokens": 500}, + {"iteration": 2, "type": "message", "input_tokens": 100, "output_tokens": 50}, + ], + }, + reasoning_content=None, + ) + assert usage.prompt_tokens == 1100 + assert usage.completion_tokens == 550 + assert usage.total_tokens == 1650 + assert usage.prompt_tokens_details.text_tokens == 1100 + assert usage.iterations is not None + assert len(usage.iterations) == 2 + assert usage.iterations[0]["type"] == "compaction" + + +def test_calculate_usage_sums_cache_tokens_across_compaction_iterations(): + usage: Final = AnthropicConfig().calculate_usage( + usage_object={ + "input_tokens": 100, + "output_tokens": 50, + "iterations": [ + { + "type": "compaction", + "input_tokens": 500, + "output_tokens": 200, + "cache_creation_input_tokens": 50, + "cache_read_input_tokens": 17000, + }, + { + "type": "message", + "input_tokens": 100, + "output_tokens": 50, + "cache_creation_input_tokens": 10, + "cache_read_input_tokens": 20, + }, + ], + }, + reasoning_content=None, + ) + assert usage.prompt_tokens == 17680 + assert usage.completion_tokens == 250 + assert usage.prompt_tokens_details.cache_creation_tokens == 60 + assert usage.prompt_tokens_details.cached_tokens == 17020 diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_router.py b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_router.py new file mode 100644 index 00000000000..c55fa476c06 --- /dev/null +++ b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_router.py @@ -0,0 +1,619 @@ +from __future__ import annotations + +import asyncio +import base64 +import json +import struct +import uuid +from collections.abc import AsyncIterable, Mapping +from typing import Final +from zlib import crc32 + +import httpx +import pytest + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.router import Router +from litellm.types.llms.anthropic import ( + AnthropicMessagesTextParam, + AnthropicMessagesTool, + AnthropicMessagesUserMessageParam, + AnthropicToolSearchToolRegex, +) +from litellm.types.utils import StandardLoggingPayload + +_ALIAS: Final = "claude-special-alias" +_ANTHROPIC_MODEL: Final = "claude-haiku-4-5-20251001" +_BEDROCK_MODEL: Final = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" +_BEDROCK_SONNET: Final = "us.anthropic.claude-sonnet-4-5-20250929-v1:0" +_OPENAI_MODEL: Final = "openai/gpt-4.1-mini" +_JOKE: Final = "why did the chicken cross the road" +_PROMPT: Final = "Hello, can you tell me a short joke?" +_ANTHROPIC_URL: Final = r".*api\.anthropic\.com/v1/messages.*" +_OPENAI_RESPONSES_URL: Final = r".*api\.openai\.com/v1/responses.*" +_BEDROCK_INVOKE_URL: Final = r".*bedrock-runtime.*/invoke$" +_BEDROCK_INVOKE_STREAM_URL: Final = r".*bedrock-runtime.*/invoke-with-response-stream$" +_BEDROCK_CONVERSE_URL: Final = r".*bedrock-runtime.*/converse$" +_BEDROCK_CONVERSE_STREAM_URL: Final = r".*bedrock-runtime.*/converse-stream$" + + +def _anthropic_body(model: str = _ANTHROPIC_MODEL, msg_id: str = "msg_1") -> Mapping[str, object]: + return { + "id": msg_id, + "type": "message", + "role": "assistant", + "model": model, + "content": [{"type": "text", "text": _JOKE}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 20}, + } + + +_CONVERSE_BODY: Final = { + "output": {"message": {"role": "assistant", "content": [{"text": _JOKE}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 10, "outputTokens": 20, "totalTokens": 30}, +} + +_OPENAI_BODY: Final = { + "id": "resp_1", + "object": "response", + "status": "completed", + "created_at": 1700000000, + "model": "gpt-4.1-mini", + "output": [ + { + "type": "message", + "id": "msg_out_1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": _JOKE, "annotations": []}], + } + ], + "usage": {"input_tokens": 10, "output_tokens": 20, "total_tokens": 30}, +} + +_STREAM_EVENTS: Final = ( + { + "type": "message_start", + "message": { + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": _ANTHROPIC_MODEL, + "content": [], + "usage": { + "input_tokens": 10, + "output_tokens": 1, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + }, + }, + }, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _JOKE}}, + {"type": "content_block_stop", "index": 0}, + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 20}}, + {"type": "message_stop"}, +) + + +class _RecordingLogger(CustomLogger): + def __init__(self, messages: list[AnthropicMessagesUserMessageParam]) -> None: + super().__init__() + self.messages: Final = messages + self.payloads: tuple[StandardLoggingPayload, ...] = () + self.received: Final = asyncio.Event() + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + payload: Final = kwargs.get("standard_logging_object") + if payload is not None and payload["messages"] == self.messages: + self.payloads = (*self.payloads, payload) + self.received.set() + + +def _unique_messages() -> list[AnthropicMessagesUserMessageParam]: + return [{"role": "user", "content": f"{_PROMPT} {uuid.uuid4().hex}"}] + + +def _router(model_name: str, model: str, **router_kwargs: object) -> Router: + return Router( + model_list=[{"model_name": model_name, "litellm_params": {"model": model, "api_key": "fake-key"}}], + **router_kwargs, + ) + + +def _event_frame(event_type: str, payload: Mapping[str, object]) -> bytes: + def header(name: str, value: str) -> bytes: + name_b: Final = name.encode() + value_b: Final = value.encode() + return ( + struct.pack("!B", len(name_b)) + name_b + struct.pack("!B", 7) + struct.pack("!H", len(value_b)) + value_b + ) + + payload_b: Final = json.dumps(payload).encode() + headers_b: Final = ( + header(":event-type", event_type) + + header(":content-type", "application/json") + + header(":message-type", "event") + ) + prelude: Final = struct.pack("!II", 16 + len(headers_b) + len(payload_b), len(headers_b)) + prelude_crc: Final = crc32(prelude) & 0xFFFFFFFF + message: Final = struct.pack("!I", prelude_crc) + headers_b + payload_b + return prelude + message + struct.pack("!I", crc32(message, prelude_crc) & 0xFFFFFFFF) + + +def _invoke_stream_body() -> bytes: + return b"".join( + _event_frame("chunk", {"bytes": base64.b64encode(json.dumps(event).encode()).decode()}) + for event in _STREAM_EVENTS + ) + + +def _anthropic_sse_body() -> bytes: + return "".join(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n" for event in _STREAM_EVENTS).encode() + + +def _converse_stream_body() -> bytes: + return ( + _event_frame("messageStart", {"role": "assistant"}) + + _event_frame("contentBlockDelta", {"contentBlockIndex": 0, "delta": {"text": _JOKE}}) + + _event_frame("contentBlockStop", {"contentBlockIndex": 0}) + + _event_frame("messageStop", {"stopReason": "end_turn"}) + + _event_frame( + "metadata", + { + "usage": { + "inputTokens": 10, + "outputTokens": 20, + "totalTokens": 530, + "cacheReadInputTokens": 500, + "cacheWriteInputTokens": 0, + }, + "metrics": {"latencyMs": 10}, + }, + ) + ) + + +def _set_fake_aws_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "fake") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "fake") + monkeypatch.setenv("AWS_REGION_NAME", "us-east-1") + + +def _sse_events(raw: str) -> tuple[Mapping[str, object], ...]: + return tuple(json.loads(line[len("data: ") :]) for line in raw.splitlines() if line.startswith("data: ")) + + +async def _stream_events(stream: object) -> tuple[Mapping[str, object], ...]: + assert isinstance(stream, AsyncIterable), type(stream) + chunks: Final = [chunk async for chunk in stream] + raw: Final = "".join(chunk.decode() for chunk in chunks if isinstance(chunk, bytes)) + dict_events: Final = tuple(chunk for chunk in chunks if isinstance(chunk, Mapping)) + return _sse_events(raw) + dict_events + + +async def _wait_for_payload(recorder: _RecordingLogger) -> None: + await asyncio.wait_for(recorder.received.wait(), timeout=30.0) + + +def _assert_anthropic_message(response: object, model: str) -> None: + assert isinstance(response, dict), type(response) + assert response["type"] == "message" + assert response["role"] == "assistant" + assert response["model"] == model + assert isinstance(response["id"], str) and response["id"] + block: Final = response["content"][0] + assert isinstance(block, dict), type(block) + assert block["type"] == "text" + assert block["text"] == _JOKE + + +@pytest.fixture(autouse=True) +def _httpx_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + +@pytest.mark.asyncio +async def test_router_aanthropic_messages_non_streaming_posts_anthropic_body(respx_mock): + route: Final = respx_mock.post(url__regex=_ANTHROPIC_URL).mock( + return_value=httpx.Response(200, json=_anthropic_body()) + ) + response: Final = await _router(_ALIAS, _ANTHROPIC_MODEL).aanthropic_messages( + messages=[{"role": "user", "content": _PROMPT}], + model=_ALIAS, + max_tokens=100, + ) + sent: Final = json.loads(route.calls.last.request.read()) + assert sent["model"] == _ANTHROPIC_MODEL + assert sent["max_tokens"] == 100 + assert sent["messages"] == [{"role": "user", "content": _PROMPT}] + _assert_anthropic_message(response, _ANTHROPIC_MODEL) + + +@pytest.mark.asyncio +async def test_router_aanthropic_messages_latency_routing_forwards_user_id(respx_mock): + route: Final = respx_mock.post(url__regex=_ANTHROPIC_URL).mock( + return_value=httpx.Response(200, json=_anthropic_body()) + ) + response: Final = await _router( + _ALIAS, _ANTHROPIC_MODEL, routing_strategy="latency-based-routing" + ).aanthropic_messages( + messages=[{"role": "user", "content": _PROMPT}], + model=_ALIAS, + max_tokens=100, + metadata={"user_id": "hello"}, + ) + sent: Final = json.loads(route.calls.last.request.read()) + assert sent["model"] == _ANTHROPIC_MODEL + assert sent["metadata"] == {"user_id": "hello"} + _assert_anthropic_message(response, _ANTHROPIC_MODEL) + + +@pytest.mark.asyncio +async def test_router_aanthropic_messages_falls_back_to_bedrock_after_anthropic_401(respx_mock, monkeypatch): + _set_fake_aws_env(monkeypatch) + anthropic_route: Final = respx_mock.post(url__regex=_ANTHROPIC_URL).mock( + return_value=httpx.Response( + 401, json={"type": "error", "error": {"type": "authentication_error", "message": "invalid x-api-key"}} + ) + ) + bedrock_route: Final = respx_mock.post(url__regex=_BEDROCK_INVOKE_URL).mock( + return_value=httpx.Response(200, json=_anthropic_body(model=_BEDROCK_SONNET, msg_id="msg_bedrock")) + ) + router: Final = Router( + model_list=[ + { + "model_name": "anthropic/claude-opus-4-7", + "litellm_params": {"model": "anthropic/claude-opus-4-7", "api_key": "bad-key"}, + }, + {"model_name": f"bedrock/{_BEDROCK_SONNET}", "litellm_params": {"model": f"bedrock/{_BEDROCK_SONNET}"}}, + ], + fallbacks=[{"anthropic/claude-opus-4-7": [f"bedrock/{_BEDROCK_SONNET}"]}], + ) + response: Final = await router.aanthropic_messages( + messages=[{"role": "user", "content": _PROMPT}], + model="anthropic/claude-opus-4-7", + max_tokens=100, + metadata={"user_id": "hello"}, + ) + assert anthropic_route.call_count == 1 + assert anthropic_route.calls.last.request.headers["x-api-key"] == "bad-key" + assert bedrock_route.call_count == 1 + assert "authorization" in bedrock_route.calls.last.request.headers + _assert_anthropic_message(response, _BEDROCK_SONNET) + assert response["id"] == "msg_bedrock" + + +@pytest.mark.asyncio +async def test_router_aanthropic_messages_bedrock_converse_and_invoke(respx_mock, monkeypatch): + _set_fake_aws_env(monkeypatch) + converse_route: Final = respx_mock.post(url__regex=_BEDROCK_CONVERSE_URL).mock( + return_value=httpx.Response(200, json=_CONVERSE_BODY) + ) + invoke_route: Final = respx_mock.post(url__regex=_BEDROCK_INVOKE_URL).mock( + return_value=httpx.Response(200, json=_anthropic_body(model=_BEDROCK_SONNET)) + ) + converse_model: Final = f"bedrock/converse/{_BEDROCK_SONNET}" + invoke_model: Final = f"bedrock/{_BEDROCK_SONNET}" + router: Final = Router( + model_list=[ + {"model_name": converse_model, "litellm_params": {"model": converse_model}}, + {"model_name": invoke_model, "litellm_params": {"model": invoke_model}}, + ] + ) + converse_response: Final = await router.aanthropic_messages( + messages=[{"role": "user", "content": _PROMPT}], model=converse_model, max_tokens=100 + ) + invoke_response: Final = await router.aanthropic_messages( + messages=[{"role": "user", "content": _PROMPT}], model=invoke_model, max_tokens=100 + ) + assert converse_route.call_count == 1 + assert invoke_route.call_count == 1 + converse_request: Final = converse_route.calls.last.request + invoke_request: Final = invoke_route.calls.last.request + assert "authorization" in converse_request.headers + assert "authorization" in invoke_request.headers + assert json.loads(converse_request.read())["messages"] == [{"role": "user", "content": [{"text": _PROMPT}]}] + assert json.loads(invoke_request.read())["messages"] == [{"role": "user", "content": _PROMPT}] + _assert_anthropic_message(converse_response, _BEDROCK_SONNET) + _assert_anthropic_message(invoke_response, _BEDROCK_SONNET) + + +def test_sync_openai_bridge_anthropic_messages_returns_content_blocks(respx_mock): + route: Final = respx_mock.post(url__regex=_OPENAI_RESPONSES_URL).mock( + return_value=httpx.Response(200, json=_OPENAI_BODY) + ) + response: Final = litellm.anthropic.messages.create( + messages=[{"role": "user", "content": _PROMPT}], + model=_OPENAI_MODEL, + max_tokens=100, + api_key="fake-key", + ) + sent: Final = json.loads(route.calls.last.request.read()) + assert sent["model"] == "gpt-4.1-mini" + assert isinstance(response, dict) + assert response["content"][0]["text"] == _JOKE + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "url", "body", "expected_model"), + [ + pytest.param(_ANTHROPIC_MODEL, _ANTHROPIC_URL, _anthropic_body(), _ANTHROPIC_MODEL, id="anthropic"), + pytest.param( + _BEDROCK_MODEL, + _BEDROCK_INVOKE_URL, + _anthropic_body(model="us.anthropic.claude-haiku-4-5-20251001-v1:0"), + "us.anthropic.claude-haiku-4-5-20251001-v1:0", + id="bedrock-invoke", + ), + pytest.param(_OPENAI_MODEL, _OPENAI_RESPONSES_URL, _OPENAI_BODY, "gpt-4.1-mini", id="openai-bridge"), + ], +) +async def test_acreate_non_streaming_returns_dict_content_blocks( + respx_mock, monkeypatch, model: str, url: str, body: Mapping[str, object], expected_model: str +): + _set_fake_aws_env(monkeypatch) + route: Final = respx_mock.post(url__regex=url).mock(return_value=httpx.Response(200, json=body)) + response: Final = await litellm.anthropic.messages.acreate( + messages=[{"role": "user", "content": _PROMPT}], + model=model, + max_tokens=100, + api_key="fake-key", + ) + assert route.call_count == 1 + _assert_anthropic_message(response, expected_model) + + +@pytest.mark.asyncio +async def test_router_aanthropic_messages_non_streaming_logs_usage_model_and_cost(respx_mock, monkeypatch): + messages: Final = _unique_messages() + recorder: Final = _RecordingLogger(messages) + monkeypatch.setattr(litellm, "callbacks", [recorder]) + respx_mock.post(url__regex=_ANTHROPIC_URL).mock(return_value=httpx.Response(200, json=_anthropic_body())) + response: Final = await _router(_ALIAS, _ANTHROPIC_MODEL).aanthropic_messages( + messages=messages, model=_ALIAS, max_tokens=100 + ) + await _wait_for_payload(recorder) + assert len(recorder.payloads) == 1 + payload: Final = recorder.payloads[0] + assert payload["status"] == "success" + assert payload["messages"] == messages + assert payload["response"] is not None + assert payload["model"] == _ANTHROPIC_MODEL + assert payload["model_group"] == _ALIAS + assert payload["response_cost"] > 0 + assert payload["prompt_tokens"] == response["usage"]["input_tokens"] == 10 + assert payload["completion_tokens"] == response["usage"]["output_tokens"] == 20 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "url", "content_type", "body", "expected_model"), + [ + pytest.param( + _ANTHROPIC_MODEL, + _ANTHROPIC_URL, + "text/event-stream", + _anthropic_sse_body(), + _ANTHROPIC_MODEL, + id="anthropic", + ), + pytest.param( + _BEDROCK_MODEL, + _BEDROCK_INVOKE_STREAM_URL, + "application/vnd.amazon.eventstream", + _invoke_stream_body(), + _BEDROCK_MODEL, + id="bedrock-invoke", + ), + ], +) +async def test_router_aanthropic_messages_streaming_logs_usage_model_and_cost( + respx_mock, monkeypatch, model: str, url: str, content_type: str, body: bytes, expected_model: str +): + _set_fake_aws_env(monkeypatch) + messages: Final = _unique_messages() + recorder: Final = _RecordingLogger(messages) + monkeypatch.setattr(litellm, "callbacks", [recorder]) + respx_mock.post(url__regex=url).mock( + return_value=httpx.Response(200, content=body, headers={"content-type": content_type}) + ) + stream: Final = await _router(_ALIAS, model).aanthropic_messages( + messages=messages, model=_ALIAS, max_tokens=100, stream=True + ) + events: Final = await _stream_events(stream) + usages: Final = tuple( + event["usage"] if "usage" in event else event["message"]["usage"] + for event in events + if "usage" in event or (event.get("type") == "message_start" and "usage" in event["message"]) + ) + assert usages, events + await _wait_for_payload(recorder) + assert len(recorder.payloads) == 1 + payload: Final = recorder.payloads[0] + assert payload["status"] == "success" + assert payload["messages"] == messages + assert payload["response"] is not None + assert payload["model"] == expected_model + assert payload["response_cost"] > 0 + assert payload["prompt_tokens"] == max(usage.get("input_tokens", 0) for usage in usages) == 10 + assert payload["completion_tokens"] == max(usage.get("output_tokens", 0) for usage in usages) == 20 + + +_LARGE_SYSTEM_PROMPT: Final = "This is a comprehensive legal agreement between Party A and Party B. " * 100 + + +def _cached_system() -> list[AnthropicMessagesTextParam]: + return [{"type": "text", "text": _LARGE_SYSTEM_PROMPT, "cache_control": {"type": "ephemeral"}}] + + +@pytest.mark.asyncio +async def test_bedrock_converse_system_prompt_caching_returns_cache_tokens(respx_mock, monkeypatch): + _set_fake_aws_env(monkeypatch) + route: Final = respx_mock.post(url__regex=_BEDROCK_CONVERSE_URL).mock( + return_value=httpx.Response( + 200, + json={ + **_CONVERSE_BODY, + "usage": { + "inputTokens": 10, + "outputTokens": 20, + "totalTokens": 580, + "cacheReadInputTokens": 500, + "cacheWriteInputTokens": 50, + }, + }, + ) + ) + response: Final = await litellm.anthropic.messages.acreate( + model=f"bedrock/converse/{_BEDROCK_SONNET}", + messages=[{"role": "user", "content": "What are the key terms?"}], + system=_cached_system(), + max_tokens=100, + ) + sent: Final = json.loads(route.calls.last.request.read()) + assert sent["system"] == [{"text": _LARGE_SYSTEM_PROMPT}, {"cachePoint": {"type": "default"}}] + assert isinstance(response, dict) + assert response["usage"]["cache_creation_input_tokens"] == 50 + assert response["usage"]["cache_read_input_tokens"] == 500 + + +@pytest.mark.asyncio +async def test_bedrock_invoke_system_prompt_caching_returns_cache_tokens(respx_mock, monkeypatch): + _set_fake_aws_env(monkeypatch) + route: Final = respx_mock.post(url__regex=_BEDROCK_INVOKE_URL).mock( + return_value=httpx.Response( + 200, + json={ + **_anthropic_body(model=_BEDROCK_SONNET), + "usage": { + "input_tokens": 10, + "output_tokens": 20, + "cache_creation_input_tokens": 50, + "cache_read_input_tokens": 500, + }, + }, + ) + ) + response: Final = await litellm.anthropic.messages.acreate( + model=f"bedrock/invoke/{_BEDROCK_SONNET}", + messages=[{"role": "user", "content": "What are the key terms?"}], + system=_cached_system(), + max_tokens=100, + ) + sent: Final = json.loads(route.calls.last.request.read()) + assert sent["system"] == _cached_system() + assert isinstance(response, dict) + assert response["usage"]["cache_creation_input_tokens"] == 50 + assert response["usage"]["cache_read_input_tokens"] == 500 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "url", "body"), + [ + pytest.param( + f"bedrock/converse/{_BEDROCK_SONNET}", _BEDROCK_CONVERSE_STREAM_URL, _converse_stream_body(), id="converse" + ), + pytest.param( + f"bedrock/invoke/{_BEDROCK_SONNET}", _BEDROCK_INVOKE_STREAM_URL, _invoke_stream_body(), id="invoke" + ), + ], +) +async def test_bedrock_streaming_message_start_carries_cache_usage_fields( + respx_mock, monkeypatch, model: str, url: str, body: bytes +): + _set_fake_aws_env(monkeypatch) + route: Final = respx_mock.post(url__regex=url).mock( + return_value=httpx.Response(200, content=body, headers={"content-type": "application/vnd.amazon.eventstream"}) + ) + stream: Final = await litellm.anthropic.messages.acreate( + model=model, + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": _LARGE_SYSTEM_PROMPT, "cache_control": {"type": "ephemeral"}}, + {"type": "text", "text": "What are the payment terms in this agreement?"}, + ], + } + ], + max_tokens=100, + stream=True, + ) + events: Final = await _stream_events(stream) + assert route.call_count == 1 + message_starts: Final = [event for event in events if event.get("type") == "message_start"] + assert len(message_starts) == 1, events + usage: Final = message_starts[0]["message"]["usage"] + assert "cache_creation_input_tokens" in usage, usage + assert "cache_read_input_tokens" in usage, usage + + +def _tool_search_tools() -> list[AnthropicToolSearchToolRegex | AnthropicMessagesTool]: + def deferred(name: str, description: str, field: str) -> AnthropicMessagesTool: + return { + "name": name, + "description": description, + "input_schema": {"type": "object", "properties": {field: {"type": "string"}}, "required": [field]}, + "defer_loading": True, + } + + return [ + {"type": "tool_search_tool_regex_20251119", "name": "tool_search_tool_regex"}, + deferred("get_weather", "Get the current weather for a location", "location"), + deferred("get_stock_price", "Get the current stock price for a ticker symbol", "ticker"), + deferred("search_web", "Search the web for information", "query"), + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("prompt", "tool_name", "tool_input"), + [ + pytest.param( + "I need to know the current weather in New York City. Please use the appropriate tool.", + "get_weather", + {"location": "New York, NY"}, + id="discovers-weather-tool", + ), + pytest.param( + "What's the stock price of Apple (AAPL)?", "get_stock_price", {"ticker": "AAPL"}, id="multiple-deferred" + ), + ], +) +async def test_tool_search_forwards_deferred_tools_and_beta_header( + respx_mock, prompt: str, tool_name: str, tool_input: Mapping[str, str] +): + route: Final = respx_mock.post(url__regex=_ANTHROPIC_URL).mock( + return_value=httpx.Response( + 200, + json={ + **_anthropic_body(model="claude-sonnet-4-5-20250929", msg_id="msg_tool"), + "content": [ + {"type": "tool_use", "id": "toolu_1", "name": tool_name, "input": tool_input}, + ], + "stop_reason": "tool_use", + }, + ) + ) + response: Final = await litellm.anthropic.messages.acreate( + model="anthropic/claude-sonnet-4-5-20250929", + messages=[{"role": "user", "content": prompt}], + tools=[dict(tool) for tool in _tool_search_tools()], + max_tokens=1024, + api_key="fake-key", + extra_headers={"anthropic-beta": "advanced-tool-use-2025-11-20"}, + ) + request: Final = route.calls.last.request + assert "advanced-tool-use-2025-11-20" in request.headers["anthropic-beta"].split(",") + assert json.loads(request.read())["tools"] == _tool_search_tools() + assert isinstance(response, dict) + assert response["stop_reason"] == "tool_use" + assert [block for block in response["content"] if block["type"] == "tool_use"] == [ + {"type": "tool_use", "id": "toolu_1", "name": tool_name, "input": tool_input} + ] diff --git a/tests/unit/llms/anthropic/pass_through/test_anthropic_native_passthrough_spend_logging.py b/tests/unit/llms/anthropic/pass_through/test_anthropic_native_passthrough_spend_logging.py new file mode 100644 index 00000000000..de2cb03bf48 --- /dev/null +++ b/tests/unit/llms/anthropic/pass_through/test_anthropic_native_passthrough_spend_logging.py @@ -0,0 +1,181 @@ +from __future__ import annotations + +import asyncio +import json +import uuid +from collections.abc import Mapping +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest +from fastapi import Request, Response + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER +from litellm.proxy._types import UserAPIKeyAuth, hash_token +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import anthropic_proxy_route +from litellm.types.utils import StandardLoggingPayload + +_UPSTREAM: Final = "https://api.anthropic.com/v1/messages" +_MODEL: Final = "claude-sonnet-4-5-20250929" +_VIRTUAL_KEY: Final = "sk-native-passthrough" + + +class _RecordingLogger(CustomLogger): + def __init__(self, message_id: str) -> None: + super().__init__() + self.message_id: Final = message_id + self.payloads: tuple[StandardLoggingPayload, ...] = () + self.received: Final = asyncio.Event() + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + payload: Final = kwargs.get("standard_logging_object") + if payload is not None and payload["id"] == self.message_id: + self.payloads = (*self.payloads, payload) + self.received.set() + + +def _proxy_request(body: Mapping[str, object]) -> Request: + request: Final = MagicMock(spec=Request) + request.method = "POST" + request.url = httpx.URL("http://proxy/anthropic/v1/messages") + request.headers = {"content-type": "application/json", "anthropic-version": "2023-06-01"} + request.scope = {"path": "/anthropic/v1/messages", "type": "http", "method": "POST", "headers": []} + request.query_params = {} + request.body = AsyncMock(return_value=json.dumps(body).encode()) + request.json = AsyncMock(return_value=body) + return request + + +def _sse(events: tuple[Mapping[str, object], ...]) -> bytes: + return "".join(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n" for event in events).encode() + + +async def _wait_for_payload(recorder: _RecordingLogger) -> None: + GLOBAL_LOGGING_WORKER.start() + await asyncio.wait_for(recorder.received.wait(), timeout=30.0) + + +def _assert_spend_payload( + payload: StandardLoggingPayload, message_id: str, tags: list[str], prompt_tokens: int, completion_tokens: int +) -> None: + assert payload["id"] == message_id + assert payload["call_type"] == "pass_through_endpoint" + assert payload["status"] == "success" + assert payload["custom_llm_provider"] == "anthropic" + assert payload["model"] == _MODEL + assert payload["prompt_tokens"] == prompt_tokens + assert payload["completion_tokens"] == completion_tokens + assert payload["total_tokens"] == prompt_tokens + completion_tokens + assert payload["response_cost"] > 0 + assert payload["request_tags"] == tags + assert payload["cache_hit"] is not True + assert payload["startTime"] <= payload["endTime"] + assert payload["metadata"]["user_api_key_hash"] == hash_token(_VIRTUAL_KEY) + + +@pytest.fixture +def recorder(monkeypatch: pytest.MonkeyPatch) -> _RecordingLogger: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setenv("ANTHROPIC_API_KEY", "synthetic-anthropic-key") + monkeypatch.delenv("ANTHROPIC_API_BASE", raising=False) + monkeypatch.delenv("ANTHROPIC_BASE_URL", raising=False) + logger: Final = _RecordingLogger(f"msg_{uuid.uuid4().hex}") + monkeypatch.setattr(litellm, "callbacks", [logger]) + monkeypatch.setattr(litellm, "_async_success_callback", [logger]) + return logger + + +@pytest.mark.asyncio +async def test_native_anthropic_passthrough_logs_usage_tags_and_spend(respx_mock, recorder: _RecordingLogger): + tags: Final = ["test-tag-1", "test-tag-2"] + route: Final = respx_mock.post(_UPSTREAM).mock( + return_value=httpx.Response( + 200, + json={ + "id": recorder.message_id, + "type": "message", + "role": "assistant", + "model": _MODEL, + "content": [{"type": "text", "text": "hello test"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 11, "output_tokens": 7}, + }, + ) + ) + response: Final = await anthropic_proxy_route( + endpoint="v1/messages", + request=_proxy_request( + { + "model": _MODEL, + "max_tokens": 10, + "messages": [{"role": "user", "content": "Say 'hello test' and nothing else"}], + "litellm_metadata": {"tags": tags}, + } + ), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=UserAPIKeyAuth(api_key=_VIRTUAL_KEY, token=_VIRTUAL_KEY), + ) + assert response.status_code == 200 + assert json.loads(response.body)["id"] == recorder.message_id + outbound: Final = route.calls.last.request + assert outbound.headers["x-api-key"] == "synthetic-anthropic-key" + assert json.loads(outbound.content) == { + "model": _MODEL, + "max_tokens": 10, + "messages": [{"role": "user", "content": "Say 'hello test' and nothing else"}], + } + await _wait_for_payload(recorder) + assert len(recorder.payloads) == 1 + payload: Final = recorder.payloads[0] + _assert_spend_payload(payload, recorder.message_id, tags, prompt_tokens=11, completion_tokens=7) + assert payload["api_base"] == _UPSTREAM + + +@pytest.mark.asyncio +async def test_native_anthropic_passthrough_streaming_logs_usage_tags_and_spend(respx_mock, recorder: _RecordingLogger): + tags: Final = ["test-tag-stream-1", "test-tag-stream-2"] + events: Final = ( + { + "type": "message_start", + "message": { + "id": recorder.message_id, + "type": "message", + "role": "assistant", + "model": _MODEL, + "content": [], + "usage": {"input_tokens": 11, "output_tokens": 1}, + }, + }, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hello stream test"}}, + {"type": "content_block_stop", "index": 0}, + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 7}}, + {"type": "message_stop"}, + ) + route: Final = respx_mock.post(_UPSTREAM).mock( + return_value=httpx.Response(200, content=_sse(events), headers={"content-type": "text/event-stream"}) + ) + response: Final = await anthropic_proxy_route( + endpoint="v1/messages", + request=_proxy_request( + { + "model": _MODEL, + "max_tokens": 10, + "stream": True, + "messages": [{"role": "user", "content": "Say 'hello stream test' and nothing else"}], + "litellm_metadata": {"tags": tags, "user": "test-user-1"}, + } + ), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=UserAPIKeyAuth(api_key=_VIRTUAL_KEY, token=_VIRTUAL_KEY), + ) + assert response.status_code == 200 + streamed: Final = b"".join([chunk async for chunk in response.body_iterator]) + assert b"hello stream test" in streamed + assert json.loads(route.calls.last.request.content)["stream"] is True + await _wait_for_payload(recorder) + assert len(recorder.payloads) == 1 + _assert_spend_payload(recorder.payloads[0], recorder.message_id, tags, prompt_tokens=11, completion_tokens=7) diff --git a/tests/unit/llms/anthropic/test_count_tokens_oauth.py b/tests/unit/llms/anthropic/test_count_tokens_oauth.py index 96d909a4b3f..7f717091754 100644 --- a/tests/unit/llms/anthropic/test_count_tokens_oauth.py +++ b/tests/unit/llms/anthropic/test_count_tokens_oauth.py @@ -11,6 +11,7 @@ Regression tests for https://github.com/BerriAI/litellm/issues/22040 and for the import os import sys +from typing import Final import httpx import pytest @@ -21,10 +22,15 @@ sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../. import litellm from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION from litellm.llms.anthropic.common_utils import AnthropicModelInfo +from litellm.llms.anthropic.count_tokens.token_counter import AnthropicTokenCounter from litellm.llms.anthropic.count_tokens.transformation import ( AnthropicCountTokensConfig, ) +from litellm.llms.azure_ai.anthropic.count_tokens.token_counter import ( + AzureAIAnthropicTokenCounter, +) from litellm.types.llms.anthropic import ANTHROPIC_OAUTH_BETA_HEADER +from litellm.types.utils import TokenCountResponse # Fake tokens for testing (not real secrets) FAKE_OAUTH_TOKEN = "sk-ant-oat01-fake-token-for-testing-123456789abcdef" @@ -269,3 +275,80 @@ class TestCountTokensUsesWorkloadIdentity: assert result is not None assert result.total_tokens == 7 assert seen["auth_header"] == {"x-api-key": vault_key} + + +@pytest.mark.parametrize( + ("counter_type", "api_base", "endpoint", "tokenizer_type"), + ( + ( + AnthropicTokenCounter, + "https://gateway.example", + "https://gateway.example/v1/messages/count_tokens", + "anthropic_api", + ), + ( + AzureAIAnthropicTokenCounter, + "https://resource.example", + "https://resource.example/anthropic/v1/messages/count_tokens", + "azure_ai_anthropic_api", + ), + ), +) +@pytest.mark.parametrize(("status_code", "expected_error"), ((200, False), (401, True))) +@pytest.mark.asyncio +async def test_count_token_counters_return_typed_success_and_error_responses( + counter_type: type[AnthropicTokenCounter] | type[AzureAIAnthropicTokenCounter], + api_base: str, + endpoint: str, + tokenizer_type: str, + status_code: int, + expected_error: bool, + httpx_transport_clients: None, +) -> None: + response_body: Final = {"input_tokens": 17} if status_code == 200 else {"error": {"message": "invalid key"}} + router: Final = respx.mock + + with router: + count_route: Final = router.post(endpoint).mock(return_value=httpx.Response(status_code, json=response_body)) + result: Final = await counter_type().count_tokens( + model_to_use="claude-test", + messages=[{"role": "user", "content": "hi"}], + contents=None, + deployment={"litellm_params": {"api_key": "sk-ant-api03-test-key", "api_base": api_base}}, + request_model="claude-test", + ) + requests: Final = tuple(router.calls) + + assert count_route.called + assert len(requests) == 1 + assert requests[0].request.url == httpx.URL(endpoint) + assert isinstance(result, TokenCountResponse) + assert result.request_model == "claude-test" + assert result.model_used == "claude-test" + assert result.tokenizer_type == tokenizer_type + + if expected_error: + assert result.error is True + assert result.status_code == status_code + assert result.total_tokens == 0 + return + + assert result.error is not True + assert result.total_tokens == 17 + + +@pytest.mark.parametrize( + ("counter_type", "provider"), + ( + (AnthropicTokenCounter, "anthropic"), + (AzureAIAnthropicTokenCounter, "azure_ai"), + ), +) +def test_count_token_counters_select_their_own_provider( + counter_type: type[AnthropicTokenCounter] | type[AzureAIAnthropicTokenCounter], + provider: str, +) -> None: + counter: Final = counter_type() + + assert counter.should_use_token_counting_api(custom_llm_provider=provider) is True + assert counter.should_use_token_counting_api(custom_llm_provider="unknown") is False diff --git a/tests/unit/llms/aws_polly/text_to_speech/test_transformation.py b/tests/unit/llms/aws_polly/text_to_speech/test_transformation.py new file mode 100644 index 00000000000..4778b6296fc --- /dev/null +++ b/tests/unit/llms/aws_polly/text_to_speech/test_transformation.py @@ -0,0 +1,19 @@ +import json +from typing import Final + +from litellm.llms.aws_polly.text_to_speech.transformation import AWSPollyTextToSpeechConfig + + +def test_installed_botocore_signs_the_speech_request() -> None: + headers, body = AWSPollyTextToSpeechConfig()._sign_polly_request( + request_body={"Text": "ping", "VoiceId": "Joanna"}, + endpoint_url="https://polly.us-west-2.amazonaws.com/v1/speech", + litellm_params={ + "aws_access_key_id": "test-key", + "aws_secret_access_key": "test-secret", + "aws_region_name": "us-west-2", + }, + ) + authorization: Final = headers["Authorization"] + assert "/us-west-2/polly/aws4_request" in authorization + assert json.loads(body) == {"Text": "ping", "VoiceId": "Joanna"} diff --git a/tests/unit/llms/azure/response/test_azure_transformation.py b/tests/unit/llms/azure/response/test_azure_transformation.py index 8db2ba4835e..fea5304a689 100644 --- a/tests/unit/llms/azure/response/test_azure_transformation.py +++ b/tests/unit/llms/azure/response/test_azure_transformation.py @@ -12,8 +12,11 @@ from litellm.llms.azure.responses.o_series_transformation import ( AzureOpenAIOSeriesResponsesAPIConfig, ) from litellm.llms.azure.responses.transformation import AzureOpenAIResponsesAPIConfig +from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager @pytest.mark.serial @@ -920,3 +923,25 @@ async def test_azure_responses_api_headers_with_llm_provider_prefix(): # Also verify openai-compatible headers are included assert "x-ratelimit-limit-tokens" in headers assert "x-ratelimit-remaining-tokens" in headers + + +@pytest.mark.parametrize( + "model", ["gpt-5", "gpt-5-turbo", "GPT-5", "azure/gpt-5", "gpt-3.5-turbo", "gpt-4", "gpt-4-turbo", "gpt-4o"] +) +def test_azure_gpt_models_resolve_to_responses_config_with_temperature(model: str) -> None: + config: Final = ProviderConfigManager.get_provider_responses_api_config(provider=LlmProviders.AZURE, model=model) + assert type(config) is AzureOpenAIResponsesAPIConfig + assert "temperature" in config.get_supported_openai_params(model) + + +@pytest.mark.parametrize("model", ["o1", "o3"]) +def test_azure_o_series_resolves_to_o_series_config_without_temperature(model: str) -> None: + config: Final = ProviderConfigManager.get_provider_responses_api_config(provider=LlmProviders.AZURE, model=model) + assert type(config) is AzureOpenAIOSeriesResponsesAPIConfig + assert "temperature" not in config.get_supported_openai_params(model) + + +def test_openai_gpt5_resolves_to_responses_config_with_temperature() -> None: + config: Final = ProviderConfigManager.get_provider_responses_api_config(provider=LlmProviders.OPENAI, model="gpt-5") + assert type(config) is OpenAIResponsesAPIConfig + assert "temperature" in config.get_supported_openai_params("gpt-5") diff --git a/tests/unit/llms/base_llm/audio_transcription/__init__.py b/tests/unit/llms/base_llm/audio_transcription/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/base_llm/audio_transcription/test_provider_audio_transcription_translation.py b/tests/unit/llms/base_llm/audio_transcription/test_provider_audio_transcription_translation.py new file mode 100644 index 00000000000..65a8ce91e64 --- /dev/null +++ b/tests/unit/llms/base_llm/audio_transcription/test_provider_audio_transcription_translation.py @@ -0,0 +1,210 @@ +import io +import json +from typing import Final, Mapping, Sequence, cast + +import httpx +import pytest +import respx +from respx import MockRouter +from typing_extensions import ReadOnly, TypedDict + +import litellm +from litellm import transcription +from litellm.litellm_core_utils.get_supported_openai_params import get_supported_openai_params +from litellm.llms.base_llm.audio_transcription.transformation import ( + AudioTranscriptionRequestData, + BaseAudioTranscriptionConfig, +) +from litellm.llms.elevenlabs.audio_transcription.transformation import ElevenLabsAudioTranscriptionConfig +from litellm.llms.mistral.audio_transcription.transformation import MistralAudioTranscriptionConfig +from litellm.llms.ovhcloud.audio_transcription.transformation import OVHCloudAudioTranscriptionConfig +from litellm.utils import ProviderConfigManager + + +class _Kwargs(TypedDict, total=False): + model: ReadOnly[str] + api_key: ReadOnly[str] + api_base: ReadOnly[str] + timestamp_granularities: ReadOnly[Sequence[str]] + + +class _Case(TypedDict): + id: ReadOnly[str] + provider: ReadOnly[str] + kwargs: ReadOnly[_Kwargs] + url: ReadOnly[str] + prefix_match: ReadOnly[bool] + base_model: ReadOnly[str] + request_markers: ReadOnly[tuple[bytes, ...]] + marker_in_url: ReadOnly[bool] + config_class: ReadOnly[type[BaseAudioTranscriptionConfig]] + + +_AUDIO_BYTES: Final = b"RIFFFAKEWAVDATA-gettysburg" + +_CASES: Final[tuple[_Case, ...]] = ( + { + "id": "openai_gpt4o", + "provider": "openai", + "kwargs": { + "model": "openai/gpt-4o-transcribe", + "api_key": "sk-offline", + "timestamp_granularities": ["word"], + }, + "url": "https://api.openai.com/v1/audio/transcriptions", + "prefix_match": False, + "base_model": "gpt-4o-transcribe", + "request_markers": (b"gpt-4o-transcribe", b'name="timestamp_granularities[]"\r\n\r\nword\r\n'), + "marker_in_url": False, + "config_class": litellm.OpenAIGPTAudioTranscriptionConfig, + }, + { + "id": "elevenlabs_scribe", + "provider": "elevenlabs", + "kwargs": {"model": "elevenlabs/scribe_v1", "api_key": "xi-offline"}, + "url": "https://api.elevenlabs.io/v1/speech-to-text", + "prefix_match": False, + "base_model": "scribe_v1", + "request_markers": (b"scribe_v1",), + "marker_in_url": False, + "config_class": ElevenLabsAudioTranscriptionConfig, + }, + { + "id": "deepgram_nova", + "provider": "deepgram", + "kwargs": {"model": "deepgram/nova-2", "api_key": "dg-offline"}, + "url": "https://api.deepgram.com/v1/listen", + "prefix_match": True, + "base_model": "nova-2", + "request_markers": (b"model=nova-2",), + "marker_in_url": True, + "config_class": litellm.DeepgramAudioTranscriptionConfig, + }, + { + "id": "mistral_voxtral", + "provider": "mistral", + "kwargs": {"model": "mistral/voxtral-mini-latest", "api_key": "mistral-offline"}, + "url": "https://api.mistral.ai/v1/audio/transcriptions", + "prefix_match": False, + "base_model": "voxtral-mini-latest", + "request_markers": (b"voxtral-mini-latest",), + "marker_in_url": False, + "config_class": MistralAudioTranscriptionConfig, + }, + { + "id": "ovhcloud_whisper", + "provider": "ovhcloud", + "kwargs": {"model": "ovhcloud/whisper-large-v3-turbo", "api_key": "ovh-offline"}, + "url": "https://oai.endpoints.kepler.ai.cloud.ovh.net/v1/audio/transcriptions", + "prefix_match": False, + "base_model": "whisper-large-v3-turbo", + "request_markers": (b"whisper-large-v3-turbo",), + "marker_in_url": False, + "config_class": OVHCloudAudioTranscriptionConfig, + }, +) + + +def _canned_response(case: _Case) -> httpx.Response: + if case["provider"] == "deepgram": + return httpx.Response( + 200, + json={ + "metadata": {"transaction_key": "offline", "duration": 1.5}, + "results": { + "channels": [ + {"alternatives": [{"transcript": "four score and seven years ago", "confidence": 0.99}]} + ] + }, + }, + ) + return httpx.Response(200, json={"text": "four score and seven years ago"}) + + +@pytest.fixture(autouse=True) +def _httpx_only_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + + +def _register(case: _Case, respx_mock: MockRouter) -> respx.Route: + if case["prefix_match"]: + return respx_mock.post(url__startswith=case["url"]).mock(return_value=_canned_response(case)) + return respx_mock.post(case["url"]).mock(return_value=_canned_response(case)) + + +def _assert_translated_request(case: _Case, request: httpx.Request) -> None: + searched: Final = request.url.query if case["marker_in_url"] else request.content + for marker in case["request_markers"]: + assert marker in searched + assert _AUDIO_BYTES in request.content + + +@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"]) +def test_audio_transcription(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock) + transcript: Final = transcription(**dict(case["kwargs"]), file=io.BytesIO(_AUDIO_BYTES)) + _assert_translated_request(case, route.calls.last.request) + assert transcript.text == "four score and seven years ago" + + +@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"]) +@pytest.mark.asyncio +async def test_audio_transcription_async(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock) + transcript: Final = await litellm.atranscription(**dict(case["kwargs"]), file=io.BytesIO(_AUDIO_BYTES)) + _assert_translated_request(case, route.calls.last.request) + assert transcript.text == "four score and seven years ago" + + +@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"]) +def test_audio_transcription_optional_params(case: _Case) -> None: + optional_params: Final = get_supported_openai_params( + model=case["kwargs"]["model"], + custom_llm_provider=case["provider"], + request_type="transcription", + ) + assert isinstance(optional_params, list) + assert optional_params == case["config_class"]().get_supported_openai_params(case["base_model"]) + assert "max_completion_tokens" not in optional_params + + +@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"]) +def test_audio_transcription_config(case: _Case) -> None: + config: Final = ProviderConfigManager.get_provider_audio_transcription_config( + model=case["kwargs"]["model"], + provider=litellm.LlmProviders(case["provider"]), + ) + assert type(config) is case["config_class"] + assert isinstance(config, BaseAudioTranscriptionConfig) + if case["provider"] == "deepgram": + complete_url: Final = config.get_complete_url( + api_base=None, + api_key=None, + model=case["base_model"], + optional_params={}, + litellm_params={}, + ) + assert "api.deepgram.com" in complete_url + assert "model=nova-2" in complete_url + else: + transformed: Final[AudioTranscriptionRequestData] = config.transform_audio_transcription_request( + model=case["base_model"], + audio_file=io.BytesIO(_AUDIO_BYTES), + optional_params={}, + litellm_params={}, + ) + assert _AUDIO_BYTES in _transformed_payload(transformed) + + +def _transformed_payload(transformed: AudioTranscriptionRequestData) -> bytes: + data: Final = transformed.data + if isinstance(data, bytes): + return data + file_entry: Final = data.get("file") if isinstance(data, dict) else None + if isinstance(file_entry, io.BytesIO): + return file_entry.getvalue() + if transformed.files is not None: + first: Final = next(iter(transformed.files.values())) + blob: Final = first[1] if isinstance(first, tuple) else first + return blob.getvalue() if isinstance(blob, io.BytesIO) else cast(bytes, blob) + return b"" diff --git a/tests/unit/llms/base_llm/chat/test_provider_chat_thinking.py b/tests/unit/llms/base_llm/chat/test_provider_chat_thinking.py new file mode 100644 index 00000000000..8f139ca0acc --- /dev/null +++ b/tests/unit/llms/base_llm/chat/test_provider_chat_thinking.py @@ -0,0 +1,385 @@ +import json +from itertools import chain +from typing import Final, Mapping, cast + +import httpx +import pytest +from pydantic import BaseModel, ConfigDict, JsonValue +from respx import MockRouter + +import litellm +from litellm import get_llm_provider +from litellm.constants import ( + DEFAULT_MAX_TOKENS, + DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, + DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET, + DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, +) +from litellm.main import stream_chunk_builder +from litellm.utils import get_optional_params + +from tests.unit.llms.base_llm.chat.test_provider_chat_translation import ( + _BY_ID, + _Case, + _aws_frame, + _call as _provider_call, + _request_body, +) + +_THINKING_BUDGET: Final = 16000 +_THINKING: Final[JsonValue] = {"type": "enabled", "budget_tokens": _THINKING_BUDGET} + +_ANTHROPIC: Final = _BY_ID["anthropic_sonnet45"] +_BEDROCK_HAIKU: Final = _BY_ID["bedrock_converse_haiku"] +_BEDROCK_SONNET: Final = _BY_ID["bedrock_converse_anthropic_thinking"] + +_THINKING_CASES: Final = (_ANTHROPIC, _BEDROCK_SONNET) +_RESPONSE_FORMAT_CASES: Final = (_ANTHROPIC, _BEDROCK_HAIKU) + +_JSON_PREFIX: Final = '{"agent_doing": "researching ' +_JSON_SUFFIX: Final = 'home automation"}' +_JSON_CONTENT: Final = _JSON_PREFIX + _JSON_SUFFIX +_REASONING: Final = "reasoning here" +_SIGNATURE: Final = "sig-1" + + +class _RFormat(BaseModel): + model_config = ConfigDict(frozen=True) + question: str + answer: str + + +_JSON_SCHEMA_ARGS: Final[Mapping[str, JsonValue]] = { + "messages": [ + {"role": "system", "content": "Summarize the agent's thinking into short descriptions."}, + {"role": "user", "content": "Here is the input data."}, + ], + "response_format": { + "type": "json_schema", + "json_schema": { + "name": "final_output", + "strict": True, + "schema": { + "properties": {"agent_doing": {"title": "Agent Doing", "type": "string"}}, + "required": ["agent_doing"], + "title": "ThinkingStep", + "type": "object", + "additionalProperties": False, + }, + }, + }, +} + +_THINKING_MESSAGES: Final[Mapping[str, JsonValue]] = { + "messages": [{"role": "user", "content": "Generate 5 question + answer pairs"}], +} + + +def _case_id(case: _Case) -> str: + return case["id"] + + +def _call(case: _Case, **extra: JsonValue) -> object: + return _provider_call(case, extra) + + +def _anthropic_sse(events: tuple[Mapping[str, JsonValue], ...]) -> str: + return "".join(f"event: {e['type']}\ndata: {json.dumps(e)}\n\n" for e in events) + + +_ANTHROPIC_START: Final[Mapping[str, JsonValue]] = { + "type": "message_start", + "message": { + "id": "msg_offline", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20250929", + "content": [], + "stop_reason": None, + "usage": {"input_tokens": 10, "output_tokens": 1}, + }, +} +_ANTHROPIC_END: Final[tuple[Mapping[str, JsonValue], ...]] = ( + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 5}}, + {"type": "message_stop"}, +) + + +def _anthropic_json_stream() -> str: + return _anthropic_sse( + ( + _ANTHROPIC_START, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _JSON_PREFIX}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _JSON_SUFFIX}}, + {"type": "content_block_stop", "index": 0}, + *_ANTHROPIC_END, + ) + ) + + +def _anthropic_thinking_stream() -> str: + return _anthropic_sse( + ( + _ANTHROPIC_START, + {"type": "content_block_start", "index": 0, "content_block": {"type": "thinking", "thinking": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "thinking_delta", "thinking": _REASONING}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "signature_delta", "signature": _SIGNATURE}}, + {"type": "content_block_stop", "index": 0}, + {"type": "content_block_start", "index": 1, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 1, "delta": {"type": "text_delta", "text": "done"}}, + {"type": "content_block_stop", "index": 1}, + *_ANTHROPIC_END, + ) + ) + + +_CONVERSE_USAGE: Final[Mapping[str, JsonValue]] = {"usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15}} + + +def _converse_stream(frames: tuple[tuple[str, Mapping[str, JsonValue]], ...]) -> bytes: + return b"".join(_aws_frame(event_type, payload) for event_type, payload in frames) + + +def _converse_json_stream() -> bytes: + return _converse_stream( + ( + ("messageStart", {"role": "assistant"}), + ("contentBlockDelta", {"delta": {"text": _JSON_PREFIX}, "contentBlockIndex": 0}), + ("contentBlockDelta", {"delta": {"text": _JSON_SUFFIX}, "contentBlockIndex": 0}), + ("contentBlockStop", {"contentBlockIndex": 0}), + ("messageStop", {"stopReason": "end_turn"}), + ("metadata", _CONVERSE_USAGE), + ) + ) + + +def _converse_thinking_stream() -> bytes: + return _converse_stream( + ( + ("messageStart", {"role": "assistant"}), + ("contentBlockDelta", {"delta": {"reasoningContent": {"text": _REASONING}}, "contentBlockIndex": 0}), + ("contentBlockDelta", {"delta": {"reasoningContent": {"signature": _SIGNATURE}}, "contentBlockIndex": 0}), + ("contentBlockStop", {"contentBlockIndex": 0}), + ("contentBlockDelta", {"delta": {"text": "done"}, "contentBlockIndex": 1}), + ("contentBlockStop", {"contentBlockIndex": 1}), + ("messageStop", {"stopReason": "end_turn"}), + ("metadata", _CONVERSE_USAGE), + ) + ) + + +def _non_stream_response(case: _Case) -> httpx.Response: + if case["shape"] == "anthropic": + return httpx.Response( + 200, + json={ + "id": "msg_offline", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20250929", + "content": [ + {"type": "thinking", "thinking": _REASONING, "signature": _SIGNATURE}, + {"type": "text", "text": _JSON_CONTENT}, + ], + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 5}, + }, + ) + return httpx.Response( + 200, + json={ + "output": { + "message": { + "role": "assistant", + "content": [ + {"reasoningContent": {"reasoningText": {"text": _REASONING, "signature": _SIGNATURE}}}, + {"text": _JSON_CONTENT}, + ], + } + }, + "stopReason": "end_turn", + "usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15}, + }, + ) + + +def _stream_response(case: _Case, *, thinking: bool) -> httpx.Response: + if case["shape"] == "anthropic": + return httpx.Response( + 200, + content=_anthropic_thinking_stream() if thinking else _anthropic_json_stream(), + headers={"content-type": "text/event-stream"}, + ) + return httpx.Response( + 200, + content=_converse_thinking_stream() if thinking else _converse_json_stream(), + headers={"content-type": "application/vnd.amazon.eventstream"}, + ) + + +def _thinking_param(case: _Case, body: Mapping[str, JsonValue]) -> JsonValue: + if case["shape"] == "anthropic": + return body["thinking"] + return cast(Mapping[str, JsonValue], body["additionalModelRequestFields"])["thinking"] + + +def _max_tokens(case: _Case, body: Mapping[str, JsonValue]) -> JsonValue: + if case["shape"] == "anthropic": + return body["max_tokens"] + return cast(Mapping[str, JsonValue], body["inferenceConfig"])["maxTokens"] + + +def _json_schema_title(case: _Case, body: Mapping[str, JsonValue]) -> str: + if case["shape"] == "anthropic": + output_format: Final = cast(Mapping[str, JsonValue], body["output_format"]) + assert output_format["type"] == "json_schema" + return cast(str, cast(Mapping[str, JsonValue], output_format["schema"])["title"]) + text_format: Final = cast( + Mapping[str, JsonValue], + cast(Mapping[str, JsonValue], body["outputConfig"])["textFormat"], + ) + assert text_format["type"] == "json_schema" + json_schema: Final = cast( + Mapping[str, JsonValue], cast(Mapping[str, JsonValue], text_format["structure"])["jsonSchema"] + ) + return cast(str, json.loads(cast(str, json_schema["schema"]))["title"]) + + +def _has_forced_tool_choice(case: _Case, body: Mapping[str, JsonValue]) -> bool: + if case["shape"] == "anthropic": + return body.get("tool_choice") is not None + return "toolConfig" in body + + +@pytest.mark.parametrize("case", _RESPONSE_FORMAT_CASES, ids=_case_id) +def test_anthropic_response_format_streaming_vs_non_streaming(case: _Case, respx_mock: MockRouter) -> None: + stream_route: Final = respx_mock.post(case["stream_url"]).mock(return_value=_stream_response(case, thinking=False)) + chunks: Final = tuple(cast(litellm.CustomStreamWrapper, _call(case, **_JSON_SCHEMA_ARGS, stream=True))) + built: Final = stream_chunk_builder(chunks=list(chunks)) + stream_body: Final = _request_body(stream_route) + + non_stream_route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case)) + non_stream: Final = cast(litellm.ModelResponse, _call(case, **_JSON_SCHEMA_ARGS)) + non_stream_body: Final = _request_body(non_stream_route) + + assert len(chunks) > 1 + assert _json_schema_title(case, stream_body) == "ThinkingStep" + assert _json_schema_title(case, non_stream_body) == "ThinkingStep" + assert built is not None + streamed_json: Final = cast( + Mapping[str, JsonValue], + json.loads(cast(str, cast(litellm.ModelResponse, built).choices[0].message.content)), + ) + non_stream_json: Final = cast(Mapping[str, JsonValue], json.loads(cast(str, non_stream.choices[0].message.content))) + assert streamed_json == non_stream_json == {"agent_doing": "researching home automation"} + + +@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id) +def test_completion_thinking_with_response_format(case: _Case, respx_mock: MockRouter) -> None: + route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case)) + response: Final = cast( + litellm.ModelResponse, + _call(case, thinking=_THINKING, **_THINKING_MESSAGES, response_format=cast(JsonValue, _RFormat)), + ) + body: Final = _request_body(route) + assert _thinking_param(case, body) == _THINKING + assert _json_schema_title(case, body) == "_RFormat" + assert not _has_forced_tool_choice(case, body) + assert response.choices[0].message.content == _JSON_CONTENT + assert response.choices[0].message.reasoning_content == _REASONING + + +@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id) +def test_completion_thinking_with_max_tokens(case: _Case, respx_mock: MockRouter) -> None: + route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case)) + response: Final = cast( + litellm.ModelResponse, + _call(case, thinking=_THINKING, **_THINKING_MESSAGES, max_completion_tokens=20000), + ) + body: Final = _request_body(route) + assert _max_tokens(case, body) == 20000 + assert _thinking_param(case, body) == _THINKING + assert response.choices[0].message.content == _JSON_CONTENT + + +@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id) +def test_completion_thinking_without_max_tokens(case: _Case, respx_mock: MockRouter) -> None: + route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case)) + response: Final = cast(litellm.ModelResponse, _call(case, thinking=_THINKING, **_THINKING_MESSAGES)) + body: Final = _request_body(route) + max_tokens: Final = cast(int, _max_tokens(case, body)) + assert max_tokens == _THINKING_BUDGET + DEFAULT_MAX_TOKENS + assert max_tokens > _THINKING_BUDGET + assert _thinking_param(case, body) == _THINKING + assert response.choices[0].message.content == _JSON_CONTENT + + +@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id) +def test_anthropic_thinking_output_stream(case: _Case, respx_mock: MockRouter) -> None: + route: Final = respx_mock.post(case["stream_url"]).mock(return_value=_stream_response(case, thinking=True)) + chunks: Final = tuple( + cast( + litellm.CustomStreamWrapper, + _call( + case, + thinking=_THINKING, + messages=[{"role": "user", "content": "Tell me a joke."}], + stream=True, + ), + ) + ) + deltas: Final = tuple(chunk.choices[0].delta for chunk in chunks) + thinking_deltas: Final = tuple( + delta + for delta in deltas + if isinstance(getattr(delta, "thinking_blocks", None), list) + and delta.thinking_blocks + and isinstance(getattr(delta, "reasoning_content", None), str) + ) + blocks: Final = chain.from_iterable(cast(list[object], delta.thinking_blocks) for delta in thinking_deltas) + signatures: Final = tuple(cast(Mapping[str, JsonValue], block).get("signature") for block in blocks) + assert _thinking_param(case, _request_body(route)) == _THINKING + assert not any(delta.tool_calls for delta in deltas) + assert "".join(cast(str, delta.reasoning_content) for delta in thinking_deltas) == _REASONING + assert _SIGNATURE in signatures + + +@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id) +def test_anthropic_reasoning_effort_thinking_translation(case: _Case, respx_mock: MockRouter) -> None: + model: Final = case["kwargs"].get("model", "") + _, provider, _, _ = get_llm_provider(model=model) + optional_params: Final = get_optional_params(model=model, custom_llm_provider=provider, reasoning_effort="high") + assert optional_params["thinking"] == { + "type": "enabled", + "budget_tokens": DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, + } + assert "reasoning_effort" not in optional_params + + route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case)) + _call(case, reasoning_effort="high", messages=[{"role": "user", "content": "hi"}]) + body: Final = _request_body(route) + assert _thinking_param(case, body) == { + "type": "enabled", + "budget_tokens": DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, + } + assert "reasoning_effort" not in json.dumps(body) + assert _max_tokens(case, body) == DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET + DEFAULT_MAX_TOKENS + + +@pytest.mark.parametrize("case", _THINKING_CASES, ids=_case_id) +@pytest.mark.parametrize( + ("effort", "budget"), + ( + ("low", DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET), + ("medium", DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET), + ), +) +def test_reasoning_effort_maps_to_distinct_thinking_budgets( + case: _Case, effort: str, budget: int, respx_mock: MockRouter +) -> None: + route: Final = respx_mock.post(case["url"]).mock(return_value=_non_stream_response(case)) + _call(case, reasoning_effort=effort, messages=[{"role": "user", "content": "hi"}]) + body: Final = _request_body(route) + assert _thinking_param(case, body) == {"type": "enabled", "budget_tokens": budget} + assert _max_tokens(case, body) == budget + DEFAULT_MAX_TOKENS diff --git a/tests/unit/llms/base_llm/chat/test_provider_chat_translation.py b/tests/unit/llms/base_llm/chat/test_provider_chat_translation.py new file mode 100644 index 00000000000..e0e07a3dd59 --- /dev/null +++ b/tests/unit/llms/base_llm/chat/test_provider_chat_translation.py @@ -0,0 +1,1644 @@ +import base64 +import copy +import itertools +import json +import struct +import zlib +from typing import Callable, Final, Iterable, Literal, Mapping, cast + +import httpx +import pytest +import respx +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter +from respx import MockRouter +from typing_extensions import ReadOnly, TypedDict + +import litellm +from litellm.constants import ( + DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, + DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET, +) +from litellm.llms.base_llm.base_utils import type_to_response_format_param +from litellm.types.utils import CallTypes +from litellm.utils import ProviderConfigManager, return_raw_request + +_Shape = Literal[ + "openai", + "anthropic", + "gemini", + "bedrock_converse", + "bedrock_invoke", + "bedrock_invoke_nova", + "bedrock_invoke_openai", +] + +_AWS_KWARGS: Final[Mapping[str, str]] = { + "aws_access_key_id": "AKIAFAKE", + "aws_secret_access_key": "fakesecret", + "aws_region_name": "us-east-1", +} + + +class _Kwargs(TypedDict, total=False): + model: ReadOnly[str] + api_key: ReadOnly[str] + api_base: ReadOnly[str] + api_version: ReadOnly[str] + aws_access_key_id: ReadOnly[str] + aws_secret_access_key: ReadOnly[str] + aws_region_name: ReadOnly[str] + + +class _Case(TypedDict): + id: ReadOnly[str] + shape: ReadOnly[_Shape] + kwargs: ReadOnly[_Kwargs] + url: ReadOnly[str] + stream_url: ReadOnly[str] + router: ReadOnly[bool] + + +def _converse_url(model_id: str, region: str = "us-east-1") -> str: + return f"https://bedrock-runtime.{region}.amazonaws.com/model/{model_id}/converse" + + +def _converse_stream_url(model_id: str, region: str = "us-east-1") -> str: + return f"https://bedrock-runtime.{region}.amazonaws.com/model/{model_id}/converse-stream" + + +def _invoke_url(model_id: str) -> str: + return f"https://bedrock-runtime.us-east-1.amazonaws.com/model/{model_id}/invoke" + + +def _invoke_stream_url(model_id: str) -> str: + return f"https://bedrock-runtime.us-east-1.amazonaws.com/model/{model_id}/invoke-with-response-stream" + + +def _bedrock_case( + case_id: str, + model: str, + invoke_model_id: str, + shape: _Shape, + region: str = "us-east-1", +) -> _Case: + kwargs: _Kwargs = {"model": model, **cast(_Kwargs, dict(_AWS_KWARGS))} + if region != "us-east-1": + kwargs = {"model": model, **cast(_Kwargs, dict(_AWS_KWARGS)), "aws_region_name": region} + is_invoke = "invoke" in model + url: Final = _invoke_url(invoke_model_id) if is_invoke else _converse_url(invoke_model_id, region) + stream_url: Final = ( + _invoke_stream_url(invoke_model_id) if is_invoke else _converse_stream_url(invoke_model_id, region) + ) + return { + "id": case_id, + "shape": shape, + "kwargs": kwargs, + "url": url, + "stream_url": stream_url, + "router": False, + } + + +def _openai_case( + case_id: str, + model: str, + url: str, + *, + router: bool = False, + extra: _Kwargs | None = None, +) -> _Case: + kwargs: _Kwargs = {"model": model, "api_key": "sk-offline"} + if extra is not None: + kwargs = {**kwargs, **extra} + return { + "id": case_id, + "shape": "openai", + "kwargs": kwargs, + "url": url, + "stream_url": url, + "router": router, + } + + +_CASES: Final[tuple[_Case, ...]] = ( + _openai_case("openai_gpt4omini", "gpt-4o-mini", "https://api.openai.com/v1/chat/completions"), + _openai_case("router_gpt4omini", "gpt-4o-mini", "https://api.openai.com/v1/chat/completions", router=True), + _openai_case("openai_o1", "o1", "https://api.openai.com/v1/chat/completions"), + _openai_case("openai_o3mini", "o3-mini", "https://api.openai.com/v1/chat/completions"), + _openai_case( + "azure_o3mini", + "azure/o3-mini", + "https://openai-gpt-4-test-v-1.openai.azure.com/openai/deployments/o3-mini/chat/completions?api-version=2024-02-15-preview", + extra={ + "api_key": "k", + "api_base": "https://openai-gpt-4-test-v-1.openai.azure.com", + "api_version": "2024-02-15-preview", + }, + ), + _openai_case( + "azure_o3mini_live", + "azure/o3-mini", + "https://openai-prod-test.openai.azure.com/openai/deployments/o3-mini/chat/completions?api-version=2024-12-01-preview", + extra={ + "api_key": "k", + "api_base": "https://openai-prod-test.openai.azure.com", + "api_version": "2024-12-01-preview", + }, + ), + { + "id": "anthropic_sonnet45", + "shape": "anthropic", + "kwargs": {"model": "anthropic/claude-sonnet-4-5-20250929", "api_key": "sk-offline"}, + "url": "https://api.anthropic.com/v1/messages", + "stream_url": "https://api.anthropic.com/v1/messages", + "router": False, + }, + { + "id": "gemini_25flash", + "shape": "gemini", + "kwargs": {"model": "gemini/gemini-2.5-flash", "api_key": "k"}, + "url": "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:generateContent", + "stream_url": "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash:streamGenerateContent", + "router": False, + }, + _openai_case("mistral_medium", "mistral/mistral-medium-latest", "https://api.mistral.ai/v1/chat/completions"), + _openai_case("together_glm", "together_ai/zai-org/GLM-5.3-Flash", "https://api.together.ai/v1/chat/completions"), + _openai_case("groq_oss120b", "groq/openai/gpt-oss-120b", "https://api.groq.com/openai/v1/chat/completions"), + _openai_case("xai_grok3mini", "xai/grok-3-mini-beta", "https://api.x.ai/v1/chat/completions"), + _openai_case( + "huggingface_llama", + "huggingface/together/meta-llama/Meta-Llama-3-8B-Instruct", + "https://router.huggingface.co/together/v1/chat/completions", + extra={"api_base": "https://router.huggingface.co/together/v1"}, + ), + _bedrock_case( + "bedrock_converse_haiku", + "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + "us.anthropic.claude-haiku-4-5-20251001-v1:0", + "bedrock_converse", + ), + _bedrock_case( + "bedrock_converse_haiku_xregion", + "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + "us.anthropic.claude-haiku-4-5-20251001-v1:0", + "bedrock_converse", + region="us-west-2", + ), + _bedrock_case( + "bedrock_converse_novalite", "bedrock/us.amazon.nova-lite-v1:0", "us.amazon.nova-lite-v1:0", "bedrock_converse" + ), + _bedrock_case( + "bedrock_converse_novamicro", + "bedrock/converse/us.amazon.nova-micro-v1:0", + "us.amazon.nova-micro-v1:0", + "bedrock_converse", + ), + _bedrock_case( + "bedrock_converse_llama33", + "bedrock/converse/us.meta.llama3-3-70b-instruct-v1:0", + "us.meta.llama3-3-70b-instruct-v1:0", + "bedrock_converse", + ), + _bedrock_case( + "bedrock_converse_gptoss", + "bedrock/converse/openai.gpt-oss-20b-1:0", + "openai.gpt-oss-20b-1:0", + "bedrock_converse", + ), + _bedrock_case( + "bedrock_converse_anthropic_thinking", + "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", + "us.anthropic.claude-sonnet-4-5-20250929-v1:0", + "bedrock_converse", + ), + _bedrock_case( + "bedrock_invoke_haiku", + "bedrock/invoke/us.anthropic.claude-haiku-4-5-20251001-v1:0", + "us.anthropic.claude-haiku-4-5-20251001-v1:0", + "bedrock_invoke", + ), + _bedrock_case( + "bedrock_invoke_novamicro", + "bedrock/invoke/us.amazon.nova-micro-v1:0", + "us.amazon.nova-micro-v1:0", + "bedrock_invoke_nova", + ), + _bedrock_case( + "bedrock_invoke_kimi", + "bedrock/invoke/moonshot.kimi-k2-thinking", + "moonshot.kimi-k2-thinking", + "bedrock_invoke_openai", + ), +) + +_BY_ID: Final[Mapping[str, _Case]] = {c["id"]: c for c in _CASES} + + +def _case_id(case: _Case) -> str: + return case["id"] + + +def _pick(*ids: str) -> tuple[_Case, ...]: + return tuple(_BY_ID[i] for i in ids) + + +_JSON: Final = TypeAdapter(dict[str, JsonValue]) +_ITEMS: Final = TypeAdapter(list[JsonValue]) + + +def _openai_response(text: str) -> httpx.Response: + return httpx.Response( + 200, + json={ + "id": "chatcmpl-offline", + "object": "chat.completion", + "created": 1, + "model": "m", + "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": text}}], + "service_tier": "default", + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + }, + ) + + +def _anthropic_response(text: str) -> httpx.Response: + return httpx.Response( + 200, + json={ + "id": "msg_offline", + "type": "message", + "role": "assistant", + "model": "m", + "content": [{"type": "text", "text": text}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 5}, + }, + ) + + +def _gemini_response(text: str) -> httpx.Response: + return httpx.Response( + 200, + json={ + "candidates": [ + { + "content": {"parts": [{"text": text}], "role": "model"}, + "finishReason": "STOP", + "index": 0, + } + ], + "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 5, "totalTokenCount": 15}, + }, + ) + + +def _converse_response(text: str) -> httpx.Response: + return httpx.Response( + 200, + json={ + "output": {"message": {"role": "assistant", "content": [{"text": text}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15}, + "metrics": {"latencyMs": 1}, + }, + ) + + +def _canned_response(case: _Case, text: str) -> httpx.Response: + match case["shape"]: + case "openai" | "bedrock_invoke_openai": + return _openai_response(text) + case "anthropic" | "bedrock_invoke": + return _anthropic_response(text) + case "gemini": + return _gemini_response(text) + case "bedrock_converse" | "bedrock_invoke_nova": + return _converse_response(text) + + +_PNG_BYTES: Final = base64.b64decode( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg==" +) +_PNG_B64: Final = base64.b64encode(_PNG_BYTES).decode() +_PNG_URLS: Final = ( + "https://cdn.jsdelivr.net/gh/BerriAI/litellm@d769e81c90d453240c61fc572cdb27fae06a89d0/ui/litellm-dashboard/public/assets/logos/litellm_logo.jpg", + "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png", +) + + +def _aws_str_header(name: str, value: str) -> bytes: + name_bytes: Final = name.encode() + value_bytes: Final = value.encode() + return ( + struct.pack("!B", len(name_bytes)) + + name_bytes + + struct.pack("!B", 7) + + struct.pack("!H", len(value_bytes)) + + value_bytes + ) + + +def _aws_frame(event_type: str, payload: Mapping[str, JsonValue]) -> bytes: + payload_bytes: Final = json.dumps(payload).encode() + headers_bytes: Final = b"".join( + ( + _aws_str_header(":event-type", event_type), + _aws_str_header(":content-type", "application/json"), + _aws_str_header(":message-type", "event"), + ) + ) + prelude: Final = struct.pack("!II", 12 + len(headers_bytes) + len(payload_bytes) + 4, len(headers_bytes)) + message: Final = prelude + struct.pack("!I", zlib.crc32(prelude) & 0xFFFFFFFF) + headers_bytes + payload_bytes + return message + struct.pack("!I", zlib.crc32(message) & 0xFFFFFFFF) + + +def _openai_sse(text: str) -> str: + chunks: Final = ( + { + "id": "chatcmpl-offline", + "object": "chat.completion.chunk", + "created": 1, + "model": "m", + "choices": [{"index": 0, "delta": {"role": "assistant", "content": text}, "finish_reason": None}], + "service_tier": "default", + }, + { + "id": "chatcmpl-offline", + "object": "chat.completion.chunk", + "created": 1, + "model": "m", + "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], + "service_tier": "default", + }, + ) + return "".join(f"data: {json.dumps(c)}\n\n" for c in chunks) + "data: [DONE]\n\n" + + +def _anthropic_sse(text: str) -> str: + events: Final = ( + ( + "message_start", + { + "type": "message_start", + "message": { + "id": "msg_offline", + "type": "message", + "role": "assistant", + "content": [], + "model": "m", + "stop_reason": None, + "usage": {"input_tokens": 10, "output_tokens": 1}, + }, + }, + ), + ( + "content_block_start", + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + ), + ( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}}, + ), + ("content_block_stop", {"type": "content_block_stop", "index": 0}), + ( + "message_delta", + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 5}}, + ), + ("message_stop", {"type": "message_stop"}), + ) + return "".join(f"event: {name}\ndata: {json.dumps(payload)}\n\n" for name, payload in events) + + +def _gemini_sse(text: str) -> str: + event: Final = { + "candidates": [{"content": {"parts": [{"text": text}], "role": "model"}, "finishReason": "STOP", "index": 0}], + "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 5, "totalTokenCount": 15}, + } + return f"data: {json.dumps(event)}\n\n" + + +def _converse_stream_bytes(text: str) -> bytes: + frames: Final = ( + _aws_frame("messageStart", {"role": "assistant"}), + _aws_frame("contentBlockDelta", {"delta": {"text": text}, "contentBlockIndex": 0}), + _aws_frame("contentBlockStop", {"contentBlockIndex": 0}), + _aws_frame("messageStop", {"stopReason": "end_turn"}), + _aws_frame("metadata", {"usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15}}), + ) + return b"".join(frames) + + +def _invoke_stream_bytes(text: str) -> bytes: + events: Final = ( + { + "type": "message_start", + "message": { + "id": "msg_invoke", + "type": "message", + "role": "assistant", + "content": [], + "model": "m", + "stop_reason": None, + "usage": {"input_tokens": 10, "output_tokens": 1}, + }, + }, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}}, + {"type": "content_block_stop", "index": 0}, + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 5}}, + {"type": "message_stop"}, + ) + return b"".join(_aws_frame("chunk", {"bytes": base64.b64encode(json.dumps(e).encode()).decode()}) for e in events) + + +def _invoke_nova_stream_bytes(text: str) -> bytes: + events: Final = ( + {"messageStart": {"role": "assistant"}}, + {"contentBlockDelta": {"delta": {"text": text}, "contentBlockIndex": 0}}, + {"contentBlockStop": {"contentBlockIndex": 0}}, + {"messageStop": {"stopReason": "end_turn"}}, + {"metadata": {"usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15}}}, + ) + return b"".join(_aws_frame("chunk", {"bytes": base64.b64encode(json.dumps(e).encode()).decode()}) for e in events) + + +def _stream_canned(case: _Case, text: str) -> httpx.Response: + match case["shape"]: + case "openai" | "bedrock_invoke_openai": + return httpx.Response(200, content=_openai_sse(text), headers={"content-type": "text/event-stream"}) + case "anthropic": + return httpx.Response(200, content=_anthropic_sse(text), headers={"content-type": "text/event-stream"}) + case "gemini": + return httpx.Response(200, content=_gemini_sse(text), headers={"content-type": "text/event-stream"}) + case "bedrock_converse": + return httpx.Response( + 200, + content=_converse_stream_bytes(text), + headers={"content-type": "application/vnd.amazon.eventstream"}, + ) + case "bedrock_invoke": + return httpx.Response( + 200, + content=_invoke_stream_bytes(text), + headers={"content-type": "application/vnd.amazon.eventstream"}, + ) + case "bedrock_invoke_nova": + return httpx.Response( + 200, + content=_invoke_nova_stream_bytes(text), + headers={"content-type": "application/vnd.amazon.eventstream"}, + ) + + +@pytest.fixture(autouse=True) +def _httpx_only_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + + +def _tool_use_response(case: _Case, payload: str) -> httpx.Response: + arguments: Final = json.loads(payload) + match case["shape"]: + case "anthropic" | "bedrock_invoke": + return httpx.Response( + 200, + json={ + "id": "msg_offline", + "type": "message", + "role": "assistant", + "model": "m", + "content": [ + {"type": "tool_use", "id": "toolu_offline", "name": "json_tool_call", "input": arguments} + ], + "stop_reason": "tool_use", + "usage": {"input_tokens": 10, "output_tokens": 5}, + }, + ) + case "bedrock_converse" | "bedrock_invoke_nova": + return httpx.Response( + 200, + json={ + "output": { + "message": { + "role": "assistant", + "content": [ + { + "toolUse": { + "toolUseId": "tooluse_offline", + "name": "json_tool_call", + "input": arguments, + } + } + ], + } + }, + "stopReason": "tool_use", + "usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15}, + "metrics": {"latencyMs": 1}, + }, + ) + case _: + return _canned_response(case, payload) + + +def _json_responder(case: _Case, payload: str) -> Callable[[httpx.Request], httpx.Response]: + def respond(request: httpx.Request) -> httpx.Response: + if b"json_tool_call" in request.content: + return _tool_use_response(case, payload) + return _canned_response(case, payload) + + return respond + + +def _register( + case: _Case, + respx_mock: MockRouter, + text: str, + *, + stream: bool = False, + json_via_tool: bool = False, +) -> respx.Route: + for url in _PNG_URLS: + respx_mock.get(url).mock(return_value=httpx.Response(200, content=_PNG_BYTES)) + target: Final = case["stream_url"] if stream else case["url"] + if json_via_tool: + return respx_mock.post(target).mock(side_effect=_json_responder(case, text)) + if case["shape"] == "gemini" and stream: + return respx_mock.post(url__startswith=target).mock(return_value=_stream_canned(case, text)) + if stream and case["id"].startswith("groq"): + return respx_mock.post(target).mock(return_value=_canned_response(case, text)) + return respx_mock.post(target).mock( + return_value=_stream_canned(case, text) if stream else _canned_response(case, text) + ) + + +_ROUTER_MODEL_LIST: Final = [ + { + "model_name": "offline-router-model", + "litellm_params": {"model": "gpt-4o-mini", "api_key": "sk-offline"}, + } +] + + +def _call(case: _Case, params: Mapping[str, JsonValue]) -> litellm.ModelResponse | litellm.CustomStreamWrapper: + extra: Final = copy.deepcopy(params) + if case["router"]: + router: Final = litellm.Router(model_list=copy.deepcopy(_ROUTER_MODEL_LIST)) + return router.completion(model="offline-router-model", **extra) + return litellm.completion(**case["kwargs"], **extra) + + +def _complete(case: _Case, params: Mapping[str, JsonValue]) -> litellm.ModelResponse: + response: Final = _call(case, params) + assert isinstance(response, litellm.ModelResponse) + return response + + +def _stream(case: _Case, params: Mapping[str, JsonValue]) -> litellm.CustomStreamWrapper: + response: Final = _call(case, {**params, "stream": True}) + assert isinstance(response, litellm.CustomStreamWrapper) + return response + + +async def _acomplete(case: _Case, params: Mapping[str, JsonValue]) -> litellm.ModelResponse: + extra: Final = copy.deepcopy(params) + if case["router"]: + router: Final = litellm.Router(model_list=copy.deepcopy(_ROUTER_MODEL_LIST)) + response: Final = await router.acompletion(model="offline-router-model", **extra) + else: + response = await litellm.acompletion(**case["kwargs"], **extra) + assert isinstance(response, litellm.ModelResponse) + return response + + +def _request_body(route: respx.Route) -> Mapping[str, JsonValue]: + assert route.calls, "provider route was never called" + return _JSON.validate_python(json.loads(route.calls.last.request.content)) + + +def _mapping(value: JsonValue) -> Mapping[str, JsonValue]: + return _JSON.validate_python(value) + + +def _items(value: JsonValue) -> tuple[JsonValue, ...]: + return tuple(_ITEMS.validate_python(value)) + + +def _mappings(value: JsonValue) -> tuple[Mapping[str, JsonValue], ...]: + return tuple(_mapping(item) for item in _items(value)) + + +def _flatten(groups: Iterable[Iterable[str]]) -> tuple[str, ...]: + return tuple(itertools.chain.from_iterable(groups)) + + +def _typed_text_parts(content: JsonValue) -> tuple[str, ...]: + if isinstance(content, str): + return (content,) + return tuple(str(p["text"]) for p in _mappings(content) if p.get("type") == "text") + + +def _keyed_text_parts(content: JsonValue) -> tuple[str, ...]: + return tuple(str(p["text"]) for p in _mappings(content) if "text" in p) + + +def _user_texts(case: _Case, body: Mapping[str, JsonValue]) -> tuple[str, ...]: + match case["shape"]: + case "openai" | "bedrock_invoke_openai": + return _flatten(_typed_text_parts(m["content"]) for m in _mappings(body["messages"]) if m["role"] == "user") + case "anthropic" | "bedrock_invoke": + return _flatten(_typed_text_parts(m["content"]) for m in _mappings(body["messages"]) if m["role"] == "user") + case "bedrock_converse" | "bedrock_invoke_nova": + return _flatten(_keyed_text_parts(m["content"]) for m in _mappings(body["messages"]) if m["role"] == "user") + case "gemini": + return _flatten( + _keyed_text_parts(m["parts"]) for m in _mappings(body["contents"]) if m.get("role") != "model" + ) + + +def _system_texts(case: _Case, body: Mapping[str, JsonValue]) -> tuple[str, ...]: + match case["shape"]: + case "openai" | "bedrock_invoke_openai": + return tuple(str(m["content"]) for m in _mappings(body["messages"]) if m["role"] == "system") + case "anthropic" | "bedrock_invoke" | "bedrock_converse" | "bedrock_invoke_nova": + return tuple(str(s["text"]) for s in _mappings(body.get("system", []))) + case "gemini": + return tuple( + str(p["text"]) for p in _mappings(_mapping(body.get("system_instruction", {})).get("parts", [])) + ) + + +def _message_roles(case: _Case, body: Mapping[str, JsonValue]) -> tuple[str, ...]: + match case["shape"]: + case "gemini": + return tuple(str(m.get("role", "user")) for m in _mappings(body["contents"])) + case _: + return tuple(str(m["role"]) for m in _mappings(body["messages"])) + + +def _tools_payload(case: _Case, body: Mapping[str, JsonValue]) -> JsonValue: + match case["shape"]: + case "openai" | "bedrock_invoke_openai" | "anthropic" | "bedrock_invoke": + return body.get("tools") + case "bedrock_converse" | "bedrock_invoke_nova": + return _mapping(body.get("toolConfig", {})).get("tools") + case "gemini": + declared: Final = _mappings(body.get("tools", [])) + if not declared: + return None + return declared[0].get("function_declarations") + + +def _tool_names(case: _Case, body: Mapping[str, JsonValue]) -> tuple[str, ...]: + entries: Final = _mappings(_tools_payload(case, body) or []) + match case["shape"]: + case "openai" | "bedrock_invoke_openai": + return tuple(str(_mapping(e["function"])["name"]) for e in entries) + case "bedrock_converse" | "bedrock_invoke_nova": + return tuple(str(_mapping(e["toolSpec"])["name"]) for e in entries) + case "anthropic" | "bedrock_invoke" | "gemini": + return tuple(str(e["name"]) for e in entries) + + +def _tool_input_schema(case: _Case, body: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]: + first: Final = _mappings(_tools_payload(case, body) or [])[0] + match case["shape"]: + case "openai" | "bedrock_invoke_openai": + return _mapping(_mapping(first["function"])["parameters"]) + case "anthropic" | "bedrock_invoke": + return _mapping(first["input_schema"]) + case "bedrock_converse" | "bedrock_invoke_nova": + return _mapping(_mapping(_mapping(first["toolSpec"])["inputSchema"])["json"]) + case "gemini": + return _mapping(first["parameters"]) + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "bedrock_converse_haiku", + "bedrock_converse_novalite", + "bedrock_converse_gptoss", + "bedrock_converse_llama33", + "bedrock_invoke_haiku", + "bedrock_invoke_novamicro", + "groq_oss120b", + "huggingface_llama", + "mistral_medium", + "openai_gpt4omini", + "router_gpt4omini", + "together_glm", + "xai_grok3mini", + ), + ids=_case_id, +) +def test_developer_role_translation(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + response: Final = _complete( + case, + { + "messages": [ + {"role": "developer", "content": "Be a good bot!"}, + {"role": "user", "content": [{"type": "text", "text": "Hello, how are you?"}]}, + ] + }, + ) + body: Final = _request_body(route) + assert "developer" not in _message_roles(case, body) + assert _system_texts(case, body) == ("Be a good bot!",) + assert "Hello, how are you?" in _user_texts(case, body) + assert response.choices[0].message.content == f"canned-{case['id']}" + + +@pytest.mark.parametrize( + "case", + _pick( + "bedrock_converse_gptoss", + "bedrock_invoke_novamicro", + "huggingface_llama", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "openai_o3mini", + "router_gpt4omini", + "together_glm", + "xai_grok3mini", + ), + ids=_case_id, +) +def test_content_list_handling(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + response: Final = _complete( + case, + {"messages": [{"role": "user", "content": [{"type": "text", "text": "Hello, how are you?"}]}]}, + ) + body: Final = _request_body(route) + assert _user_texts(case, body) == ("Hello, how are you?",) + first_content: Final = _mappings(body["messages"])[0]["content"] + if case["id"] == "mistral_medium": + assert first_content == "Hello, how are you?" + else: + assert isinstance(first_content, list) + assert response.choices[0].message.content == f"canned-{case['id']}" + + +_TOOL_ARRAY_SCHEMA: Final[Mapping[str, JsonValue]] = { + "type": "function", + "function": { + "name": "shoe_get_id", + "description": "Get information about a show by its ID or name", + "parameters": { + "type": "object", + "properties": {"shoe_id": {"type": ["string", "number"], "description": "The shoe ID or name"}}, + "required": ["shoe_id"], + "additionalProperties": False, + "$schema": "http://json-schema.org/draft-07/schema#", + }, + }, +} + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "azure_o3mini", + "bedrock_converse_haiku", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "bedrock_converse_gptoss", + "bedrock_converse_llama33", + "bedrock_invoke_haiku", + "bedrock_invoke_kimi", + "gemini_25flash", + "groq_oss120b", + "mistral_medium", + "openai_gpt4omini", + "openai_o3mini", + "router_gpt4omini", + ), + ids=_case_id, +) +def test_tool_call_with_property_type_array(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + response: Final = _complete( + case, + { + "messages": [{"role": "user", "content": "Tell me about shoes"}], + "tools": [_TOOL_ARRAY_SCHEMA], + }, + ) + body: Final = _request_body(route) + assert _tool_names(case, body) == ("shoe_get_id",) + schema: Final = _tool_input_schema(case, body) + shoe_id: Final = _mapping(_mapping(schema["properties"])["shoe_id"]) + assert schema["required"] == ["shoe_id"] + if case["shape"] == "gemini": + assert [_mapping(v)["type"] for v in _items(shoe_id["anyOf"])] == ["string", "number"] + else: + assert shoe_id["type"] == ["string", "number"] + assert response.choices[0].message.content == f"canned-{case['id']}" + + +_TOOL_ENUM_SCHEMA: Final[Mapping[str, JsonValue]] = { + "type": "function", + "function": { + "name": "litellm_product_search", + "description": "Search for product information", + "parameters": { + "properties": { + "search_mode": { + "default": "", + "description": "The search strategy to use", + "enum": ["", "product_search", "product_search_with_filters"], + "type": "string", + } + }, + "required": ["search_mode"], + "type": "object", + }, + }, +} + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "azure_o3mini", + "bedrock_converse_haiku", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "bedrock_converse_gptoss", + "bedrock_converse_llama33", + "bedrock_invoke_haiku", + "bedrock_invoke_kimi", + "gemini_25flash", + "mistral_medium", + "openai_gpt4omini", + "openai_o3mini", + "router_gpt4omini", + ), + ids=_case_id, +) +def test_tool_call_with_empty_enum_property(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + response: Final = _complete( + case, + { + "messages": [{"role": "user", "content": "Search for the latest iPhone models"}], + "tools": [_TOOL_ENUM_SCHEMA], + }, + ) + body: Final = _request_body(route) + assert _tool_names(case, body) == ("litellm_product_search",) + schema: Final = _tool_input_schema(case, body) + search_mode: Final = _mapping(_mapping(schema["properties"])["search_mode"]) + enum_values: Final = _items(search_mode["enum"]) + assert len(enum_values) == 3 + assert enum_values[0] == (None if case["shape"] == "gemini" else "") + assert enum_values[1:] == ("product_search", "product_search_with_filters") + assert response.choices[0].message.content == f"canned-{case['id']}" + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "azure_o3mini", + "bedrock_converse_haiku", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "bedrock_converse_novamicro", + "bedrock_converse_gptoss", + "bedrock_converse_llama33", + "bedrock_invoke_haiku", + "gemini_25flash", + "groq_oss120b", + "huggingface_llama", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "openai_o3mini", + "router_gpt4omini", + "together_glm", + "xai_grok3mini", + ), + ids=_case_id, +) +def test_pydantic_model_input(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + messages: Final = [litellm.Message(content="Hello, how are you?", role="user")] + response: Final = _complete(case, {"messages": messages}) + assert "Hello, how are you?" in _user_texts(case, _request_body(route)) + assert response.choices[0].message.content == f"canned-{case['id']}" + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "bedrock_converse_haiku", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "bedrock_invoke_haiku", + "gemini_25flash", + "openai_gpt4omini", + "openai_o1", + "router_gpt4omini", + ), + ids=_case_id, +) +def test_file_data_unit_test(case: _Case, respx_mock: MockRouter) -> None: + pdf_b64: Final = base64.b64encode(b"%PDF-1.4 offline dummy").decode() + file_data_url: Final = f"data:application/pdf;base64,{pdf_b64}" + raw_request: Final = return_raw_request( + endpoint=CallTypes.completion, + kwargs={ + **{k: v for k, v in dict(case["kwargs"]).items() if k != "api_key"}, + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What's this file about?"}, + {"type": "file", "file": {"file_data": file_data_url}}, + ], + } + ], + }, + ) + assert raw_request.get("error") is None + assert pdf_b64 in json.dumps(raw_request.get("raw_request_body")) + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "azure_o3mini", + "bedrock_converse_haiku", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "bedrock_converse_novamicro", + "bedrock_converse_gptoss", + "bedrock_converse_llama33", + "bedrock_invoke_haiku", + "gemini_25flash", + "groq_oss120b", + "huggingface_llama", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "openai_o3mini", + "router_gpt4omini", + "together_glm", + "xai_grok3mini", + ), + ids=_case_id, +) +def test_message_with_name(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + response: Final = _complete(case, {"messages": [{"role": "user", "content": "Hello", "name": "test_name"}]}) + body: Final = _request_body(route) + assert "Hello" in _user_texts(case, body) + if case["shape"] in ("openai", "bedrock_invoke_openai"): + first: Final = cast(Mapping[str, JsonValue], cast(list[JsonValue], body["messages"])[0]) + if case["id"] == "mistral_medium": + assert "name" not in first + else: + assert first.get("name") == "test_name" + else: + assert "test_name" not in json.dumps(body) + assert response.choices[0].message.content == f"canned-{case['id']}" + + +@pytest.mark.parametrize("response_format", ({"type": "json_object"}, {"type": "text"})) +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "azure_o3mini", + "bedrock_converse_haiku", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "bedrock_converse_novamicro", + "bedrock_invoke_haiku", + "bedrock_invoke_kimi", + "gemini_25flash", + "groq_oss120b", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "openai_o3mini", + "router_gpt4omini", + ), + ids=_case_id, +) +def test_json_response_format(case: _Case, response_format: Mapping[str, JsonValue], respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, '{"city":"San Francisco","state":"CA"}') + response: Final = _complete( + case, + { + "messages": [ + {"role": "system", "content": "Your output should be a JSON object with no additional properties."}, + {"role": "user", "content": "Respond with this in json. city=San Francisco, state=CA"}, + ], + "response_format": response_format, + }, + ) + body: Final = _request_body(route) + match case["shape"]: + case "gemini": + mime_types: Final = {"json_object": "application/json", "text": "text/plain"} + config: Final = _mapping(body["generationConfig"]) + assert config["response_mime_type"] == mime_types[str(response_format["type"])] + case "openai" | "bedrock_invoke_openai": + assert body["response_format"] == response_format + case _: + assert "response_format" not in body + assert not _tools_payload(case, body) + assert response.choices[0].message.content == '{"city":"San Francisco","state":"CA"}' + + +_WEATHER_TOOL: Final[Mapping[str, JsonValue]] = { + "type": "function", + "function": { + "name": "get_current_weather", + "description": "Get the current weather in a given location", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string", "description": "The city and state"}, + "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]}, + }, + "required": ["location"], + }, + }, +} + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "azure_o3mini", + "bedrock_converse_haiku", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "bedrock_converse_novamicro", + "bedrock_converse_gptoss", + "bedrock_converse_llama33", + "bedrock_invoke_haiku", + "gemini_25flash", + "groq_oss120b", + "huggingface_llama", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "openai_o3mini", + "router_gpt4omini", + "together_glm", + "xai_grok3mini", + ), + ids=_case_id, +) +def test_response_format_type_text_with_tool_calls_no_tool_choice(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + response: Final = _complete( + case, + { + "messages": [{"role": "user", "content": "What's the weather like in Boston today?"}], + "response_format": {"type": "text"}, + "tools": [_WEATHER_TOOL], + "drop_params": True, + }, + ) + body: Final = _request_body(route) + assert _tool_names(case, body) == ("get_current_weather",) + assert "tool_choice" not in body + assert "toolChoice" not in _mapping(body.get("toolConfig", {})) + assert response.choices[0].message.content == f"canned-{case['id']}" + + +@pytest.mark.parametrize( + "case", + _CASES, + ids=_case_id, +) +def test_response_format_type_text(case: _Case) -> None: + _, provider, _, _ = litellm.get_llm_provider(model=case["kwargs"]["model"]) + provider_config: Final = ProviderConfigManager.get_provider_chat_config( + case["kwargs"]["model"], litellm.LlmProviders(provider) + ) + translated_params: Final = provider_config.map_openai_params( + non_default_params={"response_format": {"type": "text"}}, + optional_params={}, + model=case["kwargs"]["model"], + drop_params=False, + ) + assert "tool_choice" not in translated_params + assert "tools" not in translated_params + + +class _FirstResponse(BaseModel): + model_config = ConfigDict(frozen=True) + first_response: str + + +class _CalendarEvent(BaseModel): + model_config = ConfigDict(frozen=True) + name: str + date: str + participants: list[str] + + +class _EventsList(BaseModel): + model_config = ConfigDict(frozen=True) + events: list[_CalendarEvent] + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "azure_o3mini", + "bedrock_converse_haiku", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "bedrock_converse_novamicro", + "bedrock_converse_gptoss", + "bedrock_invoke_haiku", + "bedrock_invoke_novamicro", + "bedrock_invoke_kimi", + "groq_oss120b", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "openai_o3mini", + "router_gpt4omini", + ), + ids=_case_id, +) +def test_json_response_pydantic_obj(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, '{"first_response":"paris"}', json_via_tool=True) + response: Final = _complete( + case, + { + "messages": [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "What is the capital of France?"}, + ], + "response_format": _FirstResponse, + }, + ) + body: Final = _request_body(route) + serialized: Final = json.dumps(body) + assert "first_response" in serialized + assert json.loads(response.choices[0].message.content) == {"first_response": "paris"} + assert response.choices[0].message.tool_calls is None + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "azure_o3mini", + "bedrock_converse_haiku", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "bedrock_converse_novamicro", + "bedrock_converse_gptoss", + "bedrock_invoke_haiku", + "bedrock_invoke_novamicro", + "bedrock_invoke_kimi", + "bedrock_invoke_novamicro", + "groq_oss120b", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "openai_o3mini", + "router_gpt4omini", + ), + ids=_case_id, +) +def test_json_response_nested_pydantic_obj(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, '{"events":[]}', json_via_tool=True) + response: Final = _complete( + case, + { + "messages": [{"role": "user", "content": "List 5 important events in the XIX century"}], + "response_format": _EventsList, + }, + ) + body: Final = _request_body(route) + serialized: Final = json.dumps(body) + assert "events" in serialized + assert "participants" in serialized + assert json.loads(response.choices[0].message.content) == {"events": []} + assert response.choices[0].message.tool_calls is None + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "azure_o3mini", + "bedrock_converse_haiku", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "bedrock_converse_novamicro", + "bedrock_converse_gptoss", + "bedrock_invoke_haiku", + "bedrock_invoke_novamicro", + "bedrock_invoke_kimi", + "bedrock_invoke_novamicro", + "groq_oss120b", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "openai_o3mini", + "router_gpt4omini", + ), + ids=_case_id, +) +def test_json_response_nested_json_schema(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, '{"events":[]}', json_via_tool=True) + response: Final = _complete( + case, + { + "messages": [{"role": "user", "content": "List 5 important events in the XIX century"}], + "response_format": type_to_response_format_param(_EventsList), + }, + ) + body: Final = _request_body(route) + serialized: Final = json.dumps(body) + assert "events" in serialized + assert "participants" in serialized + assert json.loads(response.choices[0].message.content) == {"events": []} + assert response.choices[0].message.tool_calls is None + + +def test_audio_input_gemini(respx_mock: MockRouter) -> None: + case: Final = _BY_ID["gemini_25flash"] + wav_b64: Final = base64.b64encode(b"RIFFFAKEWAVDATA").decode() + route: Final = _register(case, respx_mock, "canned-gemini") + response: Final = _complete( + case, + { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is in this recording?"}, + {"type": "input_audio", "input_audio": {"data": wav_b64, "format": "wav"}}, + ], + } + ] + }, + ) + body: Final = _request_body(route) + first_content: Final = cast(Mapping[str, JsonValue], cast(list[JsonValue], body["contents"])[0]) + parts: Final = cast(list[JsonValue], first_content["parts"]) + audio_part: Final = cast(Mapping[str, JsonValue], parts[1]) + assert cast(Mapping[str, JsonValue], audio_part["inline_data"])["data"] == wav_b64 + assert response.choices[0].message.content == "canned-gemini" + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "azure_o3mini", + "bedrock_converse_haiku", + "bedrock_converse_novalite", + "bedrock_converse_gptoss", + "bedrock_invoke_haiku", + "bedrock_invoke_novamicro", + "gemini_25flash", + "groq_oss120b", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "openai_o3mini", + "router_gpt4omini", + "together_glm", + ), + ids=_case_id, +) +def test_json_response_format_stream(case: _Case, respx_mock: MockRouter) -> None: + canned: Final = '{"city":"San Francisco"}' + route: Final = _register(case, respx_mock, canned, stream=True) + response: Final = _stream( + case, + { + "messages": [ + {"role": "system", "content": "Your output should be a JSON object with no additional properties."}, + {"role": "user", "content": "Respond with this in json. city=San Francisco, state=CA"}, + ], + "response_format": {"type": "json_object"}, + }, + ) + content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in response) + assert content == canned + body: Final = _request_body(route) + if case["shape"] in ("openai", "anthropic") and not case["id"].startswith("groq"): + assert body["stream"] is True + + +@pytest.mark.parametrize( + "case", + _pick( + "bedrock_invoke_haiku", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "router_gpt4omini", + "together_glm", + ), + ids=_case_id, +) +@pytest.mark.parametrize("detail", (None, "low", "high"), ids=("detail_none", "detail_low", "detail_high")) +@pytest.mark.parametrize("image_url", _PNG_URLS, ids=("litellm_logo", "awsmp_png")) +def test_image_url(case: _Case, detail: str | None, image_url: str, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + image_url_part: Final[Mapping[str, JsonValue]] = ( + {"url": image_url} if detail is None else {"url": image_url, "detail": detail} + ) + response: Final = _complete( + case, + { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What's in this image?"}, + {"type": "image_url", "image_url": image_url_part}, + ], + } + ] + }, + ) + body: Final = _request_body(route) + assert "What's in this image?" in _user_texts(case, body) + content: Final = _mappings(_mappings(body["messages"])[0]["content"]) + image_block: Final = content[1] + if case["shape"] == "bedrock_invoke": + assert image_block["type"] == "image" + source: Final = _mapping(image_block["source"]) + assert source["type"] == "base64" + assert source["data"] == _PNG_B64 + assert source["media_type"] == ("image/jpeg" if image_url.endswith(".jpg") else "image/png") + else: + assert image_block["type"] == "image_url" + assert image_block["image_url"] == image_url_part + assert response.choices[0].message.content == f"canned-{case['id']}" + + +@pytest.mark.parametrize( + "case", + _pick( + "bedrock_converse_haiku", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "gemini_25flash", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "router_gpt4omini", + "together_glm", + ), + ids=_case_id, +) +def test_image_url_string(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + response: Final = _complete( + case, + { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What's in this image?"}, + {"type": "image_url", "image_url": _PNG_URLS[1]}, + ], + } + ] + }, + ) + body: Final = _request_body(route) + assert "What's in this image?" in _user_texts(case, body) + image_block: Final = _items( + _mappings(body["contents"])[0]["parts"] + if case["shape"] == "gemini" + else _mappings(body["messages"])[0]["content"] + )[1] + match case["shape"]: + case "openai": + assert image_block == {"type": "image_url", "image_url": {"url": _PNG_URLS[1]}} + case _: + assert _PNG_B64 in json.dumps(image_block) + assert _PNG_URLS[1] not in json.dumps(image_block) + assert response.choices[0].message.content == f"canned-{case['id']}" + + +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "bedrock_converse_haiku_xregion", + "bedrock_converse_novalite", + "bedrock_converse_gptoss", + "bedrock_invoke_haiku", + "bedrock_invoke_kimi", + "gemini_25flash", + "mistral_medium", + "openai_gpt4omini", + "openai_o3mini", + "router_gpt4omini", + ), + ids=_case_id, +) +def test_empty_tools(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + response: Final = _complete( + case, + { + "messages": [{"role": "user", "content": "Hello, how are you?"}], + "tools": [], + }, + ) + body: Final = _request_body(route) + if case["shape"] in ("bedrock_converse", "bedrock_invoke_nova", "gemini"): + assert "toolConfig" not in body + assert "tools" not in body + else: + assert _tools_payload(case, body) == [] + assert response.choices[0].message.content == f"canned-{case['id']}" + + +def _cost_model_key(case: _Case) -> str | None: + model: Final = cast(str, case["kwargs"]["model"]) + stripped: Final = model.split("/", 1)[-1] + for candidate in (stripped, model): + if candidate in litellm.model_cost: + return candidate + return None + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "case", + _pick( + "anthropic_sonnet45", + "azure_o3mini", + "bedrock_converse_haiku", + "bedrock_converse_novalite", + "bedrock_converse_novamicro", + "bedrock_converse_llama33", + "bedrock_invoke_haiku", + "gemini_25flash", + "groq_oss120b", + "mistral_medium", + "openai_gpt4omini", + "openai_o1", + "openai_o3mini", + "router_gpt4omini", + "together_glm", + "xai_grok3mini", + ), + ids=_case_id, +) +async def test_completion_cost(case: _Case, respx_mock: MockRouter) -> None: + _register(case, respx_mock, f"canned-{case['id']}") + response: Final = await _acomplete( + case, + {"messages": [{"role": "user", "content": "Hello, how are you?"}]}, + ) + usage: Final = response.usage + assert usage.prompt_tokens == 10 + assert usage.completion_tokens == 5 + assert usage.total_tokens == 15 + model_key: Final = _cost_model_key(case) + cost_entry: Final = litellm.model_cost.get(model_key) if model_key is not None else None + actual_cost: Final = response._hidden_params["response_cost"] + if cost_entry is not None and "input_cost_per_token" in cost_entry: + expected: Final = 10 * cost_entry["input_cost_per_token"] + 5 * cost_entry["output_cost_per_token"] + assert actual_cost == pytest.approx(expected) + else: + try: + expected_cost: Final = litellm.completion_cost( + completion_response=response, model=cast(str, case["kwargs"]["model"]) + ) + except litellm.exceptions.ModelNotMappedError: + assert actual_cost is None + else: + assert actual_cost == expected_cost + + +@pytest.mark.parametrize("input_type", ("input_audio", "audio_url")) +def test_supports_audio_input_gemini(input_type: str) -> None: + wav_b64: Final = base64.b64encode(b"RIFFFAKEWAVDATA").decode() + audio_part: Final[Mapping[str, JsonValue]] = ( + {"type": "input_audio", "input_audio": {"data": wav_b64, "format": "wav"}} + if input_type == "input_audio" + else { + "type": "file", + "file": {"file_id": "gs://bucket/file.wav", "filename": "my-sample-audio-file"}, + } + ) + raw_request: Final = return_raw_request( + endpoint=CallTypes.completion, + kwargs={ + "model": "gemini/gemini-2.5-flash", + "modalities": ["text", "audio"], + "audio": {"voice": "alloy", "format": "wav"}, + "drop_params": True, + "messages": [ + { + "role": "user", + "content": [{"type": "text", "text": "What is in this recording?"}, audio_part], + } + ], + }, + ) + assert raw_request.get("error") is None + serialized: Final = json.dumps(raw_request.get("raw_request_body")) + if input_type == "input_audio": + assert wav_b64 in serialized + else: + assert "gs://bucket/file.wav" in serialized + + +def test_reasoning_effort_gemini(respx_mock: MockRouter) -> None: + case: Final = _BY_ID["gemini_25flash"] + route: Final = _register(case, respx_mock, "canned-gemini") + optional_params: Final = litellm.get_optional_params( + model="gemini/gemini-2.5-flash", + custom_llm_provider="gemini", + reasoning_effort="high", + ) + assert optional_params["thinkingConfig"] == { + "thinkingBudget": DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, + "includeThoughts": True, + } + response: Final = _complete( + case, + { + "messages": [{"role": "user", "content": "Hello!"}], + "reasoning_effort": "low", + }, + ) + body: Final = _request_body(route) + config: Final = cast(Mapping[str, JsonValue], body["generationConfig"]) + assert config["thinkingConfig"] == { + "thinkingBudget": DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET, + "includeThoughts": True, + } + assert response.choices[0].message.content == "canned-gemini" + + +@pytest.mark.parametrize("case", _pick("openai_o1", "openai_o3mini", "azure_o3mini", "azure_o3mini_live"), ids=_case_id) +def test_o_series_reasoning_effort_forwarded(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + _complete( + case, + { + "messages": [{"role": "user", "content": "Hello!"}], + "reasoning_effort": "low", + }, + ) + body: Final = _request_body(route) + assert body["reasoning_effort"] == "low" + + +@pytest.mark.parametrize("case", _pick("openai_o1", "openai_o3mini", "azure_o3mini", "azure_o3mini_live"), ids=_case_id) +def test_o_series_developer_role_kept(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + _complete( + case, + { + "messages": [ + {"role": "developer", "content": "Be a good bot!"}, + {"role": "user", "content": "Hello!"}, + ] + }, + ) + body: Final = _request_body(route) + first: Final = cast(Mapping[str, JsonValue], cast(list[JsonValue], body["messages"])[0]) + assert first["role"] == "developer" + assert first["content"] == "Be a good bot!" + + +@pytest.mark.parametrize("case", _pick("openai_o1", "openai_o3mini", "azure_o3mini", "azure_o3mini_live"), ids=_case_id) +def test_o_series_temperature_dropped(case: _Case, respx_mock: MockRouter) -> None: + route: Final = _register(case, respx_mock, f"canned-{case['id']}") + _complete( + case, + { + "messages": [{"role": "user", "content": "Hello, world!"}], + "temperature": 0.0, + "drop_params": True, + }, + ) + body: Final = _request_body(route) + assert "temperature" not in body diff --git a/tests/unit/llms/base_llm/decisions/test_systemone.py b/tests/unit/llms/base_llm/decisions/test_systemone.py deleted file mode 100644 index 42a1c41b3d3..00000000000 --- a/tests/unit/llms/base_llm/decisions/test_systemone.py +++ /dev/null @@ -1,262 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping, Sequence -from typing import Final - -import pytest -from pydantic import TypeAdapter - -from litellm.llms.base_llm.chat.transformation import BaseLLMException -from litellm.llms.base_llm.decisions.systemone import ( - SYSTEM_ONE_RESPONSE_ADAPTER, - question_keys, - to_decisions_response, - to_system_one_request, -) -from litellm.types.openai_decisions import ( - ChoiceAnswer, - DecisionsRequest, - DecisionsRequestBody, - PredicateAnswer, - ScoreAnswer, -) - -_INPUT: Final = "The export job hangs at 99% and never finishes" -_QUESTIONS: Final[Sequence[Mapping[str, object]]] = ( - {"type": "predicate", "name": "is_defect", "instructions": "Is this a defect?"}, - { - "type": "choice", - "name": "sentiment", - "instructions": "How does the customer feel?", - "choices": [{"value": "positive"}, {"value": "negative", "description": "unhappy"}], - }, - { - "type": "score", - "name": "severity", - "instructions": "How severe is it?", - "levels": [{"label": "none"}, {"label": "low"}, {"label": "high", "description": "blocks users"}], - }, -) -_SYSTEM_ONE_QUESTIONS: Final[Mapping[str, object]] = { - "is_defect": {"type": "noul", "instructions": "Is this a defect?"}, - "sentiment": { - "type": "choice", - "instructions": "How does the customer feel?", - "criteria": {"positive": None, "negative": "unhappy"}, - }, - "severity": {"type": "score", "instructions": "How severe is it?", "criteria": ["none", "low", "blocks users"]}, -} -_SYSTEM_ONE_RESPONSE: Final[Mapping[str, object]] = { - "model": "jev-1.13", - "answers": { - "is_defect": {"type": "noul", "noul": 0.9}, - "sentiment": { - "type": "choice", - "choice": "positive", - "confidence": 0.8, - "probabilities": {"positive": 0.8, "negative": 0.2}, - }, - "severity": { - "type": "score", - "score": 1, - "confidence": 0.7, - "legend": {"0": "none", "1": "low", "2": "high"}, - "probabilities": {"0": 0.1, "1": 0.8, "2": 0.1}, - }, - }, - "usage": {"input_tokens": 367, "output_tokens": 3}, -} -_EXPECTED_ANSWERS: Final[Sequence[Mapping[str, object]]] = ( - {"type": "predicate", "name": "is_defect", "probability": 0.9}, - { - "type": "choice", - "name": "sentiment", - "choice": "positive", - "probabilities": [{"value": "positive", "probability": 0.8}, {"value": "negative", "probability": 0.2}], - "confidence": 0.8, - }, - { - "type": "score", - "name": "severity", - "score": 1.0, - "probabilities": [ - {"value": 0, "label": "none", "probability": 0.1}, - {"value": 1, "label": "low", "probability": 0.8}, - {"value": 2, "label": "high", "probability": 0.1}, - ], - "confidence": 0.7, - }, -) -_EXPECTED_USAGE: Final[Mapping[str, object]] = { - "input_tokens": 367, - "input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 0}, - "output_tokens": 3, - "output_tokens_details": {"reasoning_tokens": 0}, - "total_tokens": 370, -} -_BODY_ADAPTER: Final[TypeAdapter[DecisionsRequestBody]] = TypeAdapter(DecisionsRequestBody) - - -def _body( - input_value: object = _INPUT, - questions: Sequence[Mapping[str, object]] = _QUESTIONS, -) -> DecisionsRequestBody: - return _BODY_ADAPTER.validate_python({"input": input_value, "questions": questions}) - - -def _request( - input_value: object = _INPUT, - questions: Sequence[Mapping[str, object]] = _QUESTIONS, - model: str = "jev-1.13", -) -> DecisionsRequest: - return DecisionsRequest(model=model, body=_body(input_value, questions)) - - -def _predicate(name: str | None = "is_defect") -> tuple[Mapping[str, object]]: - return ({"type": "predicate", "name": name, "instructions": "Is this a defect?"},) - - -def test_openai_request_becomes_the_system_one_body() -> None: - assert to_system_one_request("jev-1.13", _body(), "typesafe") == { - "model": "jev-1.13", - "state": _INPUT, - "questions": _SYSTEM_ONE_QUESTIONS, - } - - -def test_system_one_answers_become_openai_answers_in_question_order() -> None: - response: Final = to_decisions_response( - SYSTEM_ONE_RESPONSE_ADAPTER.validate_python(_SYSTEM_ONE_RESPONSE), _request(), "typesafe" - ) - - assert response.model_dump(mode="json") == { - "model": "jev-1.13", - "answers": list(_EXPECTED_ANSWERS), - "usage": _EXPECTED_USAGE, - } - assert isinstance(response.answers[0], PredicateAnswer) - assert isinstance(response.answers[1], ChoiceAnswer) - assert isinstance(response.answers[2], ScoreAnswer) - - -def test_user_messages_are_joined_into_one_system_one_state() -> None: - messages: Final = [ - {"role": "user", "content": "first"}, - { - "role": "user", - "content": [{"type": "input_text", "text": "second"}, {"type": "input_text", "text": "third"}], - }, - ] - - body: Final = to_system_one_request("jev-1.13", _body(messages, _predicate()), "typesafe") - - assert body["state"] == "first\nsecond\nthird" - - -@pytest.mark.parametrize( - ("label", "input_value", "questions"), - ( - ( - "input_image", - [{"role": "user", "content": [{"type": "input_image", "image_url": "data:image/png;base64,AA=="}]}], - _predicate(), - ), - ("boolean choice", _INPUT, [{"type": "choice", "instructions": "Refund?", "choices": [{"value": True}]}]), - ("unique name", _INPUT, [*_predicate(), *_predicate()]), - ( - "repeated choice", - _INPUT, - [{"type": "choice", "instructions": "Refund?", "choices": [{"value": "yes"}, {"value": "yes"}]}], - ), - ), -) -def test_what_system_one_cannot_express_is_a_400( - label: str, - input_value: object, - questions: Sequence[Mapping[str, object]], -) -> None: - with pytest.raises(BaseLLMException, match=label) as error: - to_system_one_request("jev-1.13", _body(input_value, questions), "perplexity") - - assert error.value.status_code == 400 - assert "perplexity" in error.value.message - - -def test_unnamed_questions_get_positional_keys_that_never_shadow_a_supplied_name() -> None: - body: Final = _body(questions=[*_predicate(None), *_predicate("q0"), *_predicate(None)]) - - assert question_keys(body.questions, "typesafe") == ("_q0", "q0", "q2") - noul: Final = {"type": "noul", "instructions": "Is this a defect?"} - assert to_system_one_request("jev-1.13", body, "typesafe")["questions"] == {"_q0": noul, "q0": noul, "q2": noul} - - -def test_positional_answers_come_back_in_question_order_without_a_name() -> None: - request: Final = _request(questions=[*_predicate(None), *_predicate("q0")]) - system_one: Final = SYSTEM_ONE_RESPONSE_ADAPTER.validate_python( - {"answers": {"_q0": {"type": "noul", "noul": 0.25}, "q0": {"type": "noul", "noul": 0.75}}} - ) - - response: Final = to_decisions_response(system_one, request, "typesafe") - - assert [answer.model_dump(mode="json") for answer in response.answers] == [ - {"type": "predicate", "name": None, "probability": 0.25}, - {"type": "predicate", "name": "q0", "probability": 0.75}, - ] - assert response.usage.model_dump(mode="json") == { - **_EXPECTED_USAGE, - "input_tokens": 0, - "output_tokens": 0, - "total_tokens": 0, - } - - -def test_a_reply_without_a_model_reports_the_requested_model() -> None: - system_one: Final = SYSTEM_ONE_RESPONSE_ADAPTER.validate_python( - {k: v for k, v in _SYSTEM_ONE_RESPONSE.items() if k != "model"} - ) - - response: Final = to_decisions_response(system_one, _request(model="typesafe/jev-1.13.0"), "typesafe") - - assert response.model == "typesafe/jev-1.13.0" - - -def test_a_choice_the_provider_left_out_of_probabilities_is_reported_at_zero() -> None: - system_one: Final = SYSTEM_ONE_RESPONSE_ADAPTER.validate_python( - { - "answers": { - "sentiment": { - "type": "choice", - "choice": "positive", - "confidence": 1.0, - "probabilities": {"positive": 1.0}, - } - } - } - ) - request: Final = _request(questions=_QUESTIONS[1:2]) - - response: Final = to_decisions_response(system_one, request, "typesafe") - - assert response.answers[0].model_dump(mode="json") == { - "type": "choice", - "name": "sentiment", - "choice": "positive", - "probabilities": [{"value": "positive", "probability": 1.0}, {"value": "negative", "probability": 0.0}], - "confidence": 1.0, - } - - -@pytest.mark.parametrize( - "answers", - ( - {}, - {"is_defect": {"type": "choice", "choice": "yes", "confidence": 1.0, "probabilities": {"yes": 1.0}}}, - ), -) -def test_a_reply_without_a_matching_answer_is_a_server_error(answers: Mapping[str, object]) -> None: - system_one: Final = SYSTEM_ONE_RESPONSE_ADAPTER.validate_python({"answers": answers}) - - with pytest.raises(BaseLLMException, match="no predicate answer for question 'is_defect'") as error: - to_decisions_response(system_one, _request(questions=_predicate()), "typesafe") - - assert error.value.status_code == 500 diff --git a/tests/unit/llms/base_llm/embedding/__init__.py b/tests/unit/llms/base_llm/embedding/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/base_llm/embedding/test_provider_embedding_translation.py b/tests/unit/llms/base_llm/embedding/test_provider_embedding_translation.py new file mode 100644 index 00000000000..24157263daa --- /dev/null +++ b/tests/unit/llms/base_llm/embedding/test_provider_embedding_translation.py @@ -0,0 +1,108 @@ +import base64 +import json +from typing import Final, cast + +import httpx +import pytest +from respx import MockRouter +from typing_extensions import ReadOnly, TypedDict + +from litellm import embedding +from litellm.utils import get_optional_params_embeddings + + +class _Kwargs(TypedDict, total=False): + model: ReadOnly[str] + api_key: ReadOnly[str] + api_base: ReadOnly[str] + api_version: ReadOnly[str] + aws_access_key_id: ReadOnly[str] + aws_secret_access_key: ReadOnly[str] + aws_region_name: ReadOnly[str] + + +class _Case(TypedDict): + id: ReadOnly[str] + provider: ReadOnly[str] + kwargs: ReadOnly[_Kwargs] + url: ReadOnly[str] + + +_AZURE_BASE: Final = "https://offline-embed.openai.azure.com" +_AZURE_URL: Final = f"{_AZURE_BASE}/openai/deployments/text-embedding-ada-002/embeddings?api-version=2024-02-15-preview" +_TITAN_URL: Final = "https://bedrock-runtime.us-west-2.amazonaws.com/model/amazon.titan-embed-image-v1/invoke" + +_CASES: Final[tuple[_Case, ...]] = ( + { + "id": "azure_text_embedding", + "provider": "azure", + "kwargs": { + "model": "azure/text-embedding-ada-002", + "api_key": "azure-offline-key", + "api_base": _AZURE_BASE, + "api_version": "2024-02-15-preview", + }, + "url": _AZURE_URL, + }, + { + "id": "bedrock_titan_image", + "provider": "bedrock", + "kwargs": cast( + _Kwargs, + { + "model": "bedrock/amazon.titan-embed-image-v1", + "aws_access_key_id": "AKIAFAKE", + "aws_secret_access_key": "fakesecret", + "aws_region_name": "us-west-2", + }, + ), + "url": _TITAN_URL, + }, +) + + +_MAX_RETRIES_KWARGS: Final[tuple[_Kwargs, ...]] = ( + *(case["kwargs"] for case in _CASES), + {"model": "volcengine/doubao-embedding-text-240715"}, + {"model": "voyage/voyage-3-lite"}, +) + + +def _canned_response(case: _Case) -> httpx.Response: + vector: Final = [0.11, 0.22, 0.33] + if case["provider"] == "azure": + return httpx.Response( + 200, + json={ + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": vector}], + "model": "text-embedding-ada-002", + "usage": {"prompt_tokens": 2, "total_tokens": 2}, + }, + ) + return httpx.Response(200, json={"embedding": vector, "inputTextTokenCount": 4}) + + +@pytest.fixture(autouse=True) +def _httpx_only_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + + +@pytest.mark.parametrize("kwargs", _MAX_RETRIES_KWARGS, ids=lambda k: k["model"]) +def test_embedding_optional_params_max_retries(kwargs: _Kwargs) -> None: + optional_params: Final = get_optional_params_embeddings(**dict(kwargs), max_retries=20) + assert optional_params["max_retries"] == 20 + + +@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"]) +def test_image_embedding(case: _Case, respx_mock: MockRouter) -> None: + png_b64: Final = base64.b64encode(b"\x89PNG\r\n\x1a\nFAKEPIXELS").decode() + data_url: Final = f"data:image/png;base64,{png_b64}" + route: Final = respx_mock.post(case["url"]).mock(return_value=_canned_response(case)) + response: Final = embedding(**dict(case["kwargs"]), input=[data_url]) + body: Final = json.loads(route.calls.last.request.content) + if case["provider"] == "azure": + assert body["input"] == [data_url] + else: + assert body["inputImage"] == png_b64 + assert response.data[0]["embedding"] == [0.11, 0.22, 0.33] diff --git a/tests/unit/llms/base_llm/rerank/__init__.py b/tests/unit/llms/base_llm/rerank/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/base_llm/rerank/test_provider_rerank_translation.py b/tests/unit/llms/base_llm/rerank/test_provider_rerank_translation.py new file mode 100644 index 00000000000..8b630f6b946 --- /dev/null +++ b/tests/unit/llms/base_llm/rerank/test_provider_rerank_translation.py @@ -0,0 +1,187 @@ +import json +from typing import Final, Mapping, cast + +import httpx +import pytest +import respx +from pydantic import JsonValue +from respx import MockRouter +from typing_extensions import ReadOnly, TypedDict + +import litellm + + +class _Kwargs(TypedDict, total=False): + model: ReadOnly[str] + api_key: ReadOnly[str] + aws_access_key_id: ReadOnly[str] + aws_secret_access_key: ReadOnly[str] + aws_region_name: ReadOnly[str] + + +class _Case(TypedDict): + id: ReadOnly[str] + provider: ReadOnly[str] + kwargs: ReadOnly[_Kwargs] + url: ReadOnly[str] + expected_cost_zero: ReadOnly[bool] + billed_units: ReadOnly[Mapping[str, int]] + response_id: ReadOnly[str | None] + + +_AWS: Final[Mapping[str, str]] = { + "aws_access_key_id": "AKIAFAKE", + "aws_secret_access_key": "fakesecret", + "aws_region_name": "us-west-2", +} + + +def _bedrock_arn(model_id: str) -> str: + return f"bedrock/arn:aws:bedrock:us-west-2::foundation-model/{model_id}" + + +_CASES: Final[tuple[_Case, ...]] = ( + { + "id": "jina_reranker", + "provider": "cohere", + "kwargs": {"model": "jina_ai/jina-reranker-v2-base-multilingual", "api_key": "jina-offline"}, + "url": "https://api.jina.ai/v1/rerank", + "expected_cost_zero": False, + "billed_units": {"total_tokens": 4}, + "response_id": "rerank-offline", + }, + { + "id": "bedrock_amazon_rerank", + "provider": "bedrock", + "kwargs": cast(_Kwargs, {"model": _bedrock_arn("amazon.rerank-v1:0"), **dict(_AWS)}), + "url": "https://bedrock-agent-runtime.us-west-2.amazonaws.com/rerank", + "expected_cost_zero": False, + "billed_units": {"search_units": 1}, + "response_id": "rerank-offline", + }, + { + "id": "bedrock_cohere_rerank", + "provider": "bedrock", + "kwargs": cast(_Kwargs, {"model": _bedrock_arn("cohere.rerank-v3-5:0"), **dict(_AWS)}), + "url": "https://bedrock-agent-runtime.us-west-2.amazonaws.com/rerank", + "expected_cost_zero": False, + "billed_units": {"search_units": 1}, + "response_id": "rerank-offline", + }, + { + "id": "nvidia_nim_rerank", + "provider": "nvidia_nim", + "kwargs": {"model": "nvidia_nim/nvidia/llama-3_2-nv-rerankqa-1b-v2", "api_key": "nvapi-offline"}, + "url": "https://ai.api.nvidia.com/v1/retrieval/nvidia/llama-3_2-nv-rerankqa-1b-v2/reranking", + "expected_cost_zero": True, + "billed_units": {"total_tokens": 4}, + "response_id": None, + }, +) + + +def _canned_response(case: _Case) -> httpx.Response: + if case["provider"] == "cohere": + return httpx.Response( + 200, + json={ + "id": "rerank-offline", + "results": [ + {"index": 0, "relevance_score": 0.95}, + {"index": 1, "relevance_score": 0.4}, + ], + "usage": {"total_tokens": 4}, + }, + ) + if case["provider"] == "bedrock": + return httpx.Response( + 200, + json={ + "id": "rerank-offline", + "results": [ + {"index": 0, "relevanceScore": 0.95}, + {"index": 1, "relevanceScore": 0.4}, + ], + "usage": {"search_units": 1}, + }, + ) + return httpx.Response( + 200, + json={ + "rankings": [ + {"index": 0, "logit": 0.95}, + {"index": 1, "logit": 0.4}, + ], + "usage": {"prompt_tokens": 4, "total_tokens": 4}, + }, + ) + + +@pytest.fixture(autouse=True) +def _httpx_only_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + + +def _request_body(route: respx.Route) -> Mapping[str, JsonValue]: + return json.loads(route.calls.last.request.content) + + +def _assert_translated_request(case: _Case, body: Mapping[str, JsonValue]) -> None: + if case["provider"] == "bedrock": + queries: Final = body["queries"] + assert queries == [{"textQuery": {"text": "hello"}, "type": "TEXT"}] + config: Final = body["rerankingConfiguration"]["bedrockRerankingConfiguration"]["modelConfiguration"] + assert config["modelArn"].endswith((".rerank-v1:0", ".rerank-v3-5:0")) + sources: Final = body["sources"] + assert len(sources) == 2 + elif case["provider"] == "nvidia_nim": + assert body["model"] == "nvidia/llama-3.2-nv-rerankqa-1b-v2" + assert body["query"] == {"text": "hello"} + assert body["passages"] == [{"text": "hello"}, {"text": "world"}] + assert body["top_k"] == 2 + else: + assert body["model"] == "jina-reranker-v2-base-multilingual" + assert body["query"] == "hello" + assert body["documents"] == ["hello", "world"] + assert body["top_n"] == 2 + + +@pytest.mark.parametrize("case", _CASES, ids=lambda c: c["id"]) +@pytest.mark.parametrize("sync_mode", (True, False)) +@pytest.mark.asyncio +async def test_basic_rerank(case: _Case, sync_mode: bool, respx_mock: MockRouter) -> None: + route: Final = respx_mock.post(case["url"]).mock(return_value=_canned_response(case)) + if sync_mode: + response: Final = litellm.rerank( + **dict(case["kwargs"]), + query="hello", + documents=["hello", "world"], + top_n=2, + ) + else: + response: Final = await litellm.arerank( + **dict(case["kwargs"]), + query="hello", + documents=["hello", "world"], + top_n=2, + ) + body: Final = _request_body(route) + _assert_translated_request(case, body) + assert route.call_count == 1 + assert isinstance(response.id, str) + if case["response_id"] is not None: + assert response.id == case["response_id"] + assert response.meta["billed_units"] == case["billed_units"] + assert len(response.results) == 2 + assert response.results[0]["index"] == 0 + assert response.results[0]["relevance_score"] == 0.95 + assert response.results[1]["index"] == 1 + assert response.results[1]["relevance_score"] == 0.4 + if case["provider"] == "nvidia_nim": + assert response.results[0]["document"] == {"text": "hello"} + assert response.results[1]["document"] == {"text": "world"} + cost: Final = response._hidden_params["response_cost"] + if case["expected_cost_zero"]: + assert cost == 0.0 + else: + assert cost > 0 diff --git a/tests/unit/llms/bedrock/chat/chat_completions/test_bedrock_chat_completions_transformation.py b/tests/unit/llms/bedrock/chat/chat_completions/test_bedrock_chat_completions_transformation.py index 16ae1114402..9d4c0fac094 100644 --- a/tests/unit/llms/bedrock/chat/chat_completions/test_bedrock_chat_completions_transformation.py +++ b/tests/unit/llms/bedrock/chat/chat_completions/test_bedrock_chat_completions_transformation.py @@ -911,7 +911,7 @@ def _stream_chunk(delta, finish_reason=None, index=0): } -def test_streaming_handler_splits_reasoning_deltas_per_choice(): +def test_streaming_handler_splits_reasoning_deltas_per_choice(local_cost_map): handler = BedrockRuntimeChatCompletionsStreamingHandler(streaming_response=iter(()), sync_stream=True) first = handler.chunk_parser(_stream_chunk({"role": "assistant", "content": "I think"})) @@ -934,7 +934,7 @@ def _reasoning_of(parsed): return getattr(parsed.choices[0].delta, "reasoning_content", None) -def test_streaming_handler_keeps_split_state_per_choice_index(): +def test_streaming_handler_keeps_split_state_per_choice_index(local_cost_map): handler = BedrockRuntimeChatCompletionsStreamingHandler(streaming_response=iter(()), sync_stream=True) opened = handler.chunk_parser(_stream_chunk({"content": "first"}, index=0)) @@ -949,7 +949,7 @@ def test_streaming_handler_keeps_split_state_per_choice_index(): assert not still_reasoning.choices[0].delta.content -def test_streaming_handler_flushes_held_text_on_an_empty_final_delta(): +def test_streaming_handler_flushes_held_text_on_an_empty_final_delta(local_cost_map): handler = BedrockRuntimeChatCompletionsStreamingHandler(streaming_response=iter(()), sync_stream=True) held = handler.chunk_parser(_stream_chunk({"content": "taggedHi"}, finish_reason="stop") @@ -1467,3 +1467,76 @@ def test_gpt56_json_object_with_response_schema_goes_to_converse_as_a_json_tool( assert body["toolConfig"]["tools"][0]["toolSpec"]["name"] == "json_tool_call" assert body["toolConfig"]["toolChoice"] == {"tool": {"name": "json_tool_call"}} assert "response_format" not in body + + +LITERAL_TAGGED_ANSWER = "not thinking Hello" + + +@pytest.mark.parametrize("model", ["openai.gpt-5.6-sol", "us.xai.grok-4.6"]) +def test_streaming_handler_keeps_a_literal_reasoning_tag_outside_gpt_oss(local_cost_map, model): + handler = BedrockRuntimeChatCompletionsStreamingHandler(streaming_response=iter(()), sync_stream=True) + + opened = handler.chunk_parser({**_stream_chunk({"content": "not thinking"}), "model": model}) + assert opened.choices[0].delta.content == "not thinking" + assert _reasoning_of(opened) is None + + closed = handler.chunk_parser({**_stream_chunk({"content": " Hello"}, finish_reason="stop"), "model": model}) + assert closed.choices[0].delta.content == " Hello" + assert _reasoning_of(closed) is None + + +@pytest.mark.parametrize( + "model, expected_content, expected_reasoning", + [ + ("bedrock/global.openai.gpt-5.6-sol", LITERAL_TAGGED_ANSWER, None), + ("bedrock/chat_completions/us.xai.grok-4.6", LITERAL_TAGGED_ANSWER, None), + ("bedrock/chat_completions/openai.gpt-oss-20b-1:0", "Hello", "not thinking"), + ("bedrock/chat_completions/openai.gpt-oss-safeguard-20b", "Hello", "not thinking"), + ], +) +def test_reasoning_tag_split_applies_to_gpt_oss_answers_only( + local_cost_map, fake_aws_env, model, expected_content, expected_reasoning +): + model_id = model.removeprefix("bedrock/").removeprefix("chat_completions/") + requests, client = _recording_client(json=_chat_completion_json(LITERAL_TAGGED_ANSWER, model_id)) + + response = litellm.completion( + model=model, messages=[{"role": "user", "content": "hello"}], client=client, max_tokens=64 + ) + + assert str(requests[0].url).endswith("/openai/v1/chat/completions") + assert response.choices[0].message.content == expected_content + assert getattr(response.choices[0].message, "reasoning_content", None) == expected_reasoning + + +@pytest.mark.parametrize( + "capability_flags, expected_content, expected_reasoning", + [ + ({"supports_bedrock_runtime_chat_completions_inline_reasoning": True}, "Hello", "not thinking"), + ({}, LITERAL_TAGGED_ANSWER, None), + ], +) +def test_reasoning_tag_split_is_read_from_the_cost_map( + monkeypatch, fake_aws_env, capability_flags, expected_content, expected_reasoning +): + model_id = SYNTHETIC_NATIVE_MODEL.removeprefix("chat_completions/") + monkeypatch.setattr(litellm, "model_cost", {model_id: {"litellm_provider": "bedrock_converse", **capability_flags}}) + requests, client = _recording_client(json=_chat_completion_json(LITERAL_TAGGED_ANSWER, model_id)) + + response = litellm.completion( + model=f"bedrock/{SYNTHETIC_NATIVE_MODEL}", + messages=[{"role": "user", "content": "hello"}], + client=client, + max_tokens=64, + ) + + assert str(requests[0].url).endswith("/openai/v1/chat/completions") + assert response.choices[0].message.content == expected_content + assert getattr(response.choices[0].message, "reasoning_content", None) == expected_reasoning + + handler = BedrockRuntimeChatCompletionsStreamingHandler(streaming_response=iter(()), sync_stream=True) + chunk = handler.chunk_parser( + {**_stream_chunk({"content": LITERAL_TAGGED_ANSWER}, finish_reason="stop"), "model": model_id} + ) + assert chunk.choices[0].delta.content == expected_content + assert _reasoning_of(chunk) == expected_reasoning diff --git a/tests/unit/llms/bedrock/chat/test_converse_transformation.py b/tests/unit/llms/bedrock/chat/test_converse_transformation.py index f43f15b8230..e42ee95f8d5 100644 --- a/tests/unit/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/unit/llms/bedrock/chat/test_converse_transformation.py @@ -742,7 +742,7 @@ def _tool_schema_properties(model, tool, litellm_params=None): "us.xai.grok-4.6", "us-gov.xai.grok-4.6", "global.xai.grok-4.7", - "xai.grok-4.7", + "us.xai.grok-4.7", ], ) def test_transform_request_drops_lookaround_regex_for_models_the_cost_map_flags(tool, model): diff --git a/tests/unit/llms/bedrock/chat/test_invoke_handler.py b/tests/unit/llms/bedrock/chat/test_invoke_handler.py index ba35b93270a..d5fe5d8cf65 100644 --- a/tests/unit/llms/bedrock/chat/test_invoke_handler.py +++ b/tests/unit/llms/bedrock/chat/test_invoke_handler.py @@ -1426,3 +1426,19 @@ def test_filter_headers_for_aws_signature(): non_aws_headers = {"x-custom-trace": "trace-123", "x-user-context": "premium", "x-request-source": "mobile-app"} filtered_non_aws = aws_llm._filter_headers_for_aws_signature(non_aws_headers) assert filtered_non_aws == {} + + +@pytest.mark.parametrize("missing", ["botocore", "unrelated_dependency"]) +def test_invoke_decoder_reports_only_missing_aws_dependency(missing): + from unittest.mock import patch + from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder + + failure = ModuleNotFoundError("missing dependency", name=missing) + with patch("builtins.__import__", side_effect=failure): + with pytest.raises(ImportError) as error: + AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0") + if missing == "botocore": + assert "pip install boto3" in str(error.value) + assert error.value.__cause__ is failure + else: + assert error.value is failure diff --git a/tests/unit/llms/bedrock/count_tokens/test_bedrock_mantle_count_tokens_handler.py b/tests/unit/llms/bedrock/count_tokens/test_bedrock_mantle_count_tokens_handler.py new file mode 100644 index 00000000000..71b451f834c --- /dev/null +++ b/tests/unit/llms/bedrock/count_tokens/test_bedrock_mantle_count_tokens_handler.py @@ -0,0 +1,109 @@ +import json +from collections.abc import Mapping +from typing import Final + +import httpx +import pytest + +from litellm.llms.bedrock.common_utils import BedrockError +from litellm.llms.bedrock.count_tokens.mantle_handler import BedrockMantleCountTokensHandler +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + +MANTLE_COUNT_URL: Final = "https://bedrock-mantle.us-east-1.api.aws/anthropic/v1/messages/count_tokens" +LITELLM_PARAMS: Final = { + "aws_access_key_id": "AKIATESTACCESSKEY", + "aws_secret_access_key": "test-secret", + "aws_region_name": "us-east-1", +} +REQUEST: Final[dict[str, object]] = { + "model": "global.anthropic.claude-opus-4-8", + "messages": [{"role": "user", "content": "The quick brown fox"}], + "system": "You are a terse assistant.", + "tools": [{"name": "get_weather", "input_schema": {"type": "object", "properties": {}}}], +} + + +class _MantleEndpoint: + def __init__(self, status_code: int, body: Mapping[str, object]) -> None: + self.status_code: Final = status_code + self.body: Final = body + self.requests: tuple[httpx.Request, ...] = () + + def __call__(self, request: httpx.Request) -> httpx.Response: + self.requests = (*self.requests, request) + return httpx.Response(self.status_code, json=dict(self.body), request=request) + + def only_request(self) -> httpx.Request: + assert len(self.requests) == 1, self.requests + return self.requests[0] + + +def _client(endpoint: _MantleEndpoint) -> AsyncHTTPHandler: + return AsyncHTTPHandler(transport=httpx.MockTransport(endpoint)) + + +@pytest.fixture(autouse=True) +def _sigv4_only(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) + + +@pytest.mark.asyncio +async def test_counts_on_mantle_with_the_base_model_and_the_deployment_credentials() -> None: + mantle: Final = _MantleEndpoint(200, {"input_tokens": 2177}) + + result: Final = await BedrockMantleCountTokensHandler().handle_count_tokens_request( + request_data=dict(REQUEST), + litellm_params=dict(LITELLM_PARAMS), + resolved_model="anthropic.claude-opus-4-8", + client=_client(mantle), + ) + + assert result == {"input_tokens": 2177} + posted: Final = mantle.only_request() + assert str(posted.url) == MANTLE_COUNT_URL + assert posted.headers["Authorization"].startswith("AWS4-HMAC-SHA256") + assert "us-east-1/bedrock/aws4_request" in posted.headers["Authorization"] + assert posted.headers["anthropic-version"] == "2023-06-01" + assert json.loads(posted.content) == { + "model": "anthropic.claude-opus-4-8", + "messages": REQUEST["messages"], + "system": REQUEST["system"], + "tools": REQUEST["tools"], + } + + +@pytest.mark.asyncio +async def test_bedrock_mantle_api_base_env_names_the_host(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("BEDROCK_MANTLE_API_BASE", "https://vpce-abc.bedrock-mantle.us-east-1.vpce.api.aws") + mantle: Final = _MantleEndpoint(200, {"input_tokens": 3}) + + await BedrockMantleCountTokensHandler().handle_count_tokens_request( + request_data=dict(REQUEST), + litellm_params=dict(LITELLM_PARAMS), + resolved_model="anthropic.claude-opus-4-8", + client=_client(mantle), + ) + + assert ( + str(mantle.only_request().url) + == "https://vpce-abc.bedrock-mantle.us-east-1.vpce.api.aws/anthropic/v1/messages/count_tokens" + ) + + +@pytest.mark.asyncio +async def test_non_200_answers_raise_bedrock_error_with_mantle_status() -> None: + mantle: Final = _MantleEndpoint( + 404, {"type": "error", "error": {"type": "not_found_error", "message": "does not exist"}} + ) + + with pytest.raises(BedrockError) as raised: + await BedrockMantleCountTokensHandler().handle_count_tokens_request( + request_data=dict(REQUEST), + litellm_params=dict(LITELLM_PARAMS), + resolved_model="anthropic.claude-sonnet-5-5", + client=_client(mantle), + ) + + assert raised.value.status_code == 404 + assert "does not exist" in raised.value.message diff --git a/tests/unit/llms/bedrock/count_tokens/test_bedrock_token_counter.py b/tests/unit/llms/bedrock/count_tokens/test_bedrock_token_counter.py new file mode 100644 index 00000000000..a3dd8021f73 --- /dev/null +++ b/tests/unit/llms/bedrock/count_tokens/test_bedrock_token_counter.py @@ -0,0 +1,140 @@ +from collections.abc import Mapping +from typing import Final + +import httpx +import pytest +from pydantic import JsonValue, TypeAdapter + +from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + +RUNTIME_HOST: Final = "bedrock-runtime.us-east-1.amazonaws.com" +MANTLE_HOST: Final = "bedrock-mantle.us-east-1.api.aws" +UNSUPPORTED: Final = {"message": "The provided model doesn't support counting tokens."} +DEPLOYMENT: Final = { + "litellm_params": { + "aws_access_key_id": "AKIATESTACCESSKEY", + "aws_secret_access_key": "test-secret", + "aws_region_name": "us-east-1", + } +} +MESSAGES: Final = [{"role": "user", "content": "The quick brown fox jumps over the lazy dog."}] +_JSON_BODY: Final = TypeAdapter(dict[str, JsonValue]) + + +class _Bedrock: + def __init__(self, runtime: tuple[int, Mapping[str, object]], mantle: tuple[int, Mapping[str, object]]) -> None: + self.runtime: Final = runtime + self.mantle: Final = mantle + self.requests: tuple[httpx.Request, ...] = () + + def __call__(self, request: httpx.Request) -> httpx.Response: + self.requests = (*self.requests, request) + status_code, body = self.runtime if request.url.host == RUNTIME_HOST else self.mantle + return httpx.Response(status_code, json=dict(body), request=request) + + def posted_hosts(self) -> tuple[str, ...]: + return tuple(request.url.host for request in self.requests) + + +def _counter(bedrock: _Bedrock) -> BedrockTokenCounter: + return BedrockTokenCounter(client=AsyncHTTPHandler(transport=httpx.MockTransport(bedrock))) + + +@pytest.fixture(autouse=True) +def _sigv4_only(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) + + +@pytest.mark.asyncio +async def test_claude_model_bedrock_runtime_cannot_count_is_counted_on_mantle() -> None: + bedrock: Final = _Bedrock(runtime=(400, UNSUPPORTED), mantle=(200, {"input_tokens": 2177})) + + result: Final = await _counter(bedrock).count_tokens( + model_to_use="global.anthropic.claude-opus-4-8", + messages=MESSAGES, + contents=None, + deployment=DEPLOYMENT, + request_model="claude-opus-4-8", + system="You are a terse assistant.", + ) + + assert result is not None + assert result.error is False + assert result.total_tokens == 2177 + assert result.tokenizer_type == "bedrock_mantle_api" + assert result.original_response == {"input_tokens": 2177} + assert bedrock.posted_hosts() == (RUNTIME_HOST, MANTLE_HOST) + mantle_body: Final = _JSON_BODY.validate_json(bedrock.requests[1].content) + assert mantle_body["model"] == "anthropic.claude-opus-4-8" + assert mantle_body["system"] == "You are a terse assistant." + + +@pytest.mark.asyncio +async def test_bedrock_runtime_count_is_kept_when_it_answers() -> None: + bedrock: Final = _Bedrock(runtime=(200, {"inputTokens": 1353}), mantle=(200, {"input_tokens": 1336})) + + result: Final = await _counter(bedrock).count_tokens( + model_to_use="global.anthropic.claude-sonnet-4-6", + messages=MESSAGES, + contents=None, + deployment=DEPLOYMENT, + request_model="claude-sonnet-4-6", + ) + + assert result is not None + assert result.total_tokens == 1353 + assert result.tokenizer_type == "bedrock_api" + assert bedrock.posted_hosts() == (RUNTIME_HOST,) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model_to_use", "runtime"), + ( + ("global.anthropic.claude-opus-4-8", (403, {"Message": "not authorized to perform: bedrock:CountTokens"})), + ("amazon.nova-pro-v1:0", (400, UNSUPPORTED)), + ), +) +async def test_other_bedrock_runtime_failures_are_not_retried_on_mantle( + model_to_use: str, runtime: tuple[int, Mapping[str, object]] +) -> None: + bedrock: Final = _Bedrock(runtime=runtime, mantle=(200, {"input_tokens": 2177})) + + result: Final = await _counter(bedrock).count_tokens( + model_to_use=model_to_use, + messages=MESSAGES, + contents=None, + deployment=DEPLOYMENT, + request_model=model_to_use, + ) + + assert result is not None + assert result.error is True + assert result.status_code == runtime[0] + assert result.tokenizer_type == "bedrock_api" + assert bedrock.posted_hosts() == (RUNTIME_HOST,) + + +@pytest.mark.asyncio +async def test_mantle_failure_is_reported_with_its_status() -> None: + bedrock: Final = _Bedrock( + runtime=(400, UNSUPPORTED), + mantle=(404, {"type": "error", "error": {"type": "not_found_error", "message": "does not exist"}}), + ) + + result: Final = await _counter(bedrock).count_tokens( + model_to_use="global.anthropic.claude-sonnet-5-5", + messages=MESSAGES, + contents=None, + deployment=DEPLOYMENT, + request_model="claude-sonnet-5-5", + ) + + assert result is not None + assert result.error is True + assert result.status_code == 404 + assert result.tokenizer_type == "bedrock_mantle_api" + assert "does not exist" in (result.error_message or "") + assert bedrock.posted_hosts() == (RUNTIME_HOST, MANTLE_HOST) diff --git a/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py b/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py index 4bfac61545d..6685deccf65 100644 --- a/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py +++ b/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py @@ -2216,6 +2216,216 @@ class TestBedrockFileContentTransformation: assert "x-amz-content-sha256" in authorization assert "X-Amz-Date" in signed_headers + def test_transform_retrieve_file_request_adds_unsigned_range(self, monkeypatch): + from litellm.llms.bedrock.files.transformation import ( + S3_SIGNED_REQUEST_HEADERS_PARAM, + BedrockFilesConfig, + ) + + monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket") + litellm_params = self._litellm_params() + url, params = BedrockFilesConfig().transform_retrieve_file_request( + file_id=self.S3_URI, + optional_params={}, + litellm_params=litellm_params, + ) + + assert url == self.EXPECTED_URL + assert params == {} + assert litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM]["Range"] == "bytes=0-0" + + @pytest.mark.parametrize( + ("file_id", "purpose"), + [ + ("s3://my-bucket/litellm-batch-outputs/job-123/output.jsonl", "batch_output"), + ("s3://my-bucket/litellm-bedrock-files-job-123/input.jsonl", "batch"), + ], + ) + def test_transform_retrieve_file_response_parses_metadata( + self, file_id: str, purpose: str, monkeypatch: pytest.MonkeyPatch + ) -> None: + import httpx + + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket") + litellm_params = self._litellm_params() + BedrockFilesConfig().transform_retrieve_file_request( + file_id=file_id, + optional_params={}, + litellm_params=litellm_params, + ) + response = BedrockFilesConfig().transform_retrieve_file_response( + raw_response=httpx.Response( + 206, + headers={ + "Content-Range": "bytes 0-0/4321", + "Last-Modified": "Wed, 21 Oct 2015 07:28:00 GMT", + }, + request=httpx.Request("GET", file_id), + ), + logging_obj=MagicMock(), + litellm_params=litellm_params, + ) + + assert response.id == file_id + assert response.bytes == 4321 + assert response.created_at == 1445412480 + assert response.filename == file_id.rsplit("/", 1)[-1] + assert response.purpose == purpose + assert response.status == "processed" + assert response.object == "file" + + @pytest.mark.parametrize( + ("file_id", "purpose"), + [ + pytest.param( + "s3://out-bucket/outpfx/litellm-batch-outputs/job-123/x.jsonl.out", + "batch_output", + id="output-bucket", + ), + pytest.param( + "s3://in-bucket/pfx/litellm-bedrock-files/job-123/input.jsonl", + "batch", + id="input-bucket-upload", + ), + pytest.param( + "s3://in-bucket/pfx/litellm-batch-outputs/job-123/x.jsonl.out", + "batch_output", + id="input-bucket-output", + ), + ], + ) + def test_transform_retrieve_file_response_uses_the_retrieved_bucket_prefix( + self, file_id: str, purpose: str, monkeypatch: pytest.MonkeyPatch + ) -> None: + import httpx + + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + monkeypatch.delenv("AWS_S3_BUCKET_NAME", raising=False) + monkeypatch.delenv("AWS_S3_OUTPUT_BUCKET_NAME", raising=False) + litellm_params: Final = _trusted_bucket_snapshot( + s3_bucket_name="in-bucket/pfx", + s3_output_bucket_name="out-bucket/outpfx", + ) + config: Final = BedrockFilesConfig() + config.transform_retrieve_file_request( + file_id=file_id, + optional_params={}, + litellm_params=litellm_params, + ) + response: Final = config.transform_retrieve_file_response( + raw_response=httpx.Response( + 206, + headers={"Content-Range": "bytes 0-0/1"}, + request=httpx.Request("GET", file_id), + ), + logging_obj=MagicMock(), + litellm_params=litellm_params, + ) + + assert response.purpose == purpose + + def test_transform_retrieve_file_response_accepts_verified_empty_object(self, monkeypatch): + import httpx + + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket") + litellm_params = self._litellm_params() + BedrockFilesConfig().transform_retrieve_file_request( + file_id=self.S3_URI, + optional_params={}, + litellm_params=litellm_params, + ) + response = BedrockFilesConfig().transform_retrieve_file_response( + raw_response=httpx.Response( + 416, + content=b"InvalidRange0", + request=httpx.Request("GET", self.S3_URI), + ), + logging_obj=MagicMock(), + litellm_params=litellm_params, + ) + + assert response.bytes == 0 + assert response.filename == "input.jsonl.out" + assert response.purpose == "batch_output" + + def test_transform_retrieve_file_response_rejects_unverified_empty_object(self, monkeypatch): + import httpx + + from litellm.llms.bedrock.common_utils import BedrockError + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket") + litellm_params = self._litellm_params() + BedrockFilesConfig().transform_retrieve_file_request( + file_id=self.S3_URI, + optional_params={}, + litellm_params=litellm_params, + ) + with pytest.raises(BedrockError): + BedrockFilesConfig().transform_retrieve_file_response( + raw_response=httpx.Response( + 416, + content=b"InvalidRange", + request=httpx.Request("GET", self.S3_URI), + ), + logging_obj=MagicMock(), + litellm_params=litellm_params, + ) + + def test_transform_retrieve_file_response_uses_content_length_when_range_is_ignored(self, monkeypatch): + import httpx + + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket") + litellm_params = self._litellm_params() + BedrockFilesConfig().transform_retrieve_file_request( + file_id=self.S3_URI, + optional_params={}, + litellm_params=litellm_params, + ) + response = BedrockFilesConfig().transform_retrieve_file_response( + raw_response=httpx.Response( + 200, + headers={"Content-Length": "4321"}, + request=httpx.Request("GET", self.S3_URI), + ), + logging_obj=MagicMock(), + litellm_params=litellm_params, + ) + + assert response.bytes == 4321 + + def test_transform_retrieve_file_response_raises_on_s3_error(self, monkeypatch): + import httpx + + from litellm.llms.bedrock.common_utils import BedrockError + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket") + litellm_params = self._litellm_params() + BedrockFilesConfig().transform_retrieve_file_request( + file_id=self.S3_URI, + optional_params={}, + litellm_params=litellm_params, + ) + + with pytest.raises(BedrockError, match="AccessDenied"): + BedrockFilesConfig().transform_retrieve_file_response( + raw_response=httpx.Response( + 403, + content=b"AccessDenied", + request=httpx.Request("GET", self.S3_URI), + ), + logging_obj=MagicMock(), + litellm_params=litellm_params, + ) + def test_transform_file_content_request_decodes_unified_file_id(self, monkeypatch): """Base64 unified ids carrying llm_output_file_id must resolve to their S3 object.""" import base64 diff --git a/tests/unit/llms/bedrock/image/test_amazon_nova_canvas_transformation.py b/tests/unit/llms/bedrock/image/test_amazon_nova_canvas_transformation.py index 2802016c04a..66001ad86be 100644 --- a/tests/unit/llms/bedrock/image/test_amazon_nova_canvas_transformation.py +++ b/tests/unit/llms/bedrock/image/test_amazon_nova_canvas_transformation.py @@ -1,4 +1,11 @@ +import json +from typing import Final + +import httpx import pytest +import respx + +import litellm from litellm.llms.bedrock.image_generation.amazon_nova_canvas_transformation import ( AmazonNovaCanvasConfig, ) @@ -79,3 +86,30 @@ def test_transform_response_dict_to_openai_response(): assert hasattr(result, "data") assert len(result.data) == 2 assert result.data[0].b64_json == "b64img1" + + +_NOVA_CANVAS_PROMPT: Final = "A serene mountain landscape at sunset with a lake reflection" +_NOVA_CANVAS_IMAGES: Final = ("b64-first-image", "b64-second-image") + + +def test_nova_canvas_image_gen_reports_positive_response_cost(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.post( + url__regex=r"^https://bedrock-runtime\.us-east-1\.amazonaws\.com/model/amazon\.nova-canvas-v1(:|%3A)0/invoke$" + ).mock(return_value=httpx.Response(200, json={"images": list(_NOVA_CANVAS_IMAGES)})) + + response: Final = litellm.image_generation( + model="bedrock/amazon.nova-canvas-v1:0", + prompt=_NOVA_CANVAS_PROMPT, + aws_region_name="us-east-1", + aws_access_key_id="fake-access-key", + aws_secret_access_key="fake-secret-key", + ) + + assert route.call_count == 1 + sent: Final = json.loads(route.calls[0].request.content) + assert sent["taskType"] == "TEXT_IMAGE" + assert sent["textToImageParams"]["text"] == _NOVA_CANVAS_PROMPT + assert [image.b64_json for image in response.data] == list(_NOVA_CANVAS_IMAGES) + per_image: Final = litellm.model_cost["amazon.nova-canvas-v1:0"]["output_cost_per_image"] + assert per_image > 0 + assert response._hidden_params["response_cost"] == pytest.approx(len(_NOVA_CANVAS_IMAGES) * per_image) # pyright: ignore[reportPrivateUsage] # cost is only surfaced on _hidden_params diff --git a/tests/unit/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py b/tests/unit/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py index fbcbda183bc..875a2fc1b1a 100644 --- a/tests/unit/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py +++ b/tests/unit/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py @@ -10,9 +10,14 @@ from datetime import datetime from typing import Final from unittest.mock import patch +import pytest + from litellm.litellm_core_utils.litellm_logging import Logging from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector -from litellm.llms.bedrock.passthrough.transformation import BedrockPassthroughConfig +from litellm.llms.bedrock.passthrough.transformation import ( + BedrockPassthroughConfig, + is_bedrock_streaming_endpoint, +) from litellm.types.utils import ModelResponse CONVERSE_MODEL = "anthropic.claude-sonnet-4-5-20250929-v1:0" @@ -20,6 +25,22 @@ CONVERSE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/converse-stream" INVOKE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/invoke-with-response-stream" +@pytest.mark.parametrize( + ("endpoint", "expected"), + [ + ("converse-stream", True), + ("invoke-with-response-stream", True), + ("converse", False), + ("invoke", False), + ("model/my-converse-stream-model/converse", False), + ("model/x/converse-stream?foo=1", True), + ("model/x/converse-stream/", True), + ], +) +def test_is_bedrock_streaming_endpoint_matches_final_action_segment(endpoint: str, expected: bool) -> None: + assert is_bedrock_streaming_endpoint(endpoint) is expected + + def test_bedrock_passthrough_get_complete_url_default_endpoint(): """Test get_complete_url with default AWS endpoint (no override)""" config = BedrockPassthroughConfig() diff --git a/tests/unit/llms/bedrock/test_base_aws_llm.py b/tests/unit/llms/bedrock/test_base_aws_llm.py index f1c95aaf7a2..cc625dcf3f3 100644 --- a/tests/unit/llms/bedrock/test_base_aws_llm.py +++ b/tests/unit/llms/bedrock/test_base_aws_llm.py @@ -745,40 +745,23 @@ def test_sign_request_with_api_key_bearer_token(): assert result_body == json.dumps(request_data).encode() -def test_get_request_headers_with_env_var_bearer_token(): - # Setup - llm = BaseAWSLLM() - credentials = Credentials("test_key", "test_secret", "test_token") - headers = {"Content-Type": "application/json"} - headers_dict = headers.copy() - - # Create mock request - mock_prepared_request = MagicMock(spec=AWSPreparedRequest) - mock_request = MagicMock(spec=AWSRequest) - mock_request.headers = headers_dict - mock_request.prepare.return_value = mock_prepared_request - - def mock_aws_request_init(method, url, data, headers): - mock_request.headers.update(headers) - return mock_request - - # Test with bearer token - with ( - patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": "test_token"}), - patch("botocore.awsrequest.AWSRequest", side_effect=mock_aws_request_init), - ): - result = llm.get_request_headers( - credentials=credentials, +@pytest.mark.parametrize("from_environment", [True, False]) +def test_get_request_headers_preserves_bearer_payload(from_environment): + with patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": "test_token"} if from_environment else {}, clear=True): + result = BaseAWSLLM().get_request_headers( + credentials=None, aws_region_name="us-west-2", extra_headers=None, endpoint_url="https://api.example.com", data='{"prompt": "test"}', - headers=headers_dict, + headers={"Content-Type": "application/json"}, + api_key=None if from_environment else "test_token", ) - - # Assert - assert mock_request.headers["Authorization"] == "Bearer test_token" - assert result == mock_prepared_request + assert result.headers["Authorization"] == "Bearer test_token" + assert result.headers["Content-Type"] == "application/json" + assert result.body == b'{"prompt": "test"}' + assert result.url == "https://api.example.com" + assert result.method == "POST" def test_get_request_headers_with_sigv4(): @@ -856,46 +839,6 @@ def test_sigv4_matches_rust_golden_vector(): ) -def test_get_request_headers_with_api_key_bearer_token(): - """ - Test that get_request_headers uses the api_key parameter as a bearer token when provided - """ - # Setup - llm = BaseAWSLLM() - credentials = Credentials("test_key", "test_secret", "test_token") - headers = {"Content-Type": "application/json"} - headers_dict = headers.copy() - api_key = "test_api_key" - - # Create mock request - mock_prepared_request = MagicMock(spec=AWSPreparedRequest) - mock_request = MagicMock(spec=AWSRequest) - mock_request.headers = headers_dict - mock_request.prepare.return_value = mock_prepared_request - - def mock_aws_request_init(method, url, data, headers): - mock_request.headers.update(headers) - return mock_request - - # Test with api_key parameter - with ( - patch.dict(os.environ, {}, clear=True), - patch("botocore.awsrequest.AWSRequest", side_effect=mock_aws_request_init), - ): - result = llm.get_request_headers( - credentials=credentials, - aws_region_name="us-west-2", - extra_headers=None, - endpoint_url="https://api.example.com", - data='{"prompt": "test"}', - headers=headers_dict, - api_key=api_key, - ) - - # Assert - assert mock_request.headers["Authorization"] == f"Bearer {api_key}" - assert result == mock_prepared_request - def test_role_assumption_without_session_name(): """ @@ -4171,3 +4114,104 @@ def test_dynamic_aws_params_propagation(model, param_name, param_value, expected # We now assert that get_credentials() was called with the dynamic param. assert dummy_get_credentials.called_kwargs.get(param_name) == expected_credentials_value + + +def test_bearer_request_preparation_does_not_require_botocore(): + import httpx + + with patch.dict("sys.modules", {"botocore.credentials": None, "botocore.awsrequest": None}): + target = BaseAWSLLM()._get_boto_credentials_from_optional_params( + {"aws_region_name": "us-west-2"}, bearer_token="test-token" + ) + request = BaseAWSLLM().get_request_headers( + credentials=None, + aws_region_name=target.aws_region_name, + extra_headers=None, + endpoint_url="https://bedrock-runtime.us-west-2.amazonaws.com/model/test/invoke", + data='{"text":"café"}', + headers={"Content-Type": "application/json"}, + api_key="test-token", + ) + assert dict(request.headers)["Authorization"] == "Bearer test-token" + assert request.body == '{"text":"café"}'.encode() + assert int(httpx.Headers(request.headers)["Content-Length"]) == len(request.body) + + +@pytest.mark.parametrize("missing", ["botocore", "unrelated_dependency"]) +def test_signing_preserves_unrelated_import_failure(missing): + llm = BaseAWSLLM() + failure = ModuleNotFoundError("missing dependency", name=missing) + with patch("builtins.__import__", side_effect=failure): + with pytest.raises(ImportError) as error: + llm.get_request_headers( + credentials=Credentials("key", "secret"), aws_region_name="us-east-1", + extra_headers=None, endpoint_url="https://bedrock-runtime.us-east-1.amazonaws.com", + data="{}", headers={}, supports_bearer_token=False, + ) + if missing == "botocore": + assert "pip install boto3" in str(error.value) + assert error.value.__cause__ is failure + else: + assert error.value is failure + + +@pytest.mark.parametrize("missing", ["botocore", "unrelated_dependency"]) +@pytest.mark.parametrize("shared_signer", [False, True]) +def test_json_signers_report_only_missing_aws_dependency(missing, shared_signer): + from litellm.llms.bedrock.base_aws_llm import sign_aws_json_post + + llm = BaseAWSLLM() + failure = ModuleNotFoundError("missing dependency", name=missing) + from functools import partial + + sign = ( + partial(sign_aws_json_post, lambda: Credentials("key", "secret"), "bedrock", "us-east-1", + "https://bedrock-runtime.us-east-1.amazonaws.com", "{}", {}) + if shared_signer else + partial(llm._sign_request, service_name="bedrock", headers={}, optional_params={}, request_data={}, + api_base="https://bedrock-runtime.us-east-1.amazonaws.com", api_key="") + ) + with patch.dict(os.environ, {}, clear=True), patch("builtins.__import__", side_effect=failure): + with pytest.raises(ImportError) as error: + sign() + if missing == "botocore": + assert "pip install boto3" in str(error.value) + assert error.value.__cause__ is failure + else: + assert error.value is failure + + +@pytest.mark.parametrize("missing", ["botocore", "unrelated_dependency"]) +def test_direct_credentials_report_missing_aws_without_masking_other_imports(missing): + import builtins + + original_import = builtins.__import__ + failure = ModuleNotFoundError("dependency unavailable", name=missing) + + def import_dependency(name, *args, **kwargs): + if name.startswith("botocore"): + raise failure + return original_import(name, *args, **kwargs) + + with patch("builtins.__import__", side_effect=import_dependency): + with pytest.raises(ImportError) as error: + BaseAWSLLM().get_credentials( + aws_access_key_id="test-key", aws_secret_access_key="test-secret", aws_session_token="test-session" + ) + if missing == "botocore": + assert "pip install boto3" in str(error.value) + assert error.value.__cause__ is failure + else: + assert error.value is failure + + +def test_shared_json_signer_preserves_body_and_signs_for_the_requested_service(): + from litellm.llms.bedrock.base_aws_llm import sign_aws_json_post + + body = '{"message":"ping"}' + request = sign_aws_json_post( + lambda: Credentials("test-key", "test-secret"), "s3", "us-west-2", + "https://s3.us-west-2.amazonaws.com", body, {"Content-Type": "application/json"}, + ) + assert request.body == body + assert "/us-west-2/s3/aws4_request" in request.headers["Authorization"] diff --git a/tests/unit/llms/bedrock/test_bedrock_common_utils.py b/tests/unit/llms/bedrock/test_bedrock_common_utils.py index b3e31b36373..3a2c851807c 100644 --- a/tests/unit/llms/bedrock/test_bedrock_common_utils.py +++ b/tests/unit/llms/bedrock/test_bedrock_common_utils.py @@ -1989,3 +1989,19 @@ class TestBedrockGovCloudSupport: """Test that GovCloud Titan models use Invoke API""" route = BedrockModelInfo.get_bedrock_route(model_name) assert route == "invoke" + + +@pytest.mark.parametrize("missing", ["botocore", "unrelated_dependency"]) +def test_event_decoder_reports_only_missing_aws_dependency(missing): + from unittest.mock import patch + from litellm.llms.bedrock.common_utils import BedrockEventStreamDecoderBase + + failure = ModuleNotFoundError("missing dependency", name=missing) + with patch("builtins.__import__", side_effect=failure): + with pytest.raises(ImportError) as error: + BedrockEventStreamDecoderBase() + if missing == "botocore": + assert "pip install boto3" in str(error.value) + assert error.value.__cause__ is failure + else: + assert error.value is failure diff --git a/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py b/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py index 9c35e3253e5..ea9cfc1fd39 100644 --- a/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py +++ b/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py @@ -1042,3 +1042,18 @@ async def test_mantle_signing_runs_off_the_event_loop(): assert "Authorization" in signed assert probe.served_during_refresh is True + + +@pytest.mark.parametrize("missing", ["botocore", "unrelated_dependency"]) +def test_mantle_signing_reports_only_missing_aws_dependency(missing): + config = BedrockMantleChatConfig() + failure = ModuleNotFoundError("missing dependency", name=missing) + with patch.dict("os.environ", {}, clear=True), patch("builtins.__import__", side_effect=failure): + with pytest.raises(ImportError) as error: + config.sign_request(headers={}, optional_params={}, request_data={}, + api_base="https://bedrock-mantle.us-east-1.api.aws/v1/chat/completions", api_key="") + if missing == "botocore": + assert "pip install boto3" in str(error.value) + assert error.value.__cause__ is failure + else: + assert error.value is failure diff --git a/tests/unit/llms/custom_httpx/test_async_client_cleanup.py b/tests/unit/llms/custom_httpx/test_async_client_cleanup.py index e8fb0808019..6b7697be945 100644 --- a/tests/unit/llms/custom_httpx/test_async_client_cleanup.py +++ b/tests/unit/llms/custom_httpx/test_async_client_cleanup.py @@ -1,4 +1,10 @@ +import re +from collections.abc import Iterator +from typing import Final + +import httpx import pytest +import respx import litellm from litellm.llms.custom_httpx.async_client_cleanup import close_litellm_async_clients @@ -19,3 +25,97 @@ async def test_second_cleanup_pass_does_not_resurrect_owned_client(): litellm.in_memory_llm_clients_cache.cache_dict.pop(cache_key, None) assert handler._client is original_client + + +GEMINI_GENERATE_URL: Final = re.compile(r"https://generativelanguage\.googleapis\.com/.*:generateContent.*") + + +def _gemini_reply(text: str) -> httpx.Response: + return httpx.Response( + 200, + json={ + "candidates": [ + {"content": {"parts": [{"text": text}], "role": "model"}, "finishReason": "STOP", "index": 0} + ], + "usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }, + ) + + +def _cached_async_handlers() -> list[AsyncHTTPHandler]: + return [ + handler + for handler in litellm.in_memory_llm_clients_cache.cache_dict.values() + if isinstance(handler, AsyncHTTPHandler) + ] + + +@pytest.fixture +def gemini_httpx_transport(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setenv("GEMINI_API_KEY", "gemini-cleanup-test") + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("gemini_httpx_transport") +async def test_acompletion_client_is_closed_by_cleanup() -> None: + with respx.mock(assert_all_called=True) as mock: + mock.post(GEMINI_GENERATE_URL).mock(return_value=_gemini_reply("Hi there!")) + response: Final = await litellm.acompletion( + model="gemini/gemini-2.0-flash-lite-001", + messages=[{"role": "user", "content": "Hello"}], + ) + assert response.choices[0].message.content == "Hi there!" + clients: Final = [handler._client for handler in _cached_async_handlers()] + assert clients + assert not any(client.is_closed for client in clients) + + await close_litellm_async_clients() + + assert all(client.is_closed for client in clients) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("gemini_httpx_transport") +async def test_repeated_acompletion_calls_reuse_one_client_that_cleanup_closes() -> None: + with respx.mock(assert_all_called=True) as mock: + route: Final = mock.post(GEMINI_GENERATE_URL).mock( + side_effect=[_gemini_reply(f"Response {index}") for index in range(3)] + ) + for index in range(3): + response = await litellm.acompletion( + model="gemini/gemini-2.0-flash-lite-001", + messages=[{"role": "user", "content": f"Hello {index}"}], + ) + assert response.choices[0].message.content == f"Response {index}" + assert route.call_count == 3 + handlers: Final = _cached_async_handlers() + assert len(handlers) == 1 + client: Final = handlers[0]._client + + await close_litellm_async_clients() + + assert client.is_closed + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("gemini_httpx_transport") +async def test_cleanup_is_idempotent_and_acompletion_works_afterwards() -> None: + with respx.mock(assert_all_called=True) as mock: + route: Final = mock.post(GEMINI_GENERATE_URL).mock(side_effect=[_gemini_reply("Hello!"), _gemini_reply("Hi!")]) + await litellm.acompletion( + model="gemini/gemini-2.0-flash-lite-001", + messages=[{"role": "user", "content": "Hello"}], + ) + for _ in range(3): + await close_litellm_async_clients() + response: Final = await litellm.acompletion( + model="gemini/gemini-2.0-flash-lite-001", + messages=[{"role": "user", "content": "Hello"}], + ) + assert response.choices[0].message.content == "Hi!" + assert route.call_count == 2 + await close_litellm_async_clients() diff --git a/tests/unit/llms/custom_httpx/test_llm_http_handler.py b/tests/unit/llms/custom_httpx/test_llm_http_handler.py index 82cc9e9d3ac..25d2d11b5ca 100644 --- a/tests/unit/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_llm_http_handler.py @@ -45,6 +45,7 @@ from litellm.llms.azure.videos.transformation import AzureVideoConfig from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import ( AmazonAnthropicClaudeMessagesConfig, ) +from litellm.llms.bedrock.files.transformation import BedrockFilesConfig from litellm.llms.anthropic.skills.transformation import AnthropicSkillsConfig from litellm.llms.openai.evals.transformation import OpenAIEvalsConfig from litellm.llms.mistral.files.transformation import MistralFilesConfig @@ -4925,3 +4926,133 @@ async def test_lookup_handlers_raise_the_provider_error_status(name: str, is_asy assert error.value.status_code == status_code assert "No such object" in error.value.message + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async", (False, True)) +@pytest.mark.parametrize( + ("status_code", "content", "headers", "expected_bytes"), + ( + ( + 416, + b"InvalidRange0", + {}, + 0, + ), + (206, b"", {"Content-Range": "bytes 0-0/4321"}, 4321), + ), +) +async def test_retrieve_file_accepts_bedrock_successful_range_responses( + is_async: bool, + status_code: int, + content: bytes, + headers: dict[str, str], + expected_bytes: int, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket") + file_id = "s3://my-bucket/litellm-batch-outputs/job-123/output.jsonl" + transport = httpx.MockTransport( + lambda request: httpx.Response( + status_code, + content=content, + headers=headers, + request=request, + ) + ) + params = { + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "secret", + "aws_region_name": "us-west-2", + } + handler = BaseLLMHTTPHandler() + + if is_async: + client = AsyncHTTPHandler() + await client.close() + async_client = httpx.AsyncClient(transport=transport) + client.client = async_client + try: + result = await handler.async_retrieve_file( + file_id=file_id, + provider_config=BedrockFilesConfig(), + litellm_params=params, + headers={}, + logging_obj=Mock(), + client=client, + ) + finally: + await async_client.aclose() + else: + sync_client = httpx.Client(transport=transport) + client = HTTPHandler(client=sync_client) + try: + result = handler.retrieve_file( + file_id=file_id, + provider_config=BedrockFilesConfig(), + litellm_params=params, + headers={}, + logging_obj=Mock(), + client=client, + ) + finally: + sync_client.close() + + assert result.bytes == expected_bytes + assert result.filename == "output.jsonl" + assert result.purpose == "batch_output" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async", (False, True)) +async def test_retrieve_file_rejects_unverified_bedrock_range_error( + is_async: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("AWS_S3_BUCKET_NAME", "my-bucket") + file_id = "s3://my-bucket/litellm-batch-outputs/job-123/output.jsonl" + transport = httpx.MockTransport( + lambda request: httpx.Response( + 416, + content=b"InvalidRange", + request=request, + ) + ) + params = { + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "secret", + "aws_region_name": "us-west-2", + } + handler = BaseLLMHTTPHandler() + + if is_async: + client = AsyncHTTPHandler() + await client.close() + async_client = httpx.AsyncClient(transport=transport) + client.client = async_client + try: + with pytest.raises(BaseLLMException, match="InvalidRange"): + await handler.async_retrieve_file( + file_id=file_id, + provider_config=BedrockFilesConfig(), + litellm_params=params, + headers={}, + logging_obj=Mock(), + client=client, + ) + finally: + await async_client.aclose() + else: + sync_client = httpx.Client(transport=transport) + client = HTTPHandler(client=sync_client) + try: + with pytest.raises(BaseLLMException, match="InvalidRange"): + handler.retrieve_file( + file_id=file_id, + provider_config=BedrockFilesConfig(), + litellm_params=params, + headers={}, + logging_obj=Mock(), + client=client, + ) + finally: + sync_client.close() diff --git a/tests/unit/llms/duckduckgo/search/test_duckduckgo_search_transformation.py b/tests/unit/llms/duckduckgo/search/test_duckduckgo_search_transformation.py index 474ffe0e519..68419891166 100644 --- a/tests/unit/llms/duckduckgo/search/test_duckduckgo_search_transformation.py +++ b/tests/unit/llms/duckduckgo/search/test_duckduckgo_search_transformation.py @@ -1,8 +1,12 @@ +from collections.abc import Iterator from typing import Final from unittest.mock import AsyncMock, MagicMock, patch -import litellm +import httpx import pytest +import respx + +import litellm class TestDuckDuckGoSearchMocked: @@ -223,3 +227,73 @@ class TestDuckDuckGoSearchMocked: urls = [result.url for result in response.results] assert any("India" in url for url in urls) assert any("Indus" in url for url in urls) + + +_DDG_INSTANT_ANSWER: Final = { + "AbstractText": "India is a country in South Asia.", + "AbstractURL": "https://en.wikipedia.org/wiki/India", + "Heading": "India", + "RelatedTopics": [ + {"FirstURL": f"https://example.com/{index}", "Text": f"Topic {index} - snippet text for topic {index}."} + for index in range(10) + ], + "Results": [], + "Type": "D", +} + + +@pytest.fixture +def httpx_transport(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.asyncio +async def test_duckduckgo_search_response_structure_and_max_results( + respx_mock: respx.MockRouter, httpx_transport: None +) -> None: + route: Final = respx_mock.get(url__startswith="https://api.duckduckgo.com/").mock( + return_value=httpx.Response(200, json=_DDG_INSTANT_ANSWER) + ) + + response: Final = await litellm.asearch(query="india", search_provider="duckduckgo", max_results=5) + + assert route.call_count == 1 + sent_params: Final = route.calls[0].request.url.params + assert sent_params["q"] == "india" + assert sent_params["format"] == "json" + assert sent_params["_max_results"] == "5" + assert response.object == "search" + assert [result.url for result in response.results] == [ + "https://en.wikipedia.org/wiki/India", + "https://example.com/0", + "https://example.com/1", + "https://example.com/2", + "https://example.com/3", + ] + first_result: Final = response.results[0] + assert first_result.title == "India" + assert first_result.snippet == "India is a country in South Asia." + assert response.results[1].title == "Topic 0" + assert response.results[1].snippet == "snippet text for topic 0." + assert response._hidden_params["response_cost"] == litellm.model_cost["duckduckgo/search"]["input_cost_per_query"] # pyright: ignore[reportPrivateUsage] # cost is only surfaced on _hidden_params + + +def test_duckduckgo_sync_search_returns_typed_results_without_a_limit(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.get(url__startswith="https://api.duckduckgo.com/").mock( + return_value=httpx.Response(200, json=_DDG_INSTANT_ANSWER) + ) + + response: Final = litellm.search(query="india", search_provider="duckduckgo") + + assert route.call_count == 1 + assert "_max_results" not in route.calls[0].request.url.params + assert response.object == "search" + assert len(response.results) == 11 + assert all( + isinstance(result.title, str) and isinstance(result.url, str) and isinstance(result.snippet, str) + for result in response.results + ) + assert response.results[-1].url == "https://example.com/9" diff --git a/tests/unit/llms/exa_ai/search/test_transformation.py b/tests/unit/llms/exa_ai/search/test_transformation.py index 5e5eb24f23b..6b556ff241f 100644 --- a/tests/unit/llms/exa_ai/search/test_transformation.py +++ b/tests/unit/llms/exa_ai/search/test_transformation.py @@ -1,11 +1,22 @@ +import json from typing import Final from unittest.mock import Mock import httpx import pytest +import respx +import litellm from litellm.llms.exa_ai.search.transformation import ExaAISearchConfig +EXA_SEARCH_URL: Final = "https://api.exa.ai/search" + + +@pytest.fixture +def exa_api_key(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("EXA_API_KEY", "test-exa-key") + monkeypatch.delenv("EXA_API_BASE", raising=False) + @pytest.mark.parametrize( ("content_fields", "expected_snippet"), @@ -31,3 +42,49 @@ def test_transform_search_response_snippet_falls_back_through_content_modes( response: Final = ExaAISearchConfig().transform_search_response(raw_response, logging_obj=Mock()) assert response.results[0].snippet == expected_snippet + + +@pytest.mark.usefixtures("exa_api_key") +def test_search_maps_exa_results_to_search_response(respx_mock: respx.MockRouter) -> None: + respx_mock.post(EXA_SEARCH_URL).respond( + json={ + "results": [ + { + "title": "AI news roundup", + "url": "https://example.com/ai-news", + "text": "The latest in artificial intelligence.", + "publishedDate": "2026-01-15T00:00:00.000Z", + }, + {"title": "Second", "url": "https://example.com/second", "text": "Second text."}, + ] + } + ) + + response: Final = litellm.search(query="artificial intelligence recent news", search_provider="exa_ai") + + assert response.object == "search" + assert isinstance(response.results, list) + assert len(response.results) == 2 + first: Final = response.results[0] + assert first.title == "AI news roundup" + assert first.url == "https://example.com/ai-news" + assert first.snippet == "The latest in artificial intelligence." + assert first.date == "2026-01-15T00:00:00.000Z" + assert response.results[1].url == "https://example.com/second" + + +@pytest.mark.usefixtures("exa_api_key") +def test_search_sends_max_results_as_num_results(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.post(EXA_SEARCH_URL).respond( + json={"results": [{"title": "ML", "url": "https://example.com/ml", "text": "Machine learning."}]} + ) + + response: Final = litellm.search(query="machine learning", search_provider="exa_ai", max_results=5) + + assert json.loads(route.calls.last.request.content) == { + "query": "machine learning", + "numResults": 5, + "contents": {"text": True}, + } + assert route.calls.last.request.headers["x-api-key"] == "test-exa-key" + assert [result.url for result in response.results] == ["https://example.com/ml"] diff --git a/tests/unit/llms/firecrawl/__init__.py b/tests/unit/llms/firecrawl/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/firecrawl/search/__init__.py b/tests/unit/llms/firecrawl/search/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/firecrawl/search/test_transformation.py b/tests/unit/llms/firecrawl/search/test_transformation.py new file mode 100644 index 00000000000..7d18b40064b --- /dev/null +++ b/tests/unit/llms/firecrawl/search/test_transformation.py @@ -0,0 +1,32 @@ +import json +from typing import Final + +import httpx +import pytest +import respx + +import litellm + + +def test_firecrawl_search_request_body(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("FIRECRAWL_API_KEY", "test-api-key") + route: Final = respx_mock.post("https://api.firecrawl.dev/v2/search").mock( + return_value=httpx.Response( + 200, + json={ + "success": True, + "data": {"web": [{"title": "Test Title", "url": "https://example.com", "markdown": "Test content"}]}, + }, + ) + ) + + response: Final = litellm.search(query="test query", search_provider="firecrawl", max_results=10, country="US") + + assert route.call_count == 1 + sent: Final = route.calls[0].request + assert sent.headers["authorization"] == "Bearer test-api-key" + body: Final = json.loads(sent.content) + assert body["query"] == "test query" + assert body["limit"] == 10 + assert body["country"] == "US" + assert [(result.title, result.url) for result in response.results] == [("Test Title", "https://example.com")] diff --git a/tests/unit/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/unit/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py index 9de8d4c4135..5a013b19a95 100644 --- a/tests/unit/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py +++ b/tests/unit/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py @@ -2008,3 +2008,126 @@ VISION_MODEL = next( for key, info in litellm.model_cost.items() if key.startswith("fireworks_ai/accounts/fireworks/models/") and info.get("supports_vision") is True ) + + +def _fireworks_chat_client() -> MagicMock: + body: Final = { + "id": "chat-user-attribution", + "object": "chat.completion", + "created": 1, + "model": "accounts/fireworks/models/kimi-k3", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } + raw_response: Final = MagicMock() + raw_response.status_code = 200 + raw_response.headers = {} + raw_response.text = json.dumps(body) + raw_response.json = lambda: body + client: Final = MagicMock(spec=HTTPHandler) + client.post.return_value = raw_response + return client + + +@pytest.mark.parametrize( + "call_kwargs, expected_user", + [ + pytest.param( + {"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": "dev-alice"}}, + "dev-alice", + id="opted-in-sends-litellm-user-id", + ), + pytest.param( + {"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": "dev-alice"}, "user": "caller"}, + "dev-alice", + id="litellm-user-id-replaces-caller-user", + ), + pytest.param( + {"fireworks_forward_user_id": True, "litellm_metadata": {"user_api_key_user_id": "dev-bob"}}, + "dev-bob", + id="reads-litellm-metadata", + ), + pytest.param( + { + "fireworks_forward_user_id": True, + "metadata": {"tags": ["caller-tag"]}, + "litellm_metadata": {"user_api_key_user_id": "dev-bob"}, + }, + "dev-bob", + id="reads-litellm-metadata-next-to-caller-metadata", + ), + pytest.param( + {"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": ""}, "user": "caller"}, + "caller", + id="empty-litellm-user-id-keeps-caller-user", + ), + pytest.param( + {"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": None}, "user": "caller"}, + "caller", + id="no-litellm-user-id-keeps-caller-user", + ), + pytest.param( + {"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": None}}, + None, + id="no-litellm-user-id-sends-no-user", + ), + pytest.param( + {"metadata": {"user_api_key_user_id": "dev-alice"}, "user": "caller"}, + "caller", + id="not-opted-in-keeps-caller-user", + ), + pytest.param( + {"metadata": {"user_api_key_user_id": "dev-alice"}}, + None, + id="not-opted-in-sends-no-user", + ), + pytest.param( + {"fireworks_forward_user_id": "true", "metadata": {"user_api_key_user_id": "dev-alice"}}, + None, + id="non-bool-flag-is-off", + ), + pytest.param( + { + "fireworks_forward_user_id": True, + "metadata": {"user_api_key_user_id": "dev-alice"}, + "extra_body": {"user": "dev-bob", "prompt_cache_max_len": 1}, + }, + "dev-alice", + id="litellm-user-id-replaces-extra-body-user", + ), + pytest.param( + {"metadata": {"user_api_key_user_id": "dev-alice"}, "extra_body": {"user": "dev-bob"}}, + "dev-bob", + id="not-opted-in-keeps-extra-body-user", + ), + ], +) +def test_completion_forwards_litellm_user_id_as_user(call_kwargs: dict[str, object], expected_user: str | None) -> None: + client: Final = _fireworks_chat_client() + litellm.completion( + model="fireworks_ai/accounts/fireworks/models/kimi-k3", + messages=[{"role": "user", "content": "hi"}], + api_key="fw-test-key", + client=client, + **call_kwargs, + ) + request_body: Final = json.loads(client.post.call_args.kwargs["data"]) + assert request_body.get("user") == expected_user + assert "fireworks_forward_user_id" not in request_body + + +def test_completion_forwards_litellm_user_id_when_streaming() -> None: + client: Final = _fireworks_chat_client() + client.post.return_value.iter_lines = lambda: iter(()) + litellm.completion( + model="fireworks_ai/accounts/fireworks/models/kimi-k3", + messages=[{"role": "user", "content": "hi"}], + api_key="fw-test-key", + client=client, + stream=True, + fireworks_forward_user_id=True, + metadata={"user_api_key_user_id": "dev-alice"}, + ) + request_body: Final = json.loads(client.post.call_args.kwargs["data"]) + assert request_body["user"] == "dev-alice" + assert request_body["stream"] is True diff --git a/tests/unit/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py b/tests/unit/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py index 735d0b1125f..878e0285cd1 100644 --- a/tests/unit/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py +++ b/tests/unit/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py @@ -608,3 +608,62 @@ def test_streaming_responses_call_hits_native_endpoint_and_yields_every_firework assert tuple(event.type for event in received) == tuple(event["type"] for event in FIREWORKS_SSE_EVENTS) assert "".join(event.delta for event in received if event.type == "response.output_text.delta") == "pong" assert received[-1].response.usage.output_tokens == 89 + + +@pytest.mark.parametrize( + "call_kwargs, expected_user", + [ + pytest.param( + {"fireworks_forward_user_id": True, "litellm_metadata": {"user_api_key_user_id": "dev-alice"}}, + "dev-alice", + id="opted-in-sends-litellm-user-id", + ), + pytest.param( + { + "fireworks_forward_user_id": True, + "litellm_metadata": {"user_api_key_user_id": "dev-alice"}, + "user": "caller", + }, + "dev-alice", + id="litellm-user-id-replaces-caller-user", + ), + pytest.param( + {"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": "dev-alice"}}, + None, + id="ignores-caller-responses-metadata", + ), + pytest.param( + {"fireworks_forward_user_id": True, "user": "caller"}, + "caller", + id="no-litellm-user-id-keeps-caller-user", + ), + pytest.param( + {"litellm_metadata": {"user_api_key_user_id": "dev-alice"}}, + None, + id="not-opted-in-sends-no-user", + ), + pytest.param( + { + "fireworks_forward_user_id": True, + "litellm_metadata": {"user_api_key_user_id": "dev-alice"}, + "extra_body": {"user": "dev-bob"}, + }, + "dev-alice", + id="litellm-user-id-replaces-extra-body-user", + ), + pytest.param( + {"litellm_metadata": {"user_api_key_user_id": "dev-alice"}, "extra_body": {"user": "dev-bob"}}, + "dev-bob", + id="not-opted-in-keeps-extra-body-user", + ), + ], +) +def test_responses_call_forwards_litellm_user_id_as_user( + call_kwargs: Mapping[str, object], expected_user: str | None +) -> None: + client: Final = _mock_http_client(_fireworks_response("accounts/fireworks/models/kimi-k3")) + with patch(HTTPX_CLIENT_FACTORY, return_value=client): + litellm.responses(model="fireworks_ai/kimi-k3", input="hi", api_key="fw-test-key", **call_kwargs) + _, _, body = _sent_request(client) + assert body.get("user") == expected_user + assert "fireworks_forward_user_id" not in body diff --git a/tests/unit/llms/litellm_proxy/responses/__init__.py b/tests/unit/llms/litellm_proxy/responses/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/litellm_proxy/responses/test_transformation.py b/tests/unit/llms/litellm_proxy/responses/test_transformation.py new file mode 100644 index 00000000000..e2742d6a40b --- /dev/null +++ b/tests/unit/llms/litellm_proxy/responses/test_transformation.py @@ -0,0 +1,29 @@ +from typing import Final + +import pytest + +from litellm.llms.litellm_proxy.responses.transformation import LiteLLMProxyResponsesAPIConfig +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager + + +def test_provider_config_manager_returns_litellm_proxy_responses_config() -> None: + config: Final = ProviderConfigManager.get_provider_responses_api_config( + model="litellm_proxy/gpt-5.5", provider=LlmProviders.LITELLM_PROXY + ) + assert isinstance(config, LiteLLMProxyResponsesAPIConfig) + assert config.custom_llm_provider == LlmProviders.LITELLM_PROXY + + +@pytest.mark.parametrize("api_base", ["https://my-proxy.example.com", "https://my-proxy.example.com/"]) +def test_get_complete_url_appends_responses_path(api_base: str) -> None: + assert ( + LiteLLMProxyResponsesAPIConfig().get_complete_url(api_base=api_base, litellm_params={}) + == "https://my-proxy.example.com/responses" + ) + + +def test_get_complete_url_requires_api_base(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("LITELLM_PROXY_API_BASE", raising=False) + with pytest.raises(ValueError, match="api_base not set"): + LiteLLMProxyResponsesAPIConfig().get_complete_url(api_base=None, litellm_params={}) diff --git a/tests/unit/llms/openai/responses/test_openai_responses_http.py b/tests/unit/llms/openai/responses/test_openai_responses_http.py new file mode 100644 index 00000000000..29af2f2c55a --- /dev/null +++ b/tests/unit/llms/openai/responses/test_openai_responses_http.py @@ -0,0 +1,531 @@ +import asyncio +import json +import threading +from collections.abc import Iterable +from types import MappingProxyType +from typing import Final, Literal, TypeAlias + +import httpx +import pytest +import respx +from pydantic import BaseModel, TypeAdapter + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.types.llms.openai import ( + IncompleteDetails, + ResponseAPIUsage, + ResponseCompletedEvent, + ResponsesAPIResponse, +) +from litellm.types.utils import StandardLoggingPayload, Usage +from tests.unit.proxy.conftest import httpx_transport + +pytestmark: Final = pytest.mark.usefixtures(httpx_transport.__name__) +_OPENAI_URL: Final = "https://api.openai.com/v1/responses" +_JSON_OBJECT: Final = TypeAdapter(dict[str, object]) +_STANDARD_LOGGING_PAYLOAD: Final = TypeAdapter(StandardLoggingPayload) +_OUTPUT_TEXT: Final = "Hello from the mocked response" +_CREATED_AT: Final = 1750000000 +_AsyncLoggingMode: TypeAlias = Literal["non_stream", "stream"] +_RESPONSE_FIELD_TYPES: Final = MappingProxyType( + { + "error": (dict, type(None)), + "incomplete_details": (IncompleteDetails, type(None)), + "instructions": (str, type(None)), + "metadata": (dict,), + "model": (str,), + "object": (str,), + "parallel_tool_calls": (bool, type(None)), + "temperature": (int, float, type(None)), + "tool_choice": (dict, str, type(None)), + "tools": (list, type(None)), + "top_p": (int, float, type(None)), + "max_output_tokens": (int, type(None)), + "previous_response_id": (str, type(None)), + "reasoning": (dict, type(None)), + "status": (str,), + "text": (dict,), + "truncation": (str, type(None)), + "user": (str, type(None)), + "store": (bool, type(None)), + } +) +_STREAM_EVENT_FIELDS: Final = MappingProxyType( + { + "response.created": ("response",), + "response.in_progress": ("response",), + "response.output_item.added": ("output_index", "item"), + "response.content_part.added": ("item_id", "output_index", "content_index", "part"), + "response.output_text.delta": ("item_id", "output_index", "content_index", "delta"), + "response.output_text.done": ("item_id", "output_index", "content_index", "text"), + "response.content_part.done": ("item_id", "output_index", "content_index", "part"), + "response.output_item.done": ("output_index", "item"), + "response.completed": ("response",), + } +) + + +class _ResponsesLoggingCapture(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.completed: Final = threading.Event() + self.payload: StandardLoggingPayload | None = None + self.usage: object = None + + async def async_log_success_event( + self, + kwargs: dict[str, object], + response_obj: object, + start_time: object, + end_time: object, + ) -> None: + assert isinstance(response_obj, ResponsesAPIResponse) + self.payload = _STANDARD_LOGGING_PAYLOAD.validate_python(kwargs["standard_logging_object"]) + self.usage = response_obj.usage + self.completed.set() + + +def _response_body(response_id: str, store: bool = False, created_at: float = _CREATED_AT) -> dict[str, object]: + return { + "id": response_id, + "object": "response", + "created_at": created_at, + "status": "completed", + "error": None, + "incomplete_details": None, + "instructions": None, + "max_output_tokens": None, + "model": "gpt-4o", + "output": [ + { + "id": f"msg_{response_id}", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": _OUTPUT_TEXT, "annotations": []}], + } + ], + "parallel_tool_calls": True, + "previous_response_id": None, + "reasoning": {"effort": None, "summary": None}, + "store": store, + "temperature": 1.0, + "text": {"format": {"type": "text"}}, + "tool_choice": "auto", + "tools": [], + "top_p": 1.0, + "truncation": "disabled", + "usage": { + "input_tokens": 2, + "input_tokens_details": {"cached_tokens": 0}, + "output_tokens": 3, + "output_tokens_details": {"reasoning_tokens": 0}, + "total_tokens": 5, + }, + "user": None, + "metadata": {}, + } + + +def _sse(events: Iterable[dict[str, object]]) -> str: + return "".join(f"data: {json.dumps(event)}\n\n" for event in events) + + +def _response_events(response_id: str) -> tuple[dict[str, object], ...]: + completed_response: Final = _response_body(response_id) + in_progress_response: Final = {**completed_response, "status": "in_progress", "output": [], "usage": None} + item_id: Final = f"msg_{response_id}" + text_part: Final = {"type": "output_text", "text": _OUTPUT_TEXT, "annotations": []} + position: Final = {"item_id": item_id, "output_index": 0, "content_index": 0} + return ( + {"type": "response.created", "sequence_number": 0, "response": in_progress_response}, + {"type": "response.in_progress", "sequence_number": 1, "response": in_progress_response}, + { + "type": "response.output_item.added", + "sequence_number": 2, + "output_index": 0, + "item": {"id": item_id, "type": "message", "status": "in_progress", "role": "assistant", "content": []}, + }, + { + "type": "response.content_part.added", + "sequence_number": 3, + **position, + "part": {"type": "output_text", "text": "", "annotations": []}, + }, + {"type": "response.output_text.delta", "sequence_number": 4, **position, "delta": _OUTPUT_TEXT}, + {"type": "response.output_text.done", "sequence_number": 5, **position, "text": _OUTPUT_TEXT}, + {"type": "response.content_part.done", "sequence_number": 6, **position, "part": text_part}, + { + "type": "response.output_item.done", + "sequence_number": 7, + "output_index": 0, + "item": completed_response["output"][0], + }, + {"type": "response.completed", "sequence_number": 8, "response": completed_response}, + ) + + +def _response_sse(response_id: str) -> str: + return _sse(_response_events(response_id)) + + +def _sse_reply(body: str) -> httpx.Response: + return httpx.Response(status_code=200, content=body, headers={"content-type": "text/event-stream"}) + + +def _assert_valid_response(response: object, final_chunk: bool) -> None: + assert isinstance(response, ResponsesAPIResponse) + assert isinstance(response.id, str) + assert isinstance(response.created_at, int) + assert isinstance(response.usage, ResponseAPIUsage if final_chunk else type(None)) + mistyped: Final = { + field: type(response[field]).__name__ + for field, expected in _RESPONSE_FIELD_TYPES.items() + if not isinstance(response[field], expected) + } + assert mistyped == {} + if final_chunk and response.status == "completed": + assert len(response.output) > 0 + + +def _json_object(value: object) -> dict[str, object]: + return _JSON_OBJECT.validate_python(value.model_dump(mode="json") if isinstance(value, BaseModel) else value) + + +def _assert_valid_stream(events: tuple[object, ...], item_id: str) -> None: + event_types: Final = tuple(getattr(event, "type", None) for event in events) + assert event_types == tuple(_STREAM_EVENT_FIELDS) + missing_fields: Final = { + event_type: tuple(name for name in fields if getattr(event, name, None) is None) + for (event_type, fields), event in zip(_STREAM_EVENT_FIELDS.items(), events) + } + assert all(fields == () for fields in missing_fields.values()), missing_fields + created: Final = getattr(events[0], "response", None) + _assert_valid_response(created, final_chunk=False) + _assert_valid_response(getattr(events[1], "response", None), final_chunk=False) + completed: Final = events[-1] + assert isinstance(completed, ResponseCompletedEvent) + _assert_valid_response(completed.response, final_chunk=True) + assert completed.response.id == getattr(created, "id", None) + assert completed.response.output_text == _OUTPUT_TEXT + assert tuple(getattr(event, "sequence_number", None) for event in events) == tuple(range(len(events))) + item_ids: Final = ( + getattr(getattr(events[2], "item", None), "id", None), + *(getattr(event, "item_id", None) for event in events[3:7]), + getattr(getattr(events[7], "item", None), "id", None), + ) + assert item_ids == (item_id,) * 6 + assert tuple(getattr(event, "output_index", None) for event in events[2:8]) == (0,) * 6 + assert tuple(getattr(event, "content_index", None) for event in events[3:7]) == (0,) * 4 + assert _json_object(getattr(events[3], "part", None)) == { + "type": "output_text", + "text": "", + "annotations": [], + } + assert getattr(events[4], "delta", None) == _OUTPUT_TEXT + assert getattr(events[5], "text", None) == _OUTPUT_TEXT + assert _json_object(getattr(events[6], "part", None)) == { + "type": "output_text", + "text": _OUTPUT_TEXT, + "annotations": [], + } + assert getattr(getattr(events[7], "item", None), "type", None) == "message" + assert tuple(item.id for item in completed.response.output) == (item_id,) + + +def _completed_response(events: tuple[object, ...]) -> ResponsesAPIResponse: + completed_events: Final = tuple(event for event in events if isinstance(event, ResponseCompletedEvent)) + assert len(completed_events) == 1 + return completed_events[0].response + + +def _install_capture(monkeypatch: pytest.MonkeyPatch) -> _ResponsesLoggingCapture: + capture: Final = _ResponsesLoggingCapture() + monkeypatch.setattr(litellm, "callbacks", [capture]) + monkeypatch.setattr(litellm, "success_callback", [capture]) + monkeypatch.setattr(litellm, "_async_success_callback", [capture]) + return capture + + +def _assert_logged_payload_matches(capture: _ResponsesLoggingCapture, response: ResponsesAPIResponse) -> None: + payload: Final = capture.payload + usage: Final = capture.usage + assert payload is not None + assert response.usage is not None + assert payload["prompt_tokens"] == response.usage.input_tokens + assert payload["completion_tokens"] == response.usage.output_tokens + assert payload["total_tokens"] == response.usage.input_tokens + response.usage.output_tokens + assert payload["response_cost"] > 0 + assert payload["id"] == response.id + assert payload["model"] == "gpt-4o-mini" + assert payload["messages"] == [{"content": "hi", "role": "user"}] + callback_usage: Final = usage.model_dump() if isinstance(usage, Usage) else _JSON_OBJECT.validate_python(usage) + assert callback_usage["prompt_tokens"] == response.usage.input_tokens + assert callback_usage["completion_tokens"] == response.usage.output_tokens + logged_response: Final = _JSON_OBJECT.validate_python(payload["response"]) + final_response: Final = response.model_dump(mode="json") + assert {key: value for key, value in logged_response.items() if key != "usage"} == { + key: value for key, value in final_response.items() if key != "usage" + } + logged_usage: Final = _JSON_OBJECT.validate_python(logged_response["usage"]) + assert logged_usage["prompt_tokens"] == response.usage.input_tokens + assert logged_usage["completion_tokens"] == response.usage.output_tokens + assert logged_usage["total_tokens"] == response.usage.total_tokens + + +async def _async_logged_openai_response(mode: _AsyncLoggingMode) -> ResponsesAPIResponse: + match mode: + case "stream": + response_stream: Final = await litellm.aresponses( + model="openai/gpt-4o-mini", api_key="sk-test", input="hi", stream=True + ) + return _completed_response(tuple([event async for event in response_stream])) + case "non_stream": + response: Final = await litellm.aresponses(model="openai/gpt-4o-mini", api_key="sk-test", input="hi") + assert isinstance(response, ResponsesAPIResponse) + return response + + +@pytest.mark.parametrize("sync_mode", (True, False)) +@pytest.mark.asyncio +async def test_responses_exposes_provider_rate_limit_headers(sync_mode: bool) -> None: + response_headers: Final = { + "x-ratelimit-limit-requests": "500", + "x-ratelimit-remaining-requests": "499", + "x-ratelimit-reset-requests": "1s", + } + + with respx.mock() as mock_router: + mock_router.post(_OPENAI_URL).mock( + return_value=httpx.Response( + status_code=200, + json=_response_body("resp_rate_limit"), + headers=response_headers, + ) + ) + response: Final = ( + litellm.responses(model="openai/gpt-4o", api_key="sk-test", input="hi") + if sync_mode + else await litellm.aresponses(model="openai/gpt-4o", api_key="sk-test", input="hi") + ) + requests: Final = tuple(mock_router.calls) + + assert isinstance(response, ResponsesAPIResponse) + assert len(requests) == 1 + additional_headers: Final = _JSON_OBJECT.validate_python(response.hidden_params["additional_headers"]) + raw_headers: Final = _JSON_OBJECT.validate_python(response.hidden_params["headers"]) + assert {name: additional_headers[f"llm_provider-{name}"] for name in response_headers} == response_headers + assert {name: raw_headers[name] for name in response_headers} == response_headers + + +@pytest.mark.asyncio +async def test_responses_converts_created_at_to_int_and_forwards_store() -> None: + response_body: Final = _response_body("resp_store", store=True, created_at=_CREATED_AT + 0.75) + + with respx.mock() as mock_router: + mock_router.post(_OPENAI_URL).mock(return_value=httpx.Response(status_code=200, json=response_body)) + response: Final = await litellm.aresponses( + model="openai/gpt-4o", + api_key="sk-test", + input="hi", + store=True, + ) + requests: Final = tuple(mock_router.calls) + + assert isinstance(response, ResponsesAPIResponse) + assert type(response.created_at) is int + assert response.created_at == _CREATED_AT + assert response.store is True + assert len(requests) == 1 + request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content) + assert request_body["store"] is True + + +@pytest.mark.asyncio +async def test_responses_mcp_followup_forwards_approval_and_previous_id() -> None: + mcp_tools: Final = [ + { + "type": "mcp", + "server_label": "weather", + "server_url": "https://mcp.example.test", + "headers": {"Authorization": "Bearer mcp-test-token"}, + } + ] + approval: Final = [{"type": "mcp_approval_response", "approve": True, "approval_request_id": "approval_123"}] + + with respx.mock() as mock_router: + route: Final = mock_router.post(_OPENAI_URL).mock( + return_value=httpx.Response(status_code=200, json=_response_body("resp_mcp")) + ) + first_response: Final = await litellm.aresponses( + model="openai/gpt-4o", + api_key="sk-test", + input="Search for recent weather", + tools=mcp_tools, + ) + assert isinstance(first_response, ResponsesAPIResponse) + second_response: Final = await litellm.aresponses( + model="openai/gpt-4o", + api_key="sk-test", + input=approval, + tools=mcp_tools, + previous_response_id=first_response.id, + ) + requests: Final = tuple(route.calls) + + assert isinstance(second_response, ResponsesAPIResponse) + assert len(requests) == 2 + first_request: Final = _JSON_OBJECT.validate_json(requests[0].request.content) + second_request: Final = _JSON_OBJECT.validate_json(requests[1].request.content) + assert first_request["tools"] == mcp_tools + assert second_request["tools"] == mcp_tools + assert second_request["input"] == approval + assert second_request["previous_response_id"] == "resp_mcp" + + +@pytest.mark.parametrize("mode", ("non_stream", "stream")) +@pytest.mark.asyncio +async def test_responses_standard_logging_matches_final_response( + mode: _AsyncLoggingMode, monkeypatch: pytest.MonkeyPatch +) -> None: + capture: Final = _install_capture(monkeypatch) + + with respx.mock() as mock_router: + mock_router.post(_OPENAI_URL).mock( + return_value=( + _sse_reply(_response_sse("resp_logging")) + if mode == "stream" + else httpx.Response(status_code=200, json=_response_body("resp_logging")) + ) + ) + response: Final = await _async_logged_openai_response(mode) + requests: Final = tuple(mock_router.calls) + logged: Final = await asyncio.to_thread(capture.completed.wait, 10) + + assert logged + assert len(requests) == 1 + _assert_logged_payload_matches(capture, response) + + +def test_sync_stream_standard_logging_matches_final_response(monkeypatch: pytest.MonkeyPatch) -> None: + capture: Final = _install_capture(monkeypatch) + + with respx.mock() as mock_router: + mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_response_sse("resp_sync_logging"))) + stream: Final = litellm.responses(model="openai/gpt-4o-mini", api_key="sk-test", input="hi", stream=True) + response: Final = _completed_response(tuple(stream)) + requests: Final = tuple(mock_router.calls) + logged: Final = capture.completed.wait(10) + + assert logged + assert len(requests) == 1 + _assert_logged_payload_matches(capture, response) + + +@pytest.mark.parametrize("sync_mode", (True, False)) +@pytest.mark.asyncio +async def test_responses_stream_emits_valid_events(sync_mode: bool) -> None: + with respx.mock() as mock_router: + route: Final = mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_response_sse("resp_stream"))) + events: Final = ( + tuple(litellm.responses(model="openai/gpt-4o", api_key="sk-test", input="hi", stream=True)) + if sync_mode + else tuple( + [ + event + async for event in await litellm.aresponses( + model="openai/gpt-4o", api_key="sk-test", input="hi", stream=True + ) + ] + ) + ) + requests: Final = tuple(route.calls) + + assert len(requests) == 1 + _assert_valid_stream(events, "msg_resp_stream") + + +@pytest.mark.asyncio +async def test_responses_stream_error_event_raises_bad_request() -> None: + message: Final = "Your input exceeds the context window of this model." + created: Final = _response_events("resp_too_long")[0] + error_event: Final = { + "type": "error", + "sequence_number": 1, + "error": { + "type": "invalid_request_error", + "code": "context_length_exceeded", + "message": message, + "param": "input", + }, + } + + with respx.mock() as mock_router: + route: Final = mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_sse((created, error_event)))) + stream: Final = await litellm.aresponses( + model="openai/gpt-5-mini", api_key="sk-test", input="oversized prompt", stream=True + ) + with pytest.raises(litellm.BadRequestError) as exc_info: + _ = [event async for event in stream] + requests: Final = tuple(route.calls) + + assert len(requests) == 1 + assert exc_info.value.status_code == 400 + assert message in str(exc_info.value) + + +def _alias_router() -> litellm.Router: + return litellm.Router( + model_list=[ + { + "model_name": "openai-offline-alias", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"}, + } + ] + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sync_mode", (True, False)) +async def test_router_responses_alias_uses_underlying_model(sync_mode: bool) -> None: + router: Final = _alias_router() + with respx.mock() as mock_router: + route: Final = mock_router.post(_OPENAI_URL).mock( + return_value=httpx.Response(status_code=200, json=_response_body("resp_router_alias")) + ) + response: Final = ( + router.responses(model="openai-offline-alias", input="hi") + if sync_mode + else await router.aresponses(model="openai-offline-alias", input="hi") + ) + requests: Final = tuple(route.calls) + + _assert_valid_response(response, final_chunk=True) + assert len(requests) == 1 + request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content) + assert request_body["model"] == "gpt-4o" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sync_mode", (True, False)) +async def test_router_responses_alias_stream_uses_underlying_model(sync_mode: bool) -> None: + router: Final = _alias_router() + with respx.mock() as mock_router: + route: Final = mock_router.post(_OPENAI_URL).mock(return_value=_sse_reply(_response_sse("resp_router_stream"))) + events: Final = ( + tuple(router.responses(model="openai-offline-alias", input="hi", stream=True)) + if sync_mode + else tuple( + [ + event + async for event in await router.aresponses(model="openai-offline-alias", input="hi", stream=True) + ] + ) + ) + requests: Final = tuple(route.calls) + + _assert_valid_stream(events, "msg_resp_router_stream") + assert len(requests) == 1 + request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content) + assert request_body["model"] == "gpt-4o" diff --git a/tests/unit/llms/openai/responses/test_openai_responses_transformation.py b/tests/unit/llms/openai/responses/test_openai_responses_transformation.py index a8f0d23e560..73332728964 100644 --- a/tests/unit/llms/openai/responses/test_openai_responses_transformation.py +++ b/tests/unit/llms/openai/responses/test_openai_responses_transformation.py @@ -6,6 +6,8 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch import httpx import pytest +import respx +from pydantic import JsonValue, TypeAdapter import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -24,6 +26,7 @@ from litellm.types.llms.openai import ( ) from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import Choices, Message, ModelResponse +from tests.unit.proxy.conftest import httpx_transport import time _ARTIFACT_FIELD_PATTERN: Final = r'^(?!__.*__$)[^\p{Cc}\p{Cf}\p{Zl}\p{Zp}"\\./[\]]{1,200}$' @@ -3445,3 +3448,127 @@ async def test_extra_body_merges_with_request_data(extra_body_mock_response_data assert "temperature" in request_body assert "custom_field" in request_body assert request_body["custom_field"] == "custom_value" + + +@pytest.mark.asyncio +@pytest.mark.usefixtures(httpx_transport.__name__) +async def test_aresponses_forwards_previous_response_id_to_openai() -> None: + first_input: Final = "remember the first turn" + second_input: Final = "continue the conversation" + first_id: Final = "resp_previous_turn" + second_id: Final = "resp_follow_up" + response_payloads: Final = ( + { + "id": first_id, + "object": "response", + "created_at": 1734366691, + "status": "completed", + "model": "gpt-4o", + "output": [ + { + "type": "message", + "id": "msg_first", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "first answer", "annotations": []}], + } + ], + "usage": { + "input_tokens": 1, + "output_tokens": 1, + "total_tokens": 2, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + }, + { + "id": second_id, + "object": "response", + "created_at": 1734366692, + "status": "completed", + "model": "gpt-4o", + "output": [ + { + "type": "message", + "id": "msg_second", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "second answer", "annotations": []}], + } + ], + "usage": { + "input_tokens": 1, + "output_tokens": 1, + "total_tokens": 2, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + }, + ) + provider_url: Final = "https://api.openai.com/v1/responses" + + with respx.mock() as router: + route: Final = router.post(provider_url).mock( + side_effect=[ + httpx.Response(status_code=200, json=response_payloads[0]), + httpx.Response(status_code=200, json=response_payloads[1]), + ] + ) + first_response: Final = await litellm.aresponses( + model="openai/gpt-4o", + api_key="sk-test", + input=first_input, + ) + assert isinstance(first_response, ResponsesAPIResponse) + second_response: Final = await litellm.aresponses( + model="openai/gpt-4o", + api_key="sk-test", + input=second_input, + previous_response_id=first_response.id, + ) + calls: Final = route.calls + + assert isinstance(second_response, ResponsesAPIResponse) + assert first_response.output[0].content[0].text == "first answer" + assert second_response.output[0].content[0].text == "second answer" + assert len(calls) == 2 + request_adapter: Final = TypeAdapter(dict[str, JsonValue]) + request_bodies: Final = tuple(request_adapter.validate_json(call.request.content) for call in calls) + assert request_bodies[0]["input"] == first_input + assert request_bodies[1]["input"] == second_input + assert request_bodies[1]["previous_response_id"] == first_id + + +def test_dict_responses_input_filters_unset_reasoning_fields() -> None: + test_input: Final = [ + {"role": "user", "content": "test"}, + { + "id": "rs_123", + "summary": [{"text": "test", "type": "summary_text"}], + "type": "reasoning", + "content": None, + "encrypted_content": None, + "status": None, + }, + { + "arguments": "{}", + "call_id": "call_123", + "name": "get_today", + "type": "function_call", + "id": "fc_123", + "status": "completed", + }, + ] + + validated_input: Final = OpenAIResponsesAPIConfig()._validate_input_param(test_input) + + assert len(validated_input) == 3 + reasoning_item: Final = validated_input[1] + assert reasoning_item["type"] == "reasoning" + assert "status" not in reasoning_item + assert "content" not in reasoning_item + assert "encrypted_content" not in reasoning_item + assert reasoning_item["id"] == "rs_123" + assert reasoning_item["summary"] == [{"text": "test", "type": "summary_text"}] + + function_call_item: Final = validated_input[2] + assert function_call_item["type"] == "function_call" + assert function_call_item["status"] == "completed" diff --git a/tests/unit/llms/pass_through/guardrail_translation/test_handler.py b/tests/unit/llms/pass_through/guardrail_translation/test_handler.py new file mode 100644 index 00000000000..82ad7702052 --- /dev/null +++ b/tests/unit/llms/pass_through/guardrail_translation/test_handler.py @@ -0,0 +1,52 @@ +import json +from typing import Final + +import pytest + +from litellm.constants import SERVER_STREAMING_CLASSIFICATION_KEY, SERVER_STREAMING_CLASSIFICATION_MARKER +from litellm.llms.pass_through.guardrail_translation.handler import PassThroughEndpointHandler + + +def test_full_payload_guardrail_text_excludes_the_server_streaming_marker(): + body: Final = {"model": "m", "stream": True, "messages": [{"role": "user", "content": "hi"}]} + + text: Final = PassThroughEndpointHandler()._extract_text_for_guardrail( + {**body, SERVER_STREAMING_CLASSIFICATION_KEY: SERVER_STREAMING_CLASSIFICATION_MARKER}, + None, + ) + + assert json.loads(text) == body, text + + +@pytest.mark.parametrize( + "marker", + [ + SERVER_STREAMING_CLASSIFICATION_MARKER, + json.loads(json.dumps(SERVER_STREAMING_CLASSIFICATION_MARKER)), + ], + ids=["enum", "json-string"], +) +def test_full_payload_guardrail_text_scans_caller_value_but_not_the_marker(marker: str): + text: Final = PassThroughEndpointHandler()._extract_text_for_guardrail( + { + "model": "m", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + SERVER_STREAMING_CLASSIFICATION_KEY: "BLOCKME caller content", + }, + None, + ) + + assert "BLOCKME caller content" in text, text + + marker_text: Final = PassThroughEndpointHandler()._extract_text_for_guardrail( + { + "model": "m", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + SERVER_STREAMING_CLASSIFICATION_KEY: marker, + }, + None, + ) + + assert SERVER_STREAMING_CLASSIFICATION_KEY not in json.loads(marker_text), marker_text diff --git a/tests/unit/llms/perplexity/search/__init__.py b/tests/unit/llms/perplexity/search/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/perplexity/search/test_perplexity_search_transformation.py b/tests/unit/llms/perplexity/search/test_perplexity_search_transformation.py new file mode 100644 index 00000000000..432ed8248e6 --- /dev/null +++ b/tests/unit/llms/perplexity/search/test_perplexity_search_transformation.py @@ -0,0 +1,57 @@ +import json +from typing import Final + +import pytest +import respx + +import litellm + +PERPLEXITY_SEARCH_URL: Final = "https://api.perplexity.ai/search" + + +@pytest.fixture(autouse=True) +def perplexity_api_key(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("PERPLEXITYAI_API_KEY", "test-perplexity-key") + monkeypatch.delenv("PERPLEXITY_API_BASE", raising=False) + + +def test_search_maps_perplexity_results_to_search_response(respx_mock: respx.MockRouter) -> None: + respx_mock.post(PERPLEXITY_SEARCH_URL).respond( + json={ + "results": [ + { + "title": "AI news roundup", + "url": "https://example.com/ai-news", + "snippet": "The latest in artificial intelligence.", + "date": "2026-01-15", + "last_updated": "2026-01-16", + }, + {"title": "Second", "url": "https://example.com/second", "snippet": "Second snippet."}, + ] + } + ) + + response: Final = litellm.search(query="artificial intelligence recent news", search_provider="perplexity") + + assert response.object == "search" + assert isinstance(response.results, list) + assert len(response.results) == 2 + first: Final = response.results[0] + assert first.title == "AI news roundup" + assert first.url == "https://example.com/ai-news" + assert first.snippet == "The latest in artificial intelligence." + assert first.date == "2026-01-15" + assert first.last_updated == "2026-01-16" + assert response.results[1].snippet == "Second snippet." + + +def test_search_sends_max_results_in_request_body(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.post(PERPLEXITY_SEARCH_URL).respond( + json={"results": [{"title": "ML", "url": "https://example.com/ml", "snippet": "Machine learning."}]} + ) + + response: Final = litellm.search(query="machine learning", search_provider="perplexity", max_results=5) + + assert json.loads(route.calls.last.request.content) == {"query": "machine learning", "max_results": 5} + assert route.calls.last.request.headers["Authorization"] == "Bearer test-perplexity-key" + assert [result.url for result in response.results] == ["https://example.com/ml"] diff --git a/tests/unit/llms/sagemaker/test_sagemaker_chat_handler.py b/tests/unit/llms/sagemaker/test_sagemaker_chat_handler.py index 598d8177e6e..4ae117b5ebf 100644 --- a/tests/unit/llms/sagemaker/test_sagemaker_chat_handler.py +++ b/tests/unit/llms/sagemaker/test_sagemaker_chat_handler.py @@ -163,3 +163,23 @@ async def test_completion_sagemaker_messages_api(sync_mode): assert json_data["max_tokens"] == 80 except Exception as e: pytest.fail(f"Error occurred: {e}") + + +def test_missing_botocore_keeps_dependency_identity(): + import pytest + + with patch.dict("sys.modules", {"botocore": None}): + with pytest.raises(ModuleNotFoundError, match="pip install boto3") as caught: + SagemakerChatHandler()._load_credentials({}) + assert caught.value.name == "botocore" + + +def test_installed_botocore_signs_the_chat_request(): + from botocore.credentials import Credentials + + request = SagemakerChatHandler()._prepare_request( + credentials=Credentials("test-key", "test-secret"), model="test-endpoint", data={"inputs": "ping"}, + optional_params={}, aws_region_name="us-west-2", + ) + assert request.body == b'{"inputs": "ping"}' + assert "/us-west-2/sagemaker/aws4_request" in request.headers["Authorization"] diff --git a/tests/unit/llms/sagemaker/test_sagemaker_completion_handler.py b/tests/unit/llms/sagemaker/test_sagemaker_completion_handler.py index 9a5f36e2081..339e02c20d5 100644 --- a/tests/unit/llms/sagemaker/test_sagemaker_completion_handler.py +++ b/tests/unit/llms/sagemaker/test_sagemaker_completion_handler.py @@ -267,3 +267,29 @@ def test_load_credentials_assumes_role_with_session_tags(monkeypatch): assert credentials.access_key == "ASIASMCOMPTAGGED" assert aws_region_name == "us-east-1" assert "aws_session_tags" not in optional_params + + +def test_missing_botocore_keeps_dependency_identity(): + from unittest.mock import patch + + import pytest + + from litellm.llms.sagemaker.completion.handler import SagemakerLLM + + with patch.dict("sys.modules", {"botocore": None}): + with pytest.raises(ModuleNotFoundError, match="pip install boto3") as caught: + SagemakerLLM()._load_credentials({}) + assert caught.value.name == "botocore" + + +def test_installed_botocore_signs_the_completion_request(): + from botocore.credentials import Credentials + + from litellm.llms.sagemaker.completion.handler import SagemakerLLM + + request = SagemakerLLM()._prepare_request( + credentials=Credentials("test-key", "test-secret"), model="test-endpoint", data={"inputs": "ping"}, + messages=[], litellm_params={}, optional_params={}, aws_region_name="us-west-2", + ) + assert request.body == b'{"inputs": "ping"}' + assert "/us-west-2/sagemaker/aws4_request" in request.headers["Authorization"] diff --git a/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py b/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py index 6f18a391f7b..dde3df6fcc0 100644 --- a/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py +++ b/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py @@ -1289,6 +1289,33 @@ class TestVertexEmbeddingsBatchInputTranslation: assert set(row["request"]) == {"content"} + def test_should_translate_a_file_block_with_video_metadata(self) -> None: + (row,) = _wrap_entries( + [ + _embeddings_entry( + body={ + "model": "gemini-embedding-2", + "input": [ + { + "type": "file", + "file": { + "file_id": "gs://my-bucket/clip.mp4", + "video_metadata": {"start_offset": "3s", "end_offset": "6s"}, + }, + } + ], + } + ) + ] + ) + + assert row["request"]["content"]["parts"] == [ + { + "file_data": {"mime_type": "video/mp4", "file_uri": "gs://my-bucket/clip.mp4"}, + "video_metadata": {"startOffset": "3s", "endOffset": "6s"}, + } + ] + def test_should_translate_multimodal_gcs_input(self): (row,) = _wrap_entries( [ diff --git a/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_handler.py b/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_handler.py new file mode 100644 index 00000000000..483b350fd43 --- /dev/null +++ b/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_handler.py @@ -0,0 +1,39 @@ +from typing import Final + +import httpx +import pytest + +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_handler import GoogleBatchEmbeddings + +FILES_URI: Final = "https://generativelanguage.googleapis.com/v1beta/files/clip123" +FILE_METADATA_URL: Final = "https://generativelanguage.googleapis.com/v1beta/files/clip123" + + +def _files_api_metadata(request: httpx.Request) -> httpx.Response: + if str(request.url) != FILE_METADATA_URL or request.headers.get("x-goog-api-key") != "gemini-key": + return httpx.Response(404, json={"error": {"message": "no such file"}}) + return httpx.Response(200, json={"mimeType": "video/mp4", "uri": FILES_URI}) + + +@pytest.mark.parametrize("reference", ["files/clip123", FILES_URI]) +def test_resolve_file_references_fetches_the_file_metadata_for_both_reference_forms(reference: str) -> None: + resolved: Final = GoogleBatchEmbeddings()._resolve_file_references( + input=[reference, "a red bus"], + api_key="gemini-key", + sync_handler=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_files_api_metadata))), + ) + assert resolved == {reference: {"mime_type": "video/mp4", "uri": FILES_URI}} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("reference", ["files/clip123", FILES_URI]) +async def test_async_resolve_file_references_fetches_the_file_metadata_for_both_reference_forms( + reference: str, +) -> None: + resolved: Final = await GoogleBatchEmbeddings()._async_resolve_file_references( + input=[reference, "a red bus"], + api_key="gemini-key", + async_handler=AsyncHTTPHandler(transport=httpx.MockTransport(_files_api_metadata)), + ) + assert resolved == {reference: {"mime_type": "video/mp4", "uri": FILES_URI}} diff --git a/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py b/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py index ba2b26bf0a2..6584c27b895 100644 --- a/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py +++ b/tests/unit/llms/vertex_ai/gemini_embeddings/test_batch_embed_content_transformation.py @@ -8,12 +8,17 @@ Covers: - Response processing with correct indices """ +from collections.abc import Callable +from typing import Final + import pytest import litellm +from litellm.exceptions import BadRequestError from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import ( _build_part_for_input, + file_reference_name, _is_multimodal_input, process_embed_content_response, process_response, @@ -25,6 +30,12 @@ from litellm.types.utils import EmbeddingResponse IMAGE_DATA_URI = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII" GCS_URL = "gs://my-bucket/image.png" +VIDEO_DATA_URI = "data:video/mp4;base64,AAAAIGZ0eXBpc29tAAACAGlzb21pc28yYXZjMW1wNDEAAAAIZnJlZQAA" +FILES_URI = "https://generativelanguage.googleapis.com/v1beta/files/clip123" + + +def _file_block(**file: object) -> dict[str, object]: + return {"type": "file", "file": file} @pytest.fixture(autouse=True) @@ -51,6 +62,7 @@ class TestIsMultimodalInput: def test_file_reference(self): assert _is_multimodal_input(["files/abc123"]) is True + assert _is_multimodal_input([FILES_URI]) is True def test_mixed_text_and_image(self): assert _is_multimodal_input(["hello", IMAGE_DATA_URI]) is True @@ -59,9 +71,16 @@ class TestIsMultimodalInput: """Nested list with text is not multimodal.""" assert _is_multimodal_input([["text_a", "text_b"]]) is False + def test_file_content_block_is_multimodal(self) -> None: + assert _is_multimodal_input([_file_block(file_data=VIDEO_DATA_URI)]) is True + def test_nested_list_with_image_is_multimodal(self): assert _is_multimodal_input([["a red shoe", IMAGE_DATA_URI]]) is True + def test_single_object_input_answers_400(self) -> None: + with pytest.raises(BadRequestError, match="string or a list"): + _is_multimodal_input(_file_block(file_data=VIDEO_DATA_URI)) + class TestBuildPartForInput: def test_text_input(self): @@ -83,15 +102,13 @@ class TestBuildPartForInput: assert part["file_data"]["file_uri"] == GCS_URL def test_file_reference_resolved(self): - resolved = { - "files/abc": {"mime_type": "image/jpeg", "uri": "https://example.com/abc"} - } + resolved = {"files/abc": {"mime_type": "image/jpeg", "uri": "https://example.com/abc"}} part = _build_part_for_input("files/abc", resolved_files=resolved) assert part["file_data"] is not None assert part["file_data"]["mime_type"] == "image/jpeg" - def test_file_reference_unresolved_raises(self): - with pytest.raises(ValueError, match="not resolved"): + def test_file_reference_unresolved_answers_400_naming_the_gemini_provider(self) -> None: + with pytest.raises(BadRequestError, match="gemini/ provider"): _build_part_for_input("files/abc") @@ -124,10 +141,7 @@ class TestTransformOpenaiInputGeminiContent: ) assert len(result["requests"]) == 2 # First request is text - assert ( - result["requests"][0]["content"]["parts"][0]["text"] - == "The food was delicious" - ) + assert result["requests"][0]["content"]["parts"][0]["text"] == "The food was delicious" # Second request is image assert result["requests"][1]["content"]["parts"][0]["inline_data"] is not None @@ -226,9 +240,7 @@ class TestProcessResponse: """Test that process_response sets correct indices.""" def test_single_embedding_index(self): - predictions: VertexAIBatchEmbeddingsResponseObject = { - "embeddings": [{"values": [0.1, 0.2]}] - } + predictions: VertexAIBatchEmbeddingsResponseObject = {"embeddings": [{"values": [0.1, 0.2]}]} model_response = EmbeddingResponse() result = process_response( input="hello", @@ -279,9 +291,7 @@ class TestProcessResponse: def test_nested_input_token_counting(self): """Nested list: only plain-text sub-elements should be counted.""" - predictions: VertexAIBatchEmbeddingsResponseObject = { - "embeddings": [{"values": [0.1, 0.2]}] - } + predictions: VertexAIBatchEmbeddingsResponseObject = {"embeddings": [{"values": [0.1, 0.2]}]} result = process_response( input=[["a red shoe", IMAGE_DATA_URI]], model_response=EmbeddingResponse(), @@ -300,7 +310,7 @@ class TestProcessResponse: ) def test_nested_non_string_element_raises(self): - with pytest.raises(ValueError, match="must be strings"): + with pytest.raises(BadRequestError, match="must be strings or file content blocks, got list"): transform_openai_input_gemini_content( input=[[["doubly", "nested"]]], model="gemini-embedding-2-preview", @@ -408,3 +418,212 @@ class TestProcessEmbedContentResponseUsage: assert result.usage.prompt_tokens > 0 +class TestFileContentBlocks: + MODEL: Final = "gemini-embedding-2-preview" + CLIP_METADATA: Final = {"fps": 2, "start_offset": "3s", "end_offset": "6s"} + CLIP_PART: Final = {"fps": 2.0, "startOffset": "3s", "endOffset": "6s"} + + def test_data_uri_block_forwards_format_and_video_metadata(self) -> None: + part: Final = _build_part_for_input( + _file_block( + file_data="data:application/octet-stream;base64,QUJD", + format="video/mp4", + video_metadata=self.CLIP_METADATA, + ) + ) + assert part == { + "inline_data": {"mime_type": "video/mp4", "data": "QUJD"}, + "video_metadata": self.CLIP_PART, + } + + def test_gcs_block_infers_mime_type_and_forwards_video_metadata(self) -> None: + part: Final = _build_part_for_input( + _file_block(file_id="gs://my-bucket/clip.mp4", video_metadata={"start_offset": "0s", "end_offset": "3s"}) + ) + assert part == { + "file_data": {"mime_type": "video/mp4", "file_uri": "gs://my-bucket/clip.mp4"}, + "video_metadata": {"startOffset": "0s", "endOffset": "3s"}, + } + + def test_file_reference_block_uses_the_resolved_file(self) -> None: + resolved_files: Final = {"files/clip123": {"mime_type": "video/mp4", "uri": FILES_URI}} + part: Final = _build_part_for_input( + _file_block(file_id="files/clip123", video_metadata={"fps": 1}), + resolved_files=resolved_files, + ) + assert part == { + "file_data": {"mime_type": "video/mp4", "file_uri": FILES_URI}, + "video_metadata": {"fps": 1.0}, + } + + def test_files_api_uri_block_uses_the_resolved_file(self) -> None: + resolved_files: Final = {FILES_URI: {"mime_type": "video/mp4", "uri": FILES_URI}} + part: Final = _build_part_for_input( + _file_block(file_id=FILES_URI, video_metadata={"start_offset": "0s", "end_offset": "3s"}), + resolved_files=resolved_files, + ) + assert part == { + "file_data": {"mime_type": "video/mp4", "file_uri": FILES_URI}, + "video_metadata": {"startOffset": "0s", "endOffset": "3s"}, + } + + @pytest.mark.parametrize("reference", ["files/clip123", FILES_URI]) + def test_file_reference_name_is_the_files_name_for_both_forms(self, reference: str) -> None: + assert file_reference_name(reference) == "files/clip123" + + def test_block_without_video_metadata_sends_no_video_metadata_key(self) -> None: + part: Final = _build_part_for_input(_file_block(file_data=IMAGE_DATA_URI, filename="dot.png")) + assert part == {"inline_data": {"mime_type": "image/png", "data": IMAGE_DATA_URI.split(",", 1)[1]}} + + def test_block_with_empty_video_metadata_sends_no_video_metadata_key(self) -> None: + part: Final = _build_part_for_input(_file_block(file_data=VIDEO_DATA_URI, video_metadata={})) + assert part == {"inline_data": {"mime_type": "video/mp4", "data": VIDEO_DATA_URI.split(",", 1)[1]}} + + def test_batch_path_nested_block_and_text_share_one_request(self) -> None: + result: Final = transform_openai_input_gemini_content( + input=[[_file_block(file_data=VIDEO_DATA_URI, video_metadata=self.CLIP_METADATA), "a solid color clip"]], + model=self.MODEL, + optional_params={"dimensions": 768}, + ) + [request] = result["requests"] + assert request["outputDimensionality"] == 768 + assert request["content"]["parts"] == [ + { + "inline_data": {"mime_type": "video/mp4", "data": VIDEO_DATA_URI.split(",", 1)[1]}, + "video_metadata": self.CLIP_PART, + }, + {"text": "a solid color clip"}, + ] + + def test_batch_path_flat_block_and_text_are_separate_requests(self) -> None: + result: Final = transform_openai_input_gemini_content( + input=[_file_block(file_data=VIDEO_DATA_URI, video_metadata=self.CLIP_METADATA), "a solid color clip"], + model=self.MODEL, + optional_params={}, + ) + assert len(result["requests"]) == 2 + assert result["requests"][0]["content"]["parts"][0]["video_metadata"] == self.CLIP_PART + assert result["requests"][1]["content"]["parts"] == [{"text": "a solid color clip"}] + + def test_embed_content_path_accepts_flat_block_and_text(self) -> None: + result: Final = transform_openai_input_gemini_embed_content( + input=[_file_block(file_data=VIDEO_DATA_URI, video_metadata=self.CLIP_METADATA), "a solid color clip"], + model=self.MODEL, + optional_params={}, + ) + parts: Final = result["content"]["parts"] + assert parts[0]["video_metadata"] == self.CLIP_PART + assert parts[1] == {"text": "a solid color clip"} + + def test_embed_content_path_still_rejects_nested_lists(self) -> None: + with pytest.raises(ValueError, match="Nested"): + transform_openai_input_gemini_embed_content( + input=[[_file_block(file_data=VIDEO_DATA_URI), "a solid color clip"]], + model=self.MODEL, + optional_params={}, + ) + + @pytest.mark.parametrize( + "block, named_in_error", + [ + ( + _file_block(file_data=VIDEO_DATA_URI, video_metadata={"fps": 1, "startOffset": "1s"}), + "video_metadata.startOffset", + ), + (_file_block(file_data=VIDEO_DATA_URI, video_metadata={"fps": "fast"}), "video_metadata.fps"), + (_file_block(file_data=VIDEO_DATA_URI, video_metadata={"fps": "1"}), "video_metadata.fps"), + (_file_block(file_data=VIDEO_DATA_URI, video_metadata={"fps": True}), "video_metadata.fps"), + (_file_block(file_data=VIDEO_DATA_URI, video_metadata={"start_offset": 5}), "video_metadata.start_offset"), + (_file_block(file_data=VIDEO_DATA_URI, detail="high"), "file.detail"), + (_file_block(file_data=VIDEO_DATA_URI, format=""), "file.format"), + (_file_block(file_id="gs://my-bucket/clip.mp4", file_data=VIDEO_DATA_URI), "not both"), + (_file_block(), "needs file.file_id or file.file_data"), + ( + _file_block(file_id="https://example.com/clip.mp4"), + "a data: URI, a gs:// URL, a files/ reference, or a Gemini Files API URI", + ), + ({"type": "image_url", "image_url": {"url": IMAGE_DATA_URI}}, "Input should be 'file'"), + ], + ) + def test_malformed_block_answers_400_naming_the_field( + self, block: dict[str, object], named_in_error: str + ) -> None: + with pytest.raises(BadRequestError, match=named_in_error): + _build_part_for_input(block) + + def test_drop_params_drops_the_block_keys_this_surface_does_not_take(self) -> None: + block: Final = _file_block( + file_data=VIDEO_DATA_URI, + detail="high", + video_metadata={"fps": 1, "start_offset": "1s", "resolution": "low"}, + ) + part: Final = _build_part_for_input({**block, "cache_control": {"type": "ephemeral"}}, drop_params=True) + assert part["inline_data"]["mime_type"] == "video/mp4" + assert part["video_metadata"] == {"fps": 1.0, "startOffset": "1s"} + + def test_drop_params_still_answers_400_for_a_malformed_value(self) -> None: + block: Final = _file_block(file_data=VIDEO_DATA_URI, detail="high", video_metadata={"fps": "fast"}) + with pytest.raises(BadRequestError, match=r"video_metadata\.fps"): + _build_part_for_input(block, drop_params=True) + + @pytest.mark.parametrize( + "transform", [transform_openai_input_gemini_content, transform_openai_input_gemini_embed_content] + ) + def test_transforms_forward_drop_params_to_every_block(self, transform: Callable[..., object]) -> None: + block: Final = _file_block(file_data=VIDEO_DATA_URI, detail="high") + with pytest.raises(BadRequestError, match=r"file\.detail"): + transform(input=[block], model="gemini-embedding-2-preview", optional_params={}) + transform(input=[block], model="gemini-embedding-2-preview", optional_params={}, drop_params=True) + + def test_batch_path_forwards_drop_params_into_nested_lists(self) -> None: + block: Final = _file_block(file_data=VIDEO_DATA_URI, detail="high") + body: Final = transform_openai_input_gemini_content( + input=[[block, "a caption"]], model="gemini-embedding-2-preview", optional_params={}, drop_params=True + ) + assert len(body["requests"][0]["content"]["parts"]) == 2 + + @pytest.mark.parametrize( + "transform", [transform_openai_input_gemini_content, transform_openai_input_gemini_embed_content] + ) + def test_single_object_input_answers_400(self, transform: Callable[..., object]) -> None: + with pytest.raises(BadRequestError, match="string or a list"): + transform( + input=_file_block(file_data=VIDEO_DATA_URI), model="gemini-embedding-2-preview", optional_params={} + ) + + def test_process_response_counts_only_the_text_tokens_next_to_a_block(self) -> None: + text: Final = "a solid color clip" + with_block: Final = process_response( + input=[_file_block(file_data=VIDEO_DATA_URI, video_metadata=self.CLIP_METADATA), text], + model_response=EmbeddingResponse(), + model=self.MODEL, + _predictions={"embeddings": [{"values": [0.1]}, {"values": [0.2]}]}, + ) + text_only: Final = process_response( + input=[text], + model_response=EmbeddingResponse(), + model=self.MODEL, + _predictions={"embeddings": [{"values": [0.2]}]}, + ) + assert with_block.usage.prompt_tokens == text_only.usage.prompt_tokens > 0 + + def test_embed_content_usage_fallback_with_a_block_does_not_estimate(self) -> None: + result: Final = process_embed_content_response( + input=[_file_block(file_data=VIDEO_DATA_URI, video_metadata=self.CLIP_METADATA)], + model_response=EmbeddingResponse(), + model=self.MODEL, + response_json={"embedding": {"values": [0.1, 0.2]}}, + ) + assert result.usage.prompt_tokens == 0 + + def test_image_block_counts_as_image_only_input(self) -> None: + result: Final = process_embed_content_response( + input=[_file_block(file_data=IMAGE_DATA_URI)], + model_response=EmbeddingResponse(), + model=self.MODEL, + response_json={ + "embedding": {"values": [0.1, 0.2]}, + "usageMetadata": {"promptTokenCount": 258, "totalTokenCount": 258}, + }, + ) + assert result.usage.prompt_tokens_details.image_tokens == 258 diff --git a/tests/unit/llms/vertex_ai/test_gemini_batch_embeddings.py b/tests/unit/llms/vertex_ai/test_gemini_batch_embeddings.py index 98abf5459df..7271fb48861 100644 --- a/tests/unit/llms/vertex_ai/test_gemini_batch_embeddings.py +++ b/tests/unit/llms/vertex_ai/test_gemini_batch_embeddings.py @@ -8,11 +8,15 @@ This test ensures that: """ import json +from contextlib import AbstractContextManager +from typing import Final from unittest.mock import MagicMock, patch import pytest import litellm +import httpx + from litellm.llms.custom_httpx.http_handler import HTTPHandler from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import ( _filter_embed_params, @@ -655,3 +659,163 @@ def test_batch_embeddings_response_has_correct_indices_and_order(): assert ( embedding.embedding == expected_values[i] ), f"embedding {i} has wrong values: {embedding.embedding}" + + +CLIP_BLOCK = { + "type": "file", + "file": { + "file_data": "data:video/mp4;base64,AAAAIGZ0eXBpc29tAAACAGlzb21pc28yYXZjMW1wNDEAAAAIZnJlZQAA", + "video_metadata": {"fps": 2, "start_offset": "3s", "end_offset": "6s"}, + }, +} +CLIP_PART_METADATA = {"fps": 2.0, "startOffset": "3s", "endOffset": "6s"} + + +def _mock_embedding_call( + response_json: dict[str, object], +) -> tuple[AbstractContextManager[MagicMock], AbstractContextManager[MagicMock], MagicMock]: + mock_get_token: Final = patch( + "litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_handler.GoogleBatchEmbeddings._get_token_and_url" + ) + mock_auth: Final = patch( + "litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_handler.GoogleBatchEmbeddings._ensure_access_token", + side_effect=lambda *args, **kwargs: (None, "test-project"), + ) + mock_response: Final = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = response_json + return mock_get_token, mock_auth, mock_response + + +def test_gemini_batch_path_sends_file_block_video_metadata() -> None: + client: Final = HTTPHandler() + mock_get_token, mock_auth, mock_response = _mock_embedding_call( + {"embeddings": [{"values": [0.1, 0.2]}, {"values": [0.3, 0.4]}]} + ) + with patch.object(client, "post", return_value=mock_response) as mock_post, mock_auth, mock_get_token as token: + token.return_value = ( + {"x-goog-api-key": "test-key"}, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-2-preview:batchEmbedContents", + ) + response: Final = litellm.embedding( + model="gemini/gemini-embedding-2-preview", + input=[CLIP_BLOCK, "a solid color clip"], + api_key="test-key", + client=client, + ) + + request_body: Final = json.loads(mock_post.call_args.kwargs["data"]) + clip_part: Final = request_body["requests"][0]["content"]["parts"][0] + assert clip_part["inline_data"]["mime_type"] == "video/mp4" + assert clip_part["video_metadata"] == CLIP_PART_METADATA + assert request_body["requests"][1]["content"]["parts"] == [{"text": "a solid color clip"}] + assert [row["index"] for row in response.data] == [0, 1] + + +def test_vertex_embed_content_path_sends_file_block_video_metadata() -> None: + client: Final = HTTPHandler() + url: Final = "https://us-central1-aiplatform.googleapis.com/v1/projects/test/locations/us-central1/publishers/google/models/gemini-embedding-2-preview:embedContent" + mock_get_token, mock_auth, mock_response = _mock_embedding_call({"embedding": {"values": [0.1, 0.2]}}) + with patch.object(client, "post", return_value=mock_response) as mock_post, mock_auth, mock_get_token as token: + token.return_value = ({"Authorization": "Bearer test-token"}, url) + response: Final = litellm.embedding( + model="vertex_ai/gemini-embedding-2-preview", + input=[CLIP_BLOCK, "a solid color clip"], + vertex_project="test-project", + vertex_location="us-central1", + client=client, + ) + + data: Final = json.loads(mock_post.call_args.kwargs["data"]) + assert data["content"]["parts"][0]["video_metadata"] == CLIP_PART_METADATA + assert data["content"]["parts"][1] == {"text": "a solid color clip"} + assert len(response.data) == 1 + + +def test_file_block_with_files_reference_is_resolved_through_the_files_api() -> None: + client: Final = HTTPHandler() + files_uri: Final = "https://generativelanguage.googleapis.com/v1beta/files/clip123" + mock_get_token, mock_auth, mock_response = _mock_embedding_call({"embeddings": [{"values": [0.1, 0.2]}]}) + file_lookup: Final = MagicMock() + file_lookup.status_code = 200 + file_lookup.json.return_value = {"mimeType": "video/mp4", "uri": files_uri} + with ( + patch.object(client, "post", return_value=mock_response) as mock_post, + patch.object(client, "get", return_value=file_lookup) as mock_get, + mock_auth, + mock_get_token as token, + ): + token.return_value = ( + {"x-goog-api-key": "test-key"}, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-embedding-2-preview:batchEmbedContents", + ) + litellm.embedding( + model="gemini/gemini-embedding-2-preview", + input=[{"type": "file", "file": {"file_id": "files/clip123", "video_metadata": {"fps": 1}}}], + api_key="test-key", + client=client, + ) + + assert mock_get.call_args.kwargs["url"] == "https://generativelanguage.googleapis.com/v1beta/files/clip123" + clip_part: Final = json.loads(mock_post.call_args.kwargs["data"])["requests"][0]["content"]["parts"][0] + assert clip_part == { + "file_data": {"mime_type": "video/mp4", "file_uri": files_uri}, + "video_metadata": {"fps": 1.0}, + } + + +CLIP_BLOCK_WITH_DETAIL: Final = {"type": "file", "file": {**CLIP_BLOCK["file"], "detail": "high"}} + + +def _recording_client(calls: list[httpx.Request], response_json: dict[str, object]) -> HTTPHandler: + def route(request: httpx.Request) -> httpx.Response: + calls.append(request) + return httpx.Response(200, json=response_json) + + return HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(route))) + + +def test_gemini_drop_params_strips_the_block_keys_embeddings_do_not_take() -> None: + calls: Final[list[httpx.Request]] = [] + client: Final = _recording_client(calls, {"embeddings": [{"values": [0.1, 0.2]}]}) + response: Final = litellm.embedding( + model="gemini/gemini-embedding-2-preview", + input=[CLIP_BLOCK_WITH_DETAIL], + api_key="test-key", + client=client, + drop_params=True, + ) + sent_part: Final = json.loads(calls[0].content)["requests"][0]["content"]["parts"][0] + assert sent_part["video_metadata"] == CLIP_PART_METADATA + assert "detail" not in calls[0].content.decode() + assert response.data[0].embedding == [0.1, 0.2] + + +def test_gemini_block_detail_answers_400_without_drop_params() -> None: + calls: Final[list[httpx.Request]] = [] + client: Final = _recording_client(calls, {"embeddings": [{"values": [0.1, 0.2]}]}) + with pytest.raises(litellm.BadRequestError, match=r"file\.detail"): + litellm.embedding( + model="gemini/gemini-embedding-2-preview", input=[CLIP_BLOCK_WITH_DETAIL], api_key="test-key", client=client + ) + assert calls == [] + + +def test_vertex_drop_params_strips_the_block_keys_embeddings_do_not_take() -> None: + client: Final = HTTPHandler() + url: Final = "https://us-central1-aiplatform.googleapis.com/v1/projects/test/locations/us-central1/publishers/google/models/gemini-embedding-2-preview:embedContent" + mock_get_token, mock_auth, mock_response = _mock_embedding_call({"embedding": {"values": [0.1, 0.2]}}) + with patch.object(client, "post", return_value=mock_response) as mock_post, mock_auth, mock_get_token as token: + token.return_value = ({"Authorization": "Bearer test-token"}, url) + litellm.embedding( + model="vertex_ai/gemini-embedding-2-preview", + input=[CLIP_BLOCK_WITH_DETAIL], + vertex_project="test-project", + vertex_location="us-central1", + client=client, + drop_params=True, + ) + + data: Final = json.loads(mock_post.call_args.kwargs["data"]) + assert data["content"]["parts"][0]["video_metadata"] == CLIP_PART_METADATA + assert "detail" not in mock_post.call_args.kwargs["data"] diff --git a/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py b/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py index 8da95f839b9..d1cd640e590 100644 --- a/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py +++ b/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py @@ -1,6 +1,11 @@ +import json +from types import MappingProxyType +from typing import Final from unittest.mock import MagicMock, patch +import httpx import pytest +import respx import litellm from litellm.llms.vertex_ai.batches.handler import VertexAIBatchPrediction @@ -197,3 +202,134 @@ async def test_litellm_cancel_batch_vertex_ai(): assert mock_instance.cancel_batch.called assert response.id == "batch_123" assert response.status == "cancelling" + + +_MOCK_GCS_FILE_RESPONSE: Final = MappingProxyType( + { + "kind": "storage#object", + "id": "litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb/1739598666670574", + "selfLink": "https://www.googleapis.com/storage/v1/b/litellm-local/o/litellm-vertex-files%2Fpublishers%2Fgoogle%2Fmodels%2Fgemini-1.5-flash-001%2F5f7b99ad-9203-4430-98bf-3b45451af4cb", + "name": "litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb", + "bucket": "litellm-local", + "generation": "1739598666670574", + "metageneration": "1", + "contentType": "application/json", + "storageClass": "STANDARD", + "size": "416", + "md5Hash": "hbBNj7C8KJ7oVH+JmyRM6A==", + "crc32c": "oDmiUA==", + "etag": "CO7D0IT+xIsDEAE=", + "timeCreated": "2025-02-15T05:51:06.741Z", + "updated": "2025-02-15T05:51:06.741Z", + "timeStorageClassUpdated": "2025-02-15T05:51:06.741Z", + "timeFinalized": "2025-02-15T05:51:06.741Z", + } +) + +_MOCK_VERTEX_BATCH_RESPONSE: Final = MappingProxyType( + { + "name": "projects/123456789/locations/us-central1/batchPredictionJobs/test-batch-id-456", + "displayName": "litellm_batch_job", + "model": "projects/123456789/locations/us-central1/models/gemini-1.5-flash-001", + "modelVersionId": "v1", + "inputConfig": { + "gcsSource": { + "uris": [ + "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/5f7b99ad-9203-4430-98bf-3b45451af4cb" + ] + } + }, + "outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://litellm-local/batch-outputs/"}}, + "dedicatedResources": { + "machineSpec": { + "machineType": "n1-standard-4", + "acceleratorType": "NVIDIA_TESLA_T4", + "acceleratorCount": 1, + }, + "startingReplicaCount": 1, + "maxReplicaCount": 1, + }, + "state": "JOB_STATE_RUNNING", + "createTime": "2025-02-15T05:51:06.741Z", + "startTime": "2025-02-15T05:51:07.741Z", + "updateTime": "2025-02-15T05:51:08.741Z", + "labels": {"key1": "value1", "key2": "value2"}, + "completionStats": {"successfulCount": 0, "failedCount": 0, "remainingCount": 100}, + } +) + + +@pytest.mark.asyncio +async def test_vertex_file_upload_create_and_retrieve_batch( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +): + monkeypatch.setenv("GCS_BUCKET_NAME", "litellm-local") + monkeypatch.setenv("VERTEXAI_PROJECT", "mock-project") + monkeypatch.setenv("VERTEXAI_LOCATION", "us-central1") + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + mock_creds: Final = MagicMock(token="mock-token", valid=True, expiry=None) + monkeypatch.setattr("google.auth.default", lambda *args, **kwargs: (mock_creds, "mock-project")) + jobs_url: Final = ( + "https://us-central1-aiplatform.googleapis.com/v1/projects/mock-project/locations/us-central1" + "/batchPredictionJobs" + ) + upload_route: Final = respx_mock.post( + url__startswith="https://storage.googleapis.com/upload/storage/v1/b/litellm-local/o" + ).mock(return_value=httpx.Response(200, json=dict(_MOCK_GCS_FILE_RESPONSE))) + create_route: Final = respx_mock.post(jobs_url).mock( + return_value=httpx.Response(200, json=dict(_MOCK_VERTEX_BATCH_RESPONSE)) + ) + retrieve_route: Final = respx_mock.get(f"{jobs_url}/test-batch-id-456").mock( + return_value=httpx.Response(200, json=dict(_MOCK_VERTEX_BATCH_RESPONSE)) + ) + gcs_object_uri: Final = ( + "gs://litellm-local/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/" + "5f7b99ad-9203-4430-98bf-3b45451af4cb" + ) + + file_obj: Final = await litellm.acreate_file( + file=( + "vertex_batch.jsonl", + b'{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", ' + b'"body": {"model": "gemini-1.5-flash-001", "messages": [{"role": "user", "content": "hi"}]}}\n', + "application/jsonl", + ), + purpose="batch", + custom_llm_provider="vertex_ai", + ) + + assert file_obj.id == gcs_object_uri + assert upload_route.call_count == 1 + upload_request: Final = upload_route.calls.last.request + assert upload_request.url.params["uploadType"] == "media" + assert upload_request.url.params["name"].startswith( + "litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/" + ) + assert upload_request.headers["Content-Type"] == "application/json" + uploaded_row: Final = json.loads(upload_request.content) + assert uploaded_row["request"]["contents"] == [{"role": "user", "parts": [{"text": "hi"}]}] + assert uploaded_row["request"]["labels"]["litellm_custom_id"] == "request-1" + + create_batch_response: Final = await litellm.acreate_batch( + completion_window="24h", + endpoint="/v1/chat/completions", + input_file_id=file_obj.id, + custom_llm_provider="vertex_ai", + metadata={"key1": "value1", "key2": "value2"}, + ) + + create_body: Final = json.loads(create_route.calls.last.request.content) + assert create_body["inputConfig"] == {"gcsSource": {"uris": [gcs_object_uri]}, "instancesFormat": "jsonl"} + assert create_body["model"] == "publishers/google/models/gemini-1.5-flash-001" + assert create_body["outputConfig"]["predictionsFormat"] == "jsonl" + assert create_body["outputConfig"]["gcsDestination"]["outputUriPrefix"].startswith("gs://litellm-local/") + assert create_batch_response.id == "test-batch-id-456" + assert create_batch_response.input_file_id == gcs_object_uri + + retrieved_batch: Final = await litellm.aretrieve_batch( + batch_id=create_batch_response.id, custom_llm_provider="vertex_ai" + ) + + assert retrieve_route.call_count == 1 + assert retrieved_batch.id == "test-batch-id-456" + assert retrieved_batch.input_file_id == gcs_object_uri diff --git a/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py b/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py index 8c73b72a65a..8ca0a4a4335 100644 --- a/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py +++ b/tests/unit/llms/vertex_ai/text_to_speech/test_transformation.py @@ -1,9 +1,13 @@ import base64 from typing import Final +import json +from collections.abc import Iterator, Mapping from unittest.mock import MagicMock, Mock, patch import httpx import pytest +import respx +import responses from pydantic import ValidationError import litellm @@ -663,3 +667,97 @@ def test_transform_text_to_speech_response_rejects_malformed_payloads_without_ec ) assert "input_value" not in str(exc_info.value) + + +_SYNTHESIZE_URL: Final = "https://texttospeech.googleapis.com/v1/text:synthesize" +_AUTHORIZED_USER: Final = json.dumps( + { + "type": "authorized_user", + "client_id": "synthetic-client-id", + "client_secret": "synthetic-client-secret", + "refresh_token": "synthetic-refresh-token", + "quota_project_id": "test-project", + } +) +_ASYNC_INPUT: Final = "async hello what llm guardrail do you have" +_UK_VOICE: Final = {"languageCode": "en-UK", "name": "en-UK-Studio-O"} +_UK_AUDIO_CONFIG: Final = {"audioEncoding": "LINEAR22", "speakingRate": "10"} + + +@pytest.fixture +def google_token_endpoint(monkeypatch: pytest.MonkeyPatch) -> Iterator[responses.RequestsMock]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + with responses.RequestsMock(assert_all_requests_are_fired=False) as token_endpoint: + token_endpoint.post( + "https://oauth2.googleapis.com/token", + json={"access_token": "minted-google-token", "expires_in": 3600, "token_type": "Bearer"}, + ) + yield token_endpoint + litellm.in_memory_llm_clients_cache.flush_cache() + + +async def _aspeech_vertex( + respx_mock: respx.MockRouter, speech_input: str, voice_params: Mapping[str, object] +) -> httpx.Request: + route: Final = respx_mock.post(_SYNTHESIZE_URL).mock( + return_value=httpx.Response(200, json={"audioContent": base64.b64encode(b"vertex-audio").decode()}) + ) + response: Final = await litellm.aspeech( + model="vertex_ai/test", + input=speech_input, + vertex_credentials=_AUTHORIZED_USER, + **voice_params, + ) + assert response.content == b"vertex-audio" + assert route.call_count == 1 + sent: Final = route.calls[0].request + assert sent.headers["x-goog-user-project"] == "test-project" + assert sent.headers["authorization"] == "Bearer minted-google-token" + return sent + + +@pytest.mark.asyncio +async def test_aspeech_vertex_ai_default_voice_posts_synthesize_request( + respx_mock: respx.MockRouter, google_token_endpoint: responses.RequestsMock +) -> None: + sent: Final = await _aspeech_vertex(respx_mock, _ASYNC_INPUT, {}) + + assert json.loads(sent.content) == { + "input": {"text": _ASYNC_INPUT}, + "voice": {"languageCode": "en-US", "name": "en-US-Studio-O"}, + "audioConfig": {"audioEncoding": "LINEAR16", "speakingRate": "1"}, + } + + +@pytest.mark.asyncio +async def test_aspeech_vertex_ai_forwards_caller_voice_and_audio_config( + respx_mock: respx.MockRouter, google_token_endpoint: responses.RequestsMock +) -> None: + sent: Final = await _aspeech_vertex(respx_mock, _ASYNC_INPUT, {"voice": _UK_VOICE, "audioConfig": _UK_AUDIO_CONFIG}) + + assert json.loads(sent.content) == { + "input": {"text": _ASYNC_INPUT}, + "voice": _UK_VOICE, + "audioConfig": _UK_AUDIO_CONFIG, + } + + +@pytest.mark.asyncio +async def test_aspeech_vertex_ai_sends_ssml_input( + respx_mock: respx.MockRouter, google_token_endpoint: responses.RequestsMock +) -> None: + ssml: Final = """ + +

Hello, world!

+

This is a test of the text-to-speech API.

+
+ """ + + sent: Final = await _aspeech_vertex(respx_mock, ssml, {"voice": _UK_VOICE, "audioConfig": _UK_AUDIO_CONFIG}) + + assert json.loads(sent.content) == { + "input": {"ssml": ssml}, + "voice": _UK_VOICE, + "audioConfig": _UK_AUDIO_CONFIG, + } diff --git a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py index 67a32cc82bc..1452f77604f 100644 --- a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py +++ b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py @@ -1,6 +1,7 @@ import copy import json import os +from typing import Final from unittest.mock import MagicMock, patch import pytest @@ -180,6 +181,40 @@ def test_no_per_message_output_config_leaves_per_turn_control_beta_out(): assert "per-turn-control-2026-07-01" not in headers.get("anthropic-beta", "") +def test_inline_tools_beta_reaches_the_vertex_messages_request(local_beta_headers_config: None) -> None: + """Vertex rejects a `tool_addition` system message unless inline-tools-2026-09-15 is on the request, so the + /v1/messages beta filter must keep the header the client sent (source and date in + tests/unit/test_anthropic_beta_headers_filtering.py).""" + from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta + + messages: Final[list[dict[str, object]]] = [ + {"role": "user", "content": "What is the weather in Paris?"}, + { + "role": "system", + "content": [ + { + "type": "tool_addition", + "tool": { + "type": "tool_definition", + "definition": { + "name": "db_query", + "description": "Run a read-only SQL query", + "input_schema": {"type": "object", "properties": {"sql": {"type": "string"}}}, + }, + }, + } + ], + }, + ] + + filtered: Final = update_headers_with_filtered_beta( + headers=_validate_vertex_headers({"anthropic-beta": "inline-tools-2026-09-15"}, messages), + provider="vertex_ai", + ) + + assert filtered["anthropic-beta"].split(",").count("inline-tools-2026-09-15") == 1 + + def test_web_search_header_not_added_without_tool(): """Test that beta header is NOT added when web search tool is not present""" config = VertexAIPartnerModelsAnthropicMessagesConfig() diff --git a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py index ca4a3dedb4a..0208d87b354 100644 --- a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py @@ -1,11 +1,13 @@ import copy import json +from typing import Final import pytest from litellm.anthropic_beta_headers_manager import ( update_headers_with_filtered_beta, + update_request_with_filtered_beta, ) from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.transformation import ( VertexAIAnthropicConfig, @@ -382,6 +384,33 @@ def test_vertex_ai_anthropic_extra_headers_beta_propagation(): assert "interleaved-thinking-2025-05-14" in headers["anthropic-beta"] +def test_vertex_ai_anthropic_inline_tools_beta_survives_the_chat_beta_filter(local_beta_headers_config: None) -> None: + """The chat path runs the Vertex beta filter over both the header and the `anthropic_beta` body field right + before the request goes out, so inline-tools-2026-09-15 must survive both (source and date in + tests/unit/test_anthropic_beta_headers_filtering.py).""" + config: Final = VertexAIAnthropicConfig() + headers: Final[dict[str, str]] = {} + optional_params: Final[dict[str, object]] = { + "max_tokens": 100, + "is_vertex_request": True, + "extra_headers": {"anthropic-beta": "inline-tools-2026-09-15"}, + } + + request_data: Final = config.transform_request( + model="claude-opus-5-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params=optional_params, + litellm_params={}, + headers=headers, + ) + filtered_headers, filtered_request = update_request_with_filtered_beta( + headers=headers, request_data=request_data, provider="vertex_ai" + ) + + assert filtered_headers["anthropic-beta"].split(",").count("inline-tools-2026-09-15") == 1 + assert "inline-tools-2026-09-15" in filtered_request["anthropic_beta"] + + def test_vertex_ai_anthropic_extra_headers_beta_merged_with_auto_betas(): """Test that extra_headers betas are merged with auto-detected betas rather than replacing them.""" diff --git a/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py b/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py index 777e52f4e2a..15689f43bd0 100644 --- a/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py +++ b/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py @@ -295,11 +295,10 @@ class TestVoyageRerankTransform: assert "top_n" in supported_params assert "return_documents" in supported_params - @patch("litellm.llms.voyage.rerank.transformation.get_secret_str") - def test_validate_environment_missing_api_key(self, mock_get_secret_str): + def test_validate_environment_missing_api_key(self, monkeypatch): """Test that validate_environment raises error when API key is missing.""" - # Mock get_secret_str to return None for both environment variables - mock_get_secret_str.return_value = None + for env_var in ("VOYAGE_API_KEY", "VOYAGE_AI_API_KEY", "VOYAGE_AI_TOKEN"): + monkeypatch.delenv(env_var, raising=False) with pytest.raises(ValueError, match="Voyage AI API key is required"): self.config.validate_environment( headers={}, diff --git a/tests/unit/llms/voyage/test_common_utils.py b/tests/unit/llms/voyage/test_common_utils.py new file mode 100644 index 00000000000..51d6c3ed0af --- /dev/null +++ b/tests/unit/llms/voyage/test_common_utils.py @@ -0,0 +1,164 @@ +import pytest + +from litellm.llms.voyage.common_utils import ( + MONGODB_API_BASE, + VOYAGE_API_BASE, + get_default_base_url, + get_voyage_api_key, +) +from litellm.llms.voyage.embedding.transformation import VoyageEmbeddingConfig +from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, +) +from litellm.llms.voyage.embedding.transformation_multimodal import ( + VoyageMultimodalEmbeddingConfig, +) +from litellm.llms.voyage.rerank.transformation import VoyageRerankConfig + +VOYAGE_KEY_ENV_VARS = ("VOYAGE_API_KEY", "VOYAGE_AI_API_KEY", "VOYAGE_AI_TOKEN") + + +@pytest.fixture(autouse=True) +def clear_voyage_env(monkeypatch): + for env_var in VOYAGE_KEY_ENV_VARS: + monkeypatch.delenv(env_var, raising=False) + + +def test_mongodb_key_routes_to_mongodb_host(): + """MongoDB-issued keys carry the `al-` prefix and are only valid on ai.mongodb.com""" + assert get_default_base_url("al-1234567890") == MONGODB_API_BASE + + +@pytest.mark.parametrize("api_key", ["pa-1234567890", "sk-voyage-legacy-key", "al", "", None]) +def test_non_mongodb_key_routes_to_voyage_host(api_key): + assert get_default_base_url(api_key) == VOYAGE_API_BASE + + +@pytest.mark.parametrize("env_var", VOYAGE_KEY_ENV_VARS) +def test_mongodb_key_from_any_supported_env_var_routes_to_mongodb_host(monkeypatch, env_var): + monkeypatch.setenv(env_var, "al-from-env") + assert get_default_base_url() == MONGODB_API_BASE + + +def test_explicit_key_wins_over_env_for_routing(monkeypatch): + monkeypatch.setenv("VOYAGE_API_KEY", "al-from-env") + assert get_default_base_url("pa-explicit") == VOYAGE_API_BASE + + +@pytest.mark.parametrize( + "config, endpoint", + [ + (VoyageEmbeddingConfig(), "embeddings"), + (VoyageContextualEmbeddingConfig(), "contextualizedembeddings"), + (VoyageMultimodalEmbeddingConfig(), "multimodalembeddings"), + ], +) +@pytest.mark.parametrize("api_key, expected_host", [("al-key", MONGODB_API_BASE), ("pa-key", VOYAGE_API_BASE)]) +def test_embedding_configs_route_by_key_prefix(config, endpoint, api_key, expected_host): + url = config.get_complete_url(None, api_key, "voyage-3", {}, {}) + assert url == f"{expected_host}/{endpoint}" + + +@pytest.mark.parametrize( + "config, endpoint", + [ + (VoyageEmbeddingConfig(), "embeddings"), + (VoyageContextualEmbeddingConfig(), "contextualizedembeddings"), + (VoyageMultimodalEmbeddingConfig(), "multimodalembeddings"), + ], +) +def test_explicit_api_base_overrides_key_routing(config, endpoint): + url = config.get_complete_url("https://gateway.internal/v1", "al-key", "voyage-3", {}, {}) + assert url == f"https://gateway.internal/v1/{endpoint}" + + +@pytest.mark.parametrize("api_key, expected_host", [("al-key", MONGODB_API_BASE), ("pa-key", VOYAGE_API_BASE)]) +def test_rerank_routes_by_request_key_prefix(api_key, expected_host): + config = VoyageRerankConfig() + config.validate_environment({}, "rerank-2.5", api_key=api_key) + + assert config.get_complete_url(None, "rerank-2.5") == f"{expected_host}/rerank" + + +@pytest.mark.parametrize("api_key, expected_host", [("al-key", MONGODB_API_BASE), ("pa-key", VOYAGE_API_BASE)]) +def test_rerank_routes_by_env_key_prefix(monkeypatch, api_key, expected_host): + monkeypatch.setenv("VOYAGE_API_KEY", api_key) + config = VoyageRerankConfig() + config.validate_environment({}, "rerank-2.5") + + assert config.get_complete_url(None, "rerank-2.5") == f"{expected_host}/rerank" + + +def test_rerank_request_key_beats_env_key_for_routing(monkeypatch): + """A MongoDB key on the request must not be posted to the Voyage host the env key names""" + monkeypatch.setenv("VOYAGE_API_KEY", "pa-from-env") + config = VoyageRerankConfig() + + headers = config.validate_environment({}, "rerank-2.5", api_key="al-on-request") + + assert headers["Authorization"] == "Bearer al-on-request" + assert config.get_complete_url(None, "rerank-2.5") == f"{MONGODB_API_BASE}/rerank" + + +def test_rerank_config_is_built_per_request_so_keys_cannot_leak(monkeypatch): + """get_complete_url reads a key off the instance, so each request must get its own instance""" + import litellm + from litellm.types.utils import LlmProviders + from litellm.utils import ProviderConfigManager + + monkeypatch.delenv("VOYAGE_API_KEY", raising=False) + first = ProviderConfigManager.get_provider_rerank_config( + model="rerank-2.5", provider=LlmProviders.VOYAGE, api_base=None, present_version_params=[] + ) + second = ProviderConfigManager.get_provider_rerank_config( + model="rerank-2.5", provider=LlmProviders.VOYAGE, api_base=None, present_version_params=[] + ) + assert isinstance(first, litellm.VoyageRerankConfig) and first is not second + + first.validate_environment({}, "rerank-2.5", api_key="al-first-request") + second.validate_environment({}, "rerank-2.5", api_key="pa-second-request") + + assert first.get_complete_url(None, "rerank-2.5") == f"{MONGODB_API_BASE}/rerank" + assert second.get_complete_url(None, "rerank-2.5") == f"{VOYAGE_API_BASE}/rerank" + + +def test_rerank_falls_back_to_env_when_validate_environment_did_not_run(monkeypatch): + """A caller that skips validate_environment keeps the pre-existing env-only behaviour""" + monkeypatch.setenv("VOYAGE_API_KEY", "al-from-env") + + assert VoyageRerankConfig().get_complete_url(None, "rerank-2.5") == f"{MONGODB_API_BASE}/rerank" + + +@pytest.mark.parametrize( + "config", + [VoyageEmbeddingConfig(), VoyageContextualEmbeddingConfig(), VoyageMultimodalEmbeddingConfig()], +) +def test_auth_header_uses_the_key_the_url_was_routed_on(monkeypatch, config): + """The host is picked from a key, so the Authorization header has to carry that same key""" + monkeypatch.setenv("VOYAGE_AI_TOKEN", "al-from-env") + + headers = config.validate_environment({}, "voyage-3", [], {}, {}) + url = config.get_complete_url(None, None, "voyage-3", {}, {}) + + assert headers["Authorization"] == "Bearer al-from-env" + assert url.startswith(MONGODB_API_BASE) + + +def test_rerank_auth_header_uses_the_key_the_url_was_routed_on(monkeypatch): + monkeypatch.setenv("VOYAGE_AI_TOKEN", "al-from-env") + config = VoyageRerankConfig() + + headers = config.validate_environment({}, "rerank-2.5") + + assert headers["Authorization"] == "Bearer al-from-env" + assert config.get_complete_url(None, "rerank-2.5").startswith(MONGODB_API_BASE) + + +def test_get_voyage_api_key_prefers_env_vars_in_documented_order(monkeypatch): + monkeypatch.setenv("VOYAGE_AI_API_KEY", "second") + monkeypatch.setenv("VOYAGE_AI_TOKEN", "third") + assert get_voyage_api_key() == "second" + + monkeypatch.setenv("VOYAGE_API_KEY", "first") + assert get_voyage_api_key() == "first" + assert get_voyage_api_key("explicit") == "explicit" diff --git a/tests/unit/llms/voyage/test_voyage_multimodal_embedding.py b/tests/unit/llms/voyage/test_voyage_multimodal_embedding.py index f3e6885cbe6..d13610ada17 100644 --- a/tests/unit/llms/voyage/test_voyage_multimodal_embedding.py +++ b/tests/unit/llms/voyage/test_voyage_multimodal_embedding.py @@ -172,15 +172,12 @@ class TestVoyageMultimodalEmbeddings: assert headers == {"Authorization": "Bearer test-key"} def test_validate_environment_uses_secret_fallback(self, monkeypatch): - import litellm.llms.voyage.embedding.transformation_multimodal as module from litellm.llms.voyage.embedding.transformation_multimodal import ( VoyageMultimodalEmbeddingConfig, ) - def fake_get_secret(name): - return "secret-key" if name == "VOYAGE_AI_API_KEY" else None - - monkeypatch.setattr(module, "get_secret_str", fake_get_secret) + monkeypatch.delenv("VOYAGE_API_KEY", raising=False) + monkeypatch.setenv("VOYAGE_AI_API_KEY", "secret-key") config = VoyageMultimodalEmbeddingConfig() headers = config.validate_environment( {}, "voyage-multimodal-3.5", [], {}, {}, api_key=None @@ -188,12 +185,12 @@ class TestVoyageMultimodalEmbeddings: assert headers == {"Authorization": "Bearer secret-key"} def test_validate_environment_raises_without_api_key(self, monkeypatch): - import litellm.llms.voyage.embedding.transformation_multimodal as module from litellm.llms.voyage.embedding.transformation_multimodal import ( VoyageMultimodalEmbeddingConfig, ) - monkeypatch.setattr(module, "get_secret_str", lambda name: None) + for env_var in ("VOYAGE_API_KEY", "VOYAGE_AI_API_KEY", "VOYAGE_AI_TOKEN"): + monkeypatch.delenv(env_var, raising=False) config = VoyageMultimodalEmbeddingConfig() with pytest.raises(ValueError, match='Voyage API key is required for multimodal embeddings\\. Set') as exc_info: config.validate_environment( diff --git a/tests/unit/messages/test_dispatch.py b/tests/unit/messages/test_dispatch.py index 8668e11ee65..faaefa1437e 100644 --- a/tests/unit/messages/test_dispatch.py +++ b/tests/unit/messages/test_dispatch.py @@ -154,13 +154,13 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None: ) assert result is expected request, call_args, call_kwargs = captured[0] - assert request.bound["model"] == "anthropic/claude-sonnet-4-5" - assert request.bound["messages"] is MESSAGES - assert request.bound["max_tokens"] == 16 - assert request.bound["stream"] is True - assert request.bound["api_key"] == "sk-test" - assert request.bound["api_base"] == "https://example.invalid" - assert request.bound["custom_llm_provider"] == "anthropic" + assert request.resolved["model"] == "anthropic/claude-sonnet-4-5" + assert request.resolved["messages"] is MESSAGES + assert request.resolved["max_tokens"] == 16 + assert request.resolved["stream"] is True + assert request.resolved["api_key"] == "sk-test" + assert request.resolved["api_base"] == "https://example.invalid" + assert request.resolved["custom_llm_provider"] == "anthropic" assert request.kwargs == kwargs assert request.kwargs["litellm_metadata"] is metadata assert call_args == args @@ -246,7 +246,7 @@ def test_anthropic_create_routes_through_dispatch(monkeypatch: pytest.MonkeyPatc finally: NATIVE_MESSAGES.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["claude-sonnet-4-5"] + assert [request.resolved["model"] for request in captured] == ["claude-sonnet-4-5"] @pytest.mark.asyncio @@ -268,7 +268,7 @@ async def test_anthropic_acreate_routes_through_dispatch(monkeypatch: pytest.Mon finally: NATIVE_AMESSAGES.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["claude-sonnet-4-5"] + assert [request.resolved["model"] for request in captured] == ["claude-sonnet-4-5"] @pytest.mark.asyncio @@ -287,10 +287,10 @@ def test_sync_messages_request_projects_public_arguments() -> None: expected: Final = AnthropicMessagesResponse(model="claude-test") def native(request: NativeCall) -> AnthropicMessagesResponse: - assert request.bound["model"] == "claude-test" - assert request.bound["messages"] == MESSAGES - assert request.bound["max_tokens"] == 10 - assert request.bound["custom_llm_provider"] == "anthropic" + assert request.resolved["model"] == "claude-test" + assert request.resolved["messages"] == MESSAGES + assert request.resolved["max_tokens"] == 10 + assert request.resolved["custom_llm_provider"] == "anthropic" return expected binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None) diff --git a/tests/unit/ocr/test_dispatch.py b/tests/unit/ocr/test_dispatch.py index 3b0d75ae07a..34ee8e40552 100644 --- a/tests/unit/ocr/test_dispatch.py +++ b/tests/unit/ocr/test_dispatch.py @@ -72,13 +72,13 @@ def test_native_receives_normalized_positional_request_and_original_call_shape() request, call_args, call_kwargs = captured[0] assert result is expected - assert request.bound["model"] == "mistral/mistral-ocr-latest" - assert request.bound["document"] is document - assert request.bound["api_key"] == "test-key" - assert request.bound["api_base"] == "https://example.invalid" - assert request.bound["timeout"] is timeout - assert request.bound["custom_llm_provider"] == "mistral" - assert request.bound["extra_headers"] is extra_headers + assert request.resolved["model"] == "mistral/mistral-ocr-latest" + assert request.resolved["document"] is document + assert request.resolved["api_key"] == "test-key" + assert request.resolved["api_base"] == "https://example.invalid" + assert request.resolved["timeout"] is timeout + assert request.resolved["custom_llm_provider"] == "mistral" + assert request.resolved["extra_headers"] is extra_headers assert request.kwargs == kwargs assert request.kwargs["pages"] is pages assert call_args is args @@ -119,8 +119,8 @@ def test_native_preserves_keyword_model_and_document_in_original_call_shape() -> request, call_args, call_kwargs = captured[0] assert result is expected - assert request.bound["model"] == "mistral/mistral-ocr-latest" - assert request.bound["document"] is document + assert request.resolved["model"] == "mistral/mistral-ocr-latest" + assert request.resolved["document"] is document assert request.kwargs == kwargs assert call_args is args assert call_kwargs is kwargs @@ -271,7 +271,7 @@ def test_public_ocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> finally: NATIVE_OCR.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["mistral/mistral-ocr-latest"] + assert [request.resolved["model"] for request in captured] == ["mistral/mistral-ocr-latest"] @pytest.mark.asyncio @@ -297,4 +297,4 @@ async def test_public_aocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPat finally: NATIVE_AOCR.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["mistral/mistral-ocr-latest"] + assert [request.resolved["model"] for request in captured] == ["mistral/mistral-ocr-latest"] diff --git a/tests/unit/passthrough/test_main.py b/tests/unit/passthrough/test_main.py new file mode 100644 index 00000000000..78fa83d126b --- /dev/null +++ b/tests/unit/passthrough/test_main.py @@ -0,0 +1,53 @@ +import json +from typing import Final + +import httpx +import pytest +import respx + +import litellm + +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.vllm.passthrough.transformation import VLLMPassthroughConfig +from litellm.passthrough.main import allm_passthrough_route +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager + +_API_BASE: Final = "http://vllm-upstream.test:8090" + + +@pytest.fixture(autouse=True) +def _httpx_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + +def test_hosted_vllm_resolves_vllm_passthrough_config() -> None: + cfg: Final = ProviderConfigManager.get_provider_passthrough_config( + model="hosted_vllm/my-deployment", + provider=LlmProviders.HOSTED_VLLM, + ) + assert isinstance(cfg, VLLMPassthroughConfig) + + +@pytest.mark.asyncio +async def test_allm_passthrough_route_hosted_vllm_sends_normalized_model(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.post(f"{_API_BASE}/v1/chat/completions").mock( + return_value=httpx.Response(200, json={"ok": True}) + ) + client: Final = AsyncHTTPHandler() + response: Final = await allm_passthrough_route( + method="POST", + endpoint="v1/chat/completions", + model="hosted_vllm/my-deployment", + api_base=_API_BASE, + json={ + "model": "anything", + "messages": [{"role": "user", "content": "Hello"}], + }, + client=client, + ) + assert response.status_code == 200 + assert route.call_count == 1 + outbound: Final = json.loads(route.calls[0].request.content) + assert outbound["model"] == "my-deployment" + assert outbound["messages"] == [{"role": "user", "content": "Hello"}] diff --git a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py index df068b60338..350e7f1f910 100644 --- a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py +++ b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py @@ -52,7 +52,7 @@ def _endpoint(body, sink=None): def _recording_persist(sink): async def persist( - user_id, server_id, access_token, refresh_token, expires_in, scopes + user_id, server_id, access_token, refresh_token, expires_in, scopes, cimd_client_id=None ): sink.append( (user_id, server_id, access_token, refresh_token, expires_in, scopes) @@ -325,3 +325,83 @@ async def test_verified_refresh_preserves_binding_proof_in_storage(): assert token.refresh_token == "rotated" assert token.identity_binding_proof == "verified-binding" assert persist.await_args.kwargs["identity_binding_proof"] == "verified-binding" + + +@pytest.mark.asyncio +async def test_cimd_refresh_on_fresh_replica_preserves_user_and_client_identity(monkeypatch): + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + server = MCPServer.model_validate( + { + "server_id": "srv", + "name": "srv", + "server_name": "srv", + "transport": "http", + "url": "https://mcp.example.com/mcp", + "auth_type": "oauth2", + "oauth2_flow": "authorization_code", + "token_url": "https://idp.example.com/token", + "client_id_metadata_document_supported": True, + } + ) + posted = [] + persisted = [] + refreshed = await _refresher( + server=server, + body={"access_token": "new-at", "expires_in": 3600}, + post_sink=posted, + persist_sink=persisted, + ).refresh("alice", "srv", OAuthToken(access_token="old-at", refresh_token="old-rt")) + assert refreshed is not None + assert refreshed.access_token == "new-at" + assert posted == [ + ( + "https://idp.example.com/token", + { + "grant_type": "refresh_token", + "refresh_token": "old-rt", + "client_id": "https://gateway.example.com/oauth/client-metadata.json", + }, + {}, + ) + ] + assert persisted == [("alice", "srv", "new-at", "old-rt", 3600, None)] + assert server.client_id is None + + +@pytest.mark.asyncio +async def test_saved_cimd_identity_refreshes_without_discovery(monkeypatch): + from litellm.proxy._experimental.mcp_server.outbound_credentials.v2_token_store import V2PerUserTokenStore + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + identity = "https://gateway.example.com/oauth/client-metadata.json" + stored = {"access_token": "old", "refresh_token": "old-rt", "cimd_client_id": identity} + + async def read_credential(user_id, server_id): + assert (user_id, server_id) == ("alice", "srv") + return stored + + async def persist(user_id, server_id, access_token, refresh_token, expires_in, scopes, **metadata): + assert (user_id, server_id, access_token) == ("alice", "srv", "new") + assert metadata["cimd_client_id"] == identity + + async def post(url, form, headers): + assert url == "https://idp.example.com/token" + assert form["client_id"] == identity + assert "client_secret" not in form + assert "Authorization" not in headers + return {"access_token": "new", "refresh_token": "rotated"} + + server = MCPServer( + server_id="srv", name="srv", transport="http", auth_type="oauth2", oauth2_flow="authorization_code", + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + ) + token = await V2PerUserTokenStore(read_credential).fetch("alice", "srv") + assert token is not None + refreshed = await AuthorizationCodeRefresher(lambda _: server, post, persist).refresh("alice", "srv", token) + assert refreshed is not None + assert refreshed.access_token == "new" + assert refreshed.refresh_token == "rotated" + assert refreshed.cimd_client_id == identity diff --git a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py index de86c1b62f4..10993280e89 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py @@ -708,7 +708,8 @@ async def test_store_user_oauth_credential_does_not_persist_plaintext(): @pytest.mark.asyncio -async def test_oauth_round_trip_returns_payload(): +@pytest.mark.parametrize("cimd_client_id", [None, "https://gateway.example.com/oauth/client-metadata.json"]) +async def test_oauth_round_trip_returns_payload(cimd_client_id): access_token = "ya29.a0AfH6SMBverysecretaccesstoken" prisma = _make_prisma_with_existing(row=None) await store_user_oauth_credential( @@ -719,6 +720,7 @@ async def test_oauth_round_trip_returns_payload(): refresh_token="rfr-xyz", scopes=["a", "b"], identity_binding_proof="verified-proof", + cimd_client_id=cimd_client_id, ) stored = _stored_value(prisma) @@ -734,6 +736,7 @@ async def test_oauth_round_trip_returns_payload(): assert result["refresh_token"] == "rfr-xyz" assert result["scopes"] == ["a", "b"] assert result["identity_binding_proof"] == "verified-proof" + assert result.get("cimd_client_id") == cimd_client_id @pytest.mark.asyncio @@ -1795,3 +1798,61 @@ async def test_unverified_legacy_cache_cannot_bypass_enforcement(monkeypatch): await mcp_per_user_token_cache.set("alice", "srv", "bob", 60) assert await module.resolve_user_oauth_access_token("alice", server) is None assert await mcp_per_user_token_cache.get("alice", "srv") is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("owner", ["native", "legacy"]) +@pytest.mark.parametrize("capability", [None, False]) +async def test_saved_cimd_grant_refreshes_from_encrypted_storage_on_a_fresh_replica(monkeypatch, respx_mock, owner, capability): + from urllib.parse import parse_qs + + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import db as db_module + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _store_per_user_token_server_side + from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import ( + LazyPerUserOAuthTokenStore, + ) + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + identity = "https://gateway.example.com/oauth/client-metadata.json" + prisma = _make_prisma_with_existing(row=None) + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + server = MCPServer( + server_id="saved-cimd", name="saved_cimd", transport="http", auth_type="oauth2", + client_id_metadata_document_supported=capability, + oauth2_flow="authorization_code", authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + ) + await _store_per_user_token_server_side( + server, "alice", {"access_token": "old", "refresh_token": "old-refresh", "expires_in": -10}, + cimd_client_id=identity, + ) + + async def read_row(**kwargs): + return SimpleNamespace(credential_b64=_stored_value(prisma), user_id="alice", server_id=server.server_id) + + prisma.db.litellm_mcpusercredentials.find_unique.side_effect = read_row + saved = await get_user_oauth_credential(prisma, "alice", server.server_id) + assert saved is not None and saved["cimd_client_id"] == identity + assert identity not in _stored_value(prisma) + monkeypatch.setenv("PROXY_BASE_URL", "https://renamed-gateway.example.com") + upstream = respx_mock.post("https://idp.example.com/token").respond( + 200, json={"access_token": "fresh", "refresh_token": "rotated", "expires_in": 3600} + ) + if owner == "native": + refreshed = await LazyPerUserOAuthTokenStore(lambda _: server).fetch("alice", server.server_id) + assert refreshed is not None and refreshed.access_token == "fresh" + assert refreshed.cimd_client_id == identity + else: + legacy = await db_module.refresh_user_oauth_token(prisma, "alice", server, saved) + assert legacy is not None and legacy["access_token"] == "fresh" + persisted = await get_user_oauth_credential(prisma, "alice", server.server_id) + assert persisted is not None + assert persisted["cimd_client_id"] == identity + assert persisted["refresh_token"] == "rotated" + assert upstream.call_count == 1 + assert parse_qs(upstream.calls[0].request.content.decode())["client_id"] == [identity] + assert "client_secret" not in parse_qs(upstream.calls[0].request.content.decode()) + assert server.client_id is None and server.client_id_metadata_document_supported is capability diff --git a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 23bdfa3eadf..4f22200145e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -20,9 +20,9 @@ if TYPE_CHECKING: import httpx from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey from fastapi import APIRouter + from respx import MockRouter from litellm.proxy.auth.handle_jwt import JWTHandler - from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -12568,10 +12568,12 @@ async def test_oauth_write_denial_does_not_erase_identity_binding( @pytest.mark.asyncio @pytest.mark.parametrize("admin_only", [False, True]) +@pytest.mark.parametrize("cimd", [False, True]) async def test_signed_oauth_callback_honors_credential_write_policy( jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], monkeypatch: pytest.MonkeyPatch, admin_only: bool, + cimd: bool, ) -> None: import httpx import litellm @@ -12586,7 +12588,8 @@ async def test_signed_oauth_callback_honors_credential_write_policy( server: Final = MCPServer( server_id="signed-server", name="signed-server", transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="client", + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id=None if cimd else "client", + client_id_metadata_document_supported=cimd, token_url="https://upstream.example.test/token", ) monkeypatch.setattr(proxy_server, "general_settings", { @@ -12594,6 +12597,7 @@ async def test_signed_oauth_callback_honors_credential_write_policy( "admin_only_routes": [f"/v1/mcp/server/{server.server_id}/oauth-user-credential"] if admin_only else [], }) monkeypatch.setenv("LITELLM_SALT_KEY", "signed-oauth-test-salt") + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") manager: Final = MagicMock() manager.get_allowed_mcp_servers = AsyncMock(return_value=[server.server_id]) manager.invalidate_user_oauth_token_cache = AsyncMock() @@ -12630,6 +12634,16 @@ async def test_signed_oauth_callback_honors_credential_write_policy( "user_id": "jwt-owner", "server_id": server.server_id, } + from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper + + payload: Final = json.loads(decrypt_value_helper( + table.upsert.call_args.kwargs["data"]["create"]["credential_b64"], "credential_b64" + )) + if cimd: + assert payload["cimd_client_id"] == "https://gateway.example.com/oauth/client-metadata.json" + else: + assert "cimd_client_id" not in payload + @pytest.mark.asyncio @pytest.mark.parametrize("allowed", [False, True]) @@ -13362,3 +13376,402 @@ async def test_register_application_type_keeps_no_registration_endpoint_fallback "redirect_uris": ["https://gateway.example/callback"], } assert len(upstream.calls) == 0 + + +def _cimd_oauth_server(): + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + return MCPServer.model_validate( + { + "server_id": "cimd-server", + "name": "cimd-server", + "server_name": "cimd-server", + "url": "https://mcp.example.com/mcp", + "transport": "http", + "auth_type": "oauth2", + "oauth2_flow": "authorization_code", + "authorization_url": "https://idp.example.com/authorize", + "token_url": "https://idp.example.com/token", + "client_id_metadata_document_supported": True, + } + ) + + +def _cimd_request(): + from starlette.requests import Request + + return Request( + { + "type": "http", + "method": "GET", + "scheme": "https", + "path": "/", + "root_path": "", + "query_string": b"", + "headers": [], + "server": ("gateway.example.com", 443), + "client": ("127.0.0.1", 10000), + } + ) + + +@pytest.mark.asyncio +async def test_cimd_registration_returns_https_identity_without_dcr(monkeypatch, respx_mock): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + response = await endpoints.register_client_with_server( + _cimd_request(), _cimd_oauth_server(), "Gateway", None, None, None + ) + body = json.loads(response.body) if hasattr(response, "body") else response + assert body["client_id"] == "https://gateway.example.com/oauth/client-metadata.json" + assert "client_secret" not in body + assert len(respx_mock.calls) == 0 + + +@pytest.mark.asyncio +async def test_cimd_authorization_uses_metadata_identity_and_s256(monkeypatch): + from urllib.parse import parse_qs, urlparse + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + monkeypatch.setenv("LITELLM_SALT_KEY", "cimd-test-state-signing-key") + response = await endpoints.authorize_with_server( + _cimd_request(), + _cimd_oauth_server(), + "placeholder", + "https://gateway.example.com/ui/", + code_challenge="a" * 43, + code_challenge_method="S256", + ) + params = parse_qs(urlparse(response.headers["location"]).query) + assert params["client_id"] == ["https://gateway.example.com/oauth/client-metadata.json"] + assert params["redirect_uri"] == ["https://gateway.example.com/callback"] + assert params["code_challenge_method"] == ["S256"] + + +@pytest.mark.asyncio +async def test_cimd_authorization_rejects_missing_pkce(monkeypatch): + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + monkeypatch.setenv("LITELLM_SALT_KEY", "cimd-test-state-signing-key") + with pytest.raises(HTTPException) as exc: + await endpoints.authorize_with_server( + _cimd_request(), _cimd_oauth_server(), "placeholder", "https://gateway.example.com/ui/" + ) + assert exc.value.status_code == 400 + + +def test_cimd_refresh_request_uses_same_identity_without_caller_secret(monkeypatch): + from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + request = build_upstream_oauth2_token_request( + _cimd_oauth_server(), auth_method=None, client_id="placeholder", client_secret="dummy" + ) + assert request.body["client_id"] == "https://gateway.example.com/oauth/client-metadata.json" + assert "client_secret" not in request.body + assert "Authorization" not in request.headers + + +@pytest.mark.parametrize( + "base", [None, "http://gateway.example.com", "invalid", "https://user:secret@gateway.example.com"] +) +@pytest.mark.asyncio +async def test_cimd_without_stable_https_origin_reports_actionable_error(monkeypatch, base, respx_mock): + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.delenv("PROXY_BASE_URL", raising=False) + if base is not None: + monkeypatch.setenv("PROXY_BASE_URL", base) + with pytest.raises(HTTPException) as exc: + await endpoints.register_client_with_server(_cimd_request(), _cimd_oauth_server(), "Gateway", None, None, None) + assert exc.value.status_code == 400 + assert "HTTPS PROXY_BASE_URL" in str(exc.value.detail) + assert len(respx_mock.calls) == 0 + + +@pytest.mark.parametrize("base", [None, "http://gateway.example.com"]) +@pytest.mark.asyncio +async def test_cimd_without_https_origin_falls_back_to_available_dcr(monkeypatch, base, respx_mock): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.delenv("PROXY_BASE_URL", raising=False) + if base is not None: + monkeypatch.setenv("PROXY_BASE_URL", base) + server = _cimd_oauth_server().model_copy(update={"registration_url": "https://idp.example.com/register"}) + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + post = respx_mock.post("https://idp.example.com/register").respond(201, json={"client_id": "registered-client"}) + response = await endpoints.register_client_with_server(_cimd_request(), server, "Gateway", None, None, None) + assert json.loads(response.body)["client_id"] == "registered-client" + assert post.call_count == 1 + + +@pytest.mark.asyncio +async def test_cimd_yields_to_dynamic_registration_when_the_authorization_server_offers_both(monkeypatch, respx_mock): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + from litellm.proxy._experimental.mcp_server.oauth_utils import get_cimd_client_id + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + server = _cimd_oauth_server().model_copy(update={"registration_url": "https://idp.example.com/register"}) + post = respx_mock.post("https://idp.example.com/register").respond(201, json={"client_id": "registered-client"}) + response = await endpoints.register_client_with_server(_cimd_request(), server, "Gateway", None, None, None) + assert json.loads(response.body)["client_id"] == "registered-client" + assert post.call_count == 1 + assert get_cimd_client_id(server) is None + + +@pytest.mark.asyncio +async def test_cimd_is_preferred_over_dynamic_registration_when_the_deployment_opts_in(monkeypatch, respx_mock): + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + monkeypatch.setattr(proxy_server, "general_settings", {"mcp_prefer_client_id_metadata_document": True}) + server = _cimd_oauth_server().model_copy(update={"registration_url": "https://idp.example.com/register"}) + post = respx_mock.post("https://idp.example.com/register").respond(201, json={"client_id": "registered-client"}) + response = await endpoints.register_client_with_server(_cimd_request(), server, "Gateway", None, None, None) + body = json.loads(response.body) if hasattr(response, "body") else response + assert body["client_id"] == "https://gateway.example.com/oauth/client-metadata.json" + assert "client_secret" not in body + assert post.call_count == 0 + + +@pytest.mark.parametrize( + "updates", + [ + {"client_id": "static-client", "client_secret": "static-secret"}, + {"client_id": "persisted-dcr-client", "dcr_issuer": "https://idp.example.com"}, + {"client_id_metadata_document_supported": False}, + {"auth_type": "oauth_delegate", "dcr_bridge": True}, + {"auth_type": "true_passthrough", "dcr_bridge": True}, + {"delegate_auth_to_upstream": True}, + {"oauth2_flow": "client_credentials"}, + {"client_secret": "configured-secret"}, + {"token_endpoint_auth_method": "client_secret_basic"}, + ], +) +def test_cimd_preserves_existing_identity_and_other_auth_modes(monkeypatch, updates): + from litellm.proxy._experimental.mcp_server.oauth_utils import get_cimd_client_id + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + assert get_cimd_client_id(_cimd_oauth_server().model_copy(update=updates)) is None + + +@pytest.mark.asyncio +async def test_cimd_document_is_public_and_binds_configured_origin(monkeypatch): + import httpx + from fastapi import FastAPI + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com/proxy") + app = FastAPI() + app.include_router(endpoints.router) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="https://attacker.example") as client: + response = await client.get("/oauth/client-metadata.json", headers={"X-Forwarded-Host": "attacker.example"}) + assert response.status_code == 200 + assert response.json()["client_id"] == "https://gateway.example.com/proxy/oauth/client-metadata.json" + assert response.json()["redirect_uris"] == ["https://gateway.example.com/proxy/callback"] + assert response.json()["token_endpoint_auth_method"] == "none" + assert "client_secret" not in response.json() + assert response.headers["cache-control"] == "public, max-age=300" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("base", [None, "https://[invalid"]) +async def test_cimd_document_is_unavailable_without_configured_https_origin(monkeypatch, base): + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.delenv("PROXY_BASE_URL", raising=False) + if base is not None: + monkeypatch.setenv("PROXY_BASE_URL", base) + with pytest.raises(HTTPException) as exc: + await endpoints.oauth_client_metadata() + assert exc.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_cimd_upstream_client_rejection_is_gateway_fault(monkeypatch, respx_mock): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + upstream = respx_mock.post("https://idp.example.com/token").respond( + 401, json={"error": "invalid_client", "error_description": "provider-private-detail"} + ) + response = await endpoints.exchange_token_with_server( + request=_cimd_request(), + mcp_server=_cimd_oauth_server(), + grant_type="authorization_code", + code="code", + redirect_uri="https://gateway.example.com/callback", + client_id="caller-placeholder", + client_secret="dummy", + code_verifier="verifier", + ) + assert upstream.call_count == 1 + assert response.status_code == 502 + body = json.loads(response.body) + assert body["error"] == "server_error" + assert "provider-private-detail" not in body["error_description"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("flow", ["refresh", "authorize"]) +async def test_optional_cimd_discovery_preserves_the_callers_configured_endpoint( + monkeypatch: pytest.MonkeyPatch, respx_mock: "MockRouter", flow: str +) -> None: + from urllib.parse import parse_qs + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "0") + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + monkeypatch.setenv("LITELLM_SALT_KEY", "cimd-test-state-signing-key") + token: Final = respx_mock.post("https://idp.example.com/token").respond( + 200, json={"access_token": "upstream-access", "token_type": "Bearer", "expires_in": 3600} + ) + discovery: Final = respx_mock.route().respond(503) + manager: Final = manager_module.MCPServerManager() + monkeypatch.setattr(manager_module, "global_mcp_server_manager", manager) + await manager.load_servers_from_config({"manual": { + "url": "https://mcp.example.com/mcp", "transport": "http", "auth_type": "oauth2", + "oauth2_flow": "authorization_code", + **({"token_url": "https://idp.example.com/token"} if flow == "refresh" + else {"authorization_url": "https://idp.example.com/authorize"}), + }}) + server: Final = next(iter(manager.config_mcp_servers.values())) + async with manager.catalog.operation(): + if flow == "refresh": + response: Final = await endpoints.exchange_token_with_server( + request=_cimd_request(), mcp_server=server, grant_type="refresh_token", + refresh_token="existing-refresh", client_id="existing-client", + code=None, redirect_uri=None, client_secret=None, code_verifier=None, + ) + assert response.status_code == 200 + assert json.loads(response.body)["access_token"] == "upstream-access" + assert parse_qs(token.calls[0].request.content.decode())["refresh_token"] == ["existing-refresh"] + else: + redirect: Final = await endpoints.authorize_with_server( + _cimd_request(), server, "existing-client", "https://gateway.example.com/ui/", + code_challenge="a" * 43, code_challenge_method="S256", + ) + assert redirect.status_code == 307 + assert redirect.headers["location"].startswith("https://idp.example.com/authorize?") + assert discovery.call_count > 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("origin", [None, "https://renamed-gateway.example.com"]) +async def test_token_route_preserves_saved_cimd_grant_after_origin_change( + jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"], + monkeypatch: pytest.MonkeyPatch, + respx_mock: "MockRouter", + origin: str | None, +) -> None: + from types import SimpleNamespace + from urllib.parse import parse_qs + + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import db, mcp_server_manager + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + monkeypatch.setenv("LITELLM_SALT_KEY", "saved-cimd-route-test-salt") + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + server: Final = _cimd_oauth_server() + mcp_server_manager.global_mcp_server_manager.registry[server.server_id] = server + identity: Final = "https://gateway.example.com/oauth/client-metadata.json" + table: Final = proxy_server.prisma_client.db.litellm_mcpusercredentials + table.find_unique = AsyncMock(return_value=None) + table.upsert = AsyncMock() + await endpoints._store_per_user_token_server_side( + server, "jwt-owner", {"access_token": "expired", "refresh_token": "saved-refresh", "expires_in": -1}, + cimd_client_id=identity, + ) + + async def saved_row(**kwargs): + return SimpleNamespace(credential_b64=table.upsert.call_args.kwargs["data"]["create"]["credential_b64"]) + + table.find_unique.side_effect = saved_row + if origin is None: + monkeypatch.delenv("PROXY_BASE_URL") + else: + monkeypatch.setenv("PROXY_BASE_URL", origin) + _, key = jwt_oauth_identity + upstream: Final = respx_mock.post(server.token_url).respond( + 200, json={"access_token": "fresh", "refresh_token": "rotated", "expires_in": 3600} + ) + response: Final = await endpoints.exchange_token_with_server( + _token_request({"Authorization": f"Bearer {_oauth_identity_jwt(key, scope='litellm_proxy_admin')}"}), + server, "refresh_token", None, None, identity, None, None, refresh_token="saved-refresh", + ) + assert response.status_code == 200 + assert parse_qs(upstream.calls[0].request.content.decode())["client_id"] == [identity] + persisted: Final = await db.get_user_oauth_credential(proxy_server.prisma_client, "jwt-owner", server.server_id) + assert persisted is not None and persisted["cimd_client_id"] == identity + assert persisted["refresh_token"] == "rotated" + assert table.upsert.await_count == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("state", [ + "matching", "unicode_grant", "foreign_refresh", "missing_refresh", "missing_grant", "database_missing", + "database_outage", "static_client", "anonymous", +]) +async def test_saved_cimd_refresh_identity_is_bound_to_the_callers_stored_grant( + monkeypatch: pytest.MonkeyPatch, state: str, +) -> None: + from types import SimpleNamespace + + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import db + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.setenv("LITELLM_SALT_KEY", "saved-cimd-owner-test-salt") + server: Final = _cimd_oauth_server() + identity: Final = "https://original-gateway.example.com/oauth/client-metadata.json" + grant: Final = "alice-refresh-\u00e9" if state == "unicode_grant" else "alice-refresh" + database: Final = MagicMock() + table: Final = database.db.litellm_mcpusercredentials + table.find_unique = AsyncMock(return_value=None) + table.upsert = AsyncMock() + monkeypatch.setattr(proxy_server, "prisma_client", database) + await db.store_user_oauth_credential( + database, "alice", server.server_id, "expired", + refresh_token=None if state == "missing_refresh" else grant, cimd_client_id=identity, + ) + table.find_unique.reset_mock() + table.find_unique.return_value = ( + None if state == "missing_grant" else SimpleNamespace( + credential_b64=table.upsert.call_args.kwargs["data"]["create"]["credential_b64"] + ) + ) + if state == "database_missing": + monkeypatch.setattr(proxy_server, "prisma_client", None) + if state == "database_outage": + table.find_unique.side_effect = RuntimeError("database unavailable") + if state == "static_client": + server.client_id = "configured-client" + resolved: Final = await endpoints._saved_cimd_refresh_client_id( + server, None if state == "anonymous" else "alice", + "foreign-refresh" if state == "foreign_refresh" else grant, + ) + assert resolved == (identity if state in ("matching", "unicode_grant") else None) + if state in ("anonymous", "static_client", "database_missing"): + table.find_unique.assert_not_awaited() + else: + table.find_unique.assert_awaited_once_with(where={"user_id_server_id": {"user_id": "alice", "server_id": server.server_id}}) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py index a59b02ec01d..dd8dddf106e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py @@ -7,12 +7,18 @@ well as the decline paths used in tool-bridge mode or when the downstream client lacks the requested elicitation capability. """ +import asyncio from types import SimpleNamespace from unittest.mock import AsyncMock import pytest from mcp.types import ( + ClientCapabilities, + ElicitationCapability, + FormElicitationCapability, + UrlElicitationCapability, + REQUEST_TIMEOUT, ElicitRequestFormParams, ElicitRequestURLParams, ElicitResult, @@ -43,23 +49,21 @@ def _url_params(message: str = "please authorize") -> ElicitRequestURLParams: ) -def _caps(*, url=True, form=True) -> SimpleNamespace: - elicit = SimpleNamespace( - url=object() if url else None, - form=object() if form else None, - ) - return SimpleNamespace(elicitation=elicit) +def _caps(*, url=True, form=True) -> ClientCapabilities: + return ClientCapabilities(elicitation=ElicitationCapability( + url=UrlElicitationCapability() if url else None, + form=FormElicitationCapability() if form else None, + )) class TestHandleElicitationRequest: - async def test_should_decline_when_no_downstream_session(self): + async def test_should_error_when_no_downstream_session(self): result = await handle_elicitation_request( context=SimpleNamespace(), params=_form_params(), downstream_session=None, ) - assert isinstance(result, ElicitResult) - assert result.action == "decline" + assert isinstance(result, ErrorData) async def test_should_relay_to_downstream_when_session_present(self): accepted = ElicitResult(action="accept", content={"name": "ada"}) @@ -69,7 +73,7 @@ class TestHandleElicitationRequest: context=SimpleNamespace(), params=_form_params(), downstream_session=session, - downstream_capabilities=None, + downstream_capabilities=_caps(), ) assert result is accepted @@ -85,21 +89,6 @@ class TestHandleElicitationRequest: assert isinstance(result, ErrorData) assert "not available" in result.message - async def test_should_return_error_data_on_unexpected_failure(self): - class _ExplodingParams: - mode = "form" - - @property - def message(self): - raise RuntimeError("boom") - - result = await handle_elicitation_request( - context=SimpleNamespace(), - params=_ExplodingParams(), - downstream_session=None, - ) - assert isinstance(result, ErrorData) - assert "boom" in result.message class TestRelayElicitationToDownstream: @@ -136,25 +125,17 @@ class TestRelayElicitationToDownstream: assert kwargs["url"] == "https://example.com/oauth" assert kwargs["elicitation_id"] == "elc-1" - async def test_should_use_generic_elicit_for_unknown_param_type(self): - accepted = ElicitResult(action="accept") - session = SimpleNamespace(elicit=AsyncMock(return_value=accepted)) - - # A bare params object that is neither Form nor URL params triggers - # the generic fallback path. - params = SimpleNamespace(mode="form", message="hi", requested_schema={}) - result = await _relay_elicitation_to_downstream( - params=params, - downstream_session=session, - downstream_capabilities=None, - ) - - assert result is accepted - session.elicit.assert_awaited_once() - - async def test_should_decline_when_elicitation_unsupported(self): + async def test_should_reject_invalid_elicitation_parameters(self): session = SimpleNamespace(elicit_form=AsyncMock()) - caps = SimpleNamespace(elicitation=None) + result = await _relay_elicitation_to_downstream( + params=SimpleNamespace(mode="form"), downstream_session=session, downstream_capabilities=_caps(), + ) + assert isinstance(result, ErrorData) + session.elicit_form.assert_not_awaited() + + async def test_should_error_when_elicitation_unsupported(self): + session = SimpleNamespace(elicit_form=AsyncMock()) + caps = ClientCapabilities() result = await _relay_elicitation_to_downstream( params=_form_params(), @@ -162,11 +143,10 @@ class TestRelayElicitationToDownstream: downstream_capabilities=caps, ) - assert isinstance(result, ElicitResult) - assert result.action == "decline" + assert isinstance(result, ErrorData) session.elicit_form.assert_not_awaited() - async def test_should_decline_url_mode_when_url_unsupported(self): + async def test_should_error_url_mode_when_url_unsupported(self): session = SimpleNamespace(elicit_url=AsyncMock()) result = await _relay_elicitation_to_downstream( @@ -175,11 +155,10 @@ class TestRelayElicitationToDownstream: downstream_capabilities=_caps(url=False, form=True), ) - assert isinstance(result, ElicitResult) - assert result.action == "decline" + assert isinstance(result, ErrorData) session.elicit_url.assert_not_awaited() - async def test_should_decline_form_mode_when_form_unsupported(self): + async def test_should_error_form_mode_when_form_unsupported(self): session = SimpleNamespace(elicit_form=AsyncMock()) result = await _relay_elicitation_to_downstream( @@ -188,24 +167,112 @@ class TestRelayElicitationToDownstream: downstream_capabilities=_caps(url=True, form=False), ) - assert isinstance(result, ElicitResult) - assert result.action == "decline" + assert isinstance(result, ErrorData) session.elicit_form.assert_not_awaited() - async def test_should_decline_when_downstream_relay_raises(self): + async def test_should_error_when_downstream_relay_raises(self): session = SimpleNamespace( elicit_form=AsyncMock(side_effect=RuntimeError("transport closed")) ) - result = await _relay_elicitation_to_downstream( + result = await handle_elicitation_request( + context=None, params=_form_params(), downstream_session=session, downstream_capabilities=_caps(form=True), ) - assert isinstance(result, ElicitResult) - assert result.action == "decline" + assert isinstance(result, ErrorData) if __name__ == "__main__": pytest.main([__file__, "-v"]) + + +@pytest.mark.asyncio +async def test_relay_failure_is_not_user_decline(): + session = SimpleNamespace(elicit_form=AsyncMock(side_effect=RuntimeError("private upstream credential"))) + result = await handle_elicitation_request( + context=None, params=_form_params(), downstream_session=session, downstream_capabilities=_caps(), + ) + assert isinstance(result, ErrorData), "A failed relay must not claim that the user declined" + assert "private upstream credential" not in result.message + + +@pytest.mark.asyncio +async def test_unknown_client_capabilities_prevent_relay(): + session = SimpleNamespace(elicit_form=AsyncMock(return_value=ElicitResult(action="accept"))) + result = await handle_elicitation_request( + context=None, params=_form_params(), downstream_session=session, downstream_capabilities=None, + ) + assert isinstance(result, ErrorData), "Unknown capabilities must fail explicitly" + session.elicit_form.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["form", "url"]) +@pytest.mark.parametrize("action", ["accept", "decline", "cancel"]) +async def test_relay_preserves_response_and_request_correlation(mode, action): + response = ElicitResult(action=action) + request = AsyncMock(return_value=response) + session = SimpleNamespace(elicit_form=request, elicit_url=request) + params = _form_params() if mode == "form" else _url_params() + result = await handle_elicitation_request( + context=SimpleNamespace(request_id="upstream-id"), params=params, + downstream_session=session, downstream_capabilities=_caps(), related_request_id=0, + ) + assert result is response + assert request.await_args.kwargs["related_request_id"] == 0 + if mode == "url": + assert request.await_args.kwargs["url"] == params.url + assert request.await_args.kwargs["elicitation_id"] == params.elicitation_id + else: + assert request.await_args.kwargs["requested_schema"] == params.requested_schema + + +@pytest.mark.asyncio +async def test_expired_relay_deadline_returns_explicit_timeout(): + request = AsyncMock() + result = await handle_elicitation_request( + context=None, params=_form_params(), downstream_session=SimpleNamespace(elicit_form=request), + downstream_capabilities=_caps(), timeout=0, + ) + assert isinstance(result, ErrorData) + assert result.code == REQUEST_TIMEOUT + assert "timed out" in result.message + request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_relay_cancellation_releases_downstream_waiter(): + started = asyncio.Event() + finished = asyncio.Event() + + async def wait_for_user(**kwargs): + started.set() + try: + await asyncio.Event().wait() + finally: + finished.set() + + task = asyncio.create_task(handle_elicitation_request( + context=None, params=_form_params(), downstream_session=SimpleNamespace(elicit_form=wait_for_user), + downstream_capabilities=_caps(), related_request_id="tool-call", + )) + await started.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert finished.is_set() + + +@pytest.mark.asyncio +async def test_legacy_empty_elicitation_capability_supports_form(): + accepted = ElicitResult(action="accept") + session = SimpleNamespace(elicit_form=AsyncMock(return_value=accepted)) + result = await handle_elicitation_request( + context=None, params=_form_params(), downstream_session=session, + downstream_capabilities=ClientCapabilities(elicitation=ElicitationCapability()), + related_request_id="legacy-call", + ) + assert result is accepted diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index 23defa6f513..5ae3c595c53 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -464,7 +464,7 @@ async def test_sse_mcp_handler_mock(): mock_sse = MagicMock() mock_sse.connect_sse.side_effect = connect_sse - run = AsyncMock() + serve = AsyncMock() # Mock scope, receive, send with proper ASGI scope format mock_scope = { @@ -489,7 +489,7 @@ async def test_sse_mcp_handler_mock(): ) with ( - patch("litellm.proxy._experimental.mcp_server.server.serve_loop", run), + patch("litellm.proxy._experimental.mcp_server.server.serve_loop", serve), patch( "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True, @@ -505,13 +505,21 @@ async def test_sse_mcp_handler_mock(): patch( "litellm.proxy._experimental.mcp_server.server.set_auth_context", ), + patch( + "litellm.proxy._experimental.mcp_server.server._raise_preemptive_401_for_unauthenticated_servers", + new=AsyncMock(), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._check_passthrough_upstream_auth", + new=AsyncMock(), + ), ): from litellm.proxy._experimental.mcp_server.server import handle_sse_mcp # Call the handler await handle_sse_mcp(mock_scope, mock_receive, mock_send) - assert run.await_args.args[1:3] == (read_stream, write_stream) + assert serve.await_args.args[1:3] == (read_stream, write_stream) assert mock_sse.connect_sse.call_args.args[0]["path"] == "/mcp/sse" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 5bd4ff5d894..c7b477a6cc6 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -155,18 +155,26 @@ async def test_elicitation_callback_keeps_initiating_session(): from litellm.proxy._experimental.mcp_server import server as legacy_server from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_elicitation_callback - initiating = MagicMock() - replacement = MagicMock() - recorder = AsyncMock() + from types import SimpleNamespace + from mcp.types import ClientCapabilities, ElicitationCapability, FormElicitationCapability, ElicitRequestFormParams, ElicitResult + from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var + + capabilities = ClientCapabilities(elicitation=ElicitationCapability(form=FormElicitationCapability())) + accepted = ElicitResult(action="accept") + request = AsyncMock(return_value=accepted) + initiating = SimpleNamespace(client_params=SimpleNamespace(capabilities=capabilities), elicit_form=request) token = legacy_server.active_mcp_session_var.set(initiating) + request_token = active_mcp_request_ctx_var.set(SimpleNamespace(session=initiating, request_id="initiating-call")) try: callback = _create_elicitation_callback() - legacy_server.active_mcp_session_var.set(replacement) - with patch("litellm.proxy._experimental.mcp_server.elicitation_handler.handle_elicitation_request", recorder): - await callback(None, None) - assert recorder.await_args.kwargs["downstream_session"] is initiating - assert recorder.await_args.kwargs["downstream_capabilities"] is initiating.capabilities + legacy_server.active_mcp_session_var.set(SimpleNamespace()) + active_mcp_request_ctx_var.set(None) + capabilities.elicitation = None + result = await callback(None, ElicitRequestFormParams(message="Confirm", requested_schema={"type":"object"})) + assert result is accepted + assert request.await_args.kwargs["related_request_id"] == "initiating-call" finally: + active_mcp_request_ctx_var.reset(request_token) legacy_server.active_mcp_session_var.reset(token) @@ -1461,6 +1469,7 @@ class TestMCPServerManager: manager = MCPServerManager() metadata = MCPOAuthMetadata( + client_id_metadata_document_supported=True, authorization_url="https://attacker.example.com/authorize", token_url="https://attacker.example.com/token", scopes=["read", "admin"], @@ -1478,6 +1487,8 @@ class TestMCPServerManager: assert server.token_url is None assert server.scopes == ["read", "admin"] + assert server.client_id_metadata_document_supported is False + @pytest.mark.asyncio async def test_load_servers_from_config_fills_token_url_when_metadata_corroborates_manual_authorization_url(self): """Corroborated metadata keeps the self-heal on the config path: when the discovered document @@ -1487,6 +1498,7 @@ class TestMCPServerManager: manager = MCPServerManager() metadata = MCPOAuthMetadata( + client_id_metadata_document_supported=True, authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", scopes=["read", "admin"], @@ -1503,6 +1515,8 @@ class TestMCPServerManager: assert server.token_url == "https://idp.example.com/token" assert server.scopes == ["read", "admin"] + assert server.client_id_metadata_document_supported is True + @pytest.mark.asyncio @pytest.mark.parametrize("blank_authorization_url", ["", " "]) async def test_load_servers_from_config_blank_authorization_url_is_not_a_pin(self, blank_authorization_url): @@ -2535,6 +2549,7 @@ class TestMCPServerManager: ) metadata = MCPOAuthMetadata( + client_id_metadata_document_supported=True, authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", @@ -2548,6 +2563,8 @@ class TestMCPServerManager: assert built.registration_url == "https://idp.example.com/register" assert built.scopes == ["read"] + assert built.client_id_metadata_document_supported is True + @pytest.mark.asyncio async def test_build_from_table_uses_issuer_anchored_endpoints_when_issuer_configured(self): """When an admin configures an issuer, the build takes its endpoints from the issuer-anchored @@ -4507,6 +4524,7 @@ class TestMCPServerManager: "authorization_endpoint": "https://idp.example.com/authorize", "token_endpoint": "https://idp.example.com/token", "scopes_supported": ["read", "write"], + "client_id_metadata_document_supported": True, }, ) mock_client = MagicMock() @@ -4522,6 +4540,8 @@ class TestMCPServerManager: assert result.token_url == "https://idp.example.com/token" assert result.scopes == ["read", "write"] + assert result.client_id_metadata_document_supported is True + @pytest.mark.asyncio async def test_fetch_single_authorization_server_metadata_rejects_issuer_mismatch(self): """RFC 8414 §3.3 fail-closed: a document self-attesting a DIFFERENT issuer than the one it was @@ -19108,3 +19128,150 @@ def test_discovery_keys_bind_static_auth_to_caller_and_configuration() -> None: manager._discovery_key(updated, first, None, None, None, None), ) assert len(set(keys)) == 3 + + +@pytest.mark.parametrize("advertised", [False, True]) +def test_untrusted_metadata_cannot_enable_cimd_without_matching_authorization_endpoint(advertised): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _restrict_discovery_to_corroborated_authorization_server, + ) + + metadata = MCPOAuthMetadata(scopes=["read"], client_id_metadata_document_supported=advertised) + result = _restrict_discovery_to_corroborated_authorization_server( + metadata, "https://trusted.example.com/authorize", "server", False + ) + assert result is not None + assert result.client_id_metadata_document_supported is False + assert result.scopes == ["read"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("source", ["config", "database"]) +@pytest.mark.parametrize("advertised", [True, False]) +@pytest.mark.parametrize("startup", [True, False]) +async def test_manual_oauth_endpoints_discover_client_metadata_once( + source: str, advertised: bool, startup: bool, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch +) -> None: + from starlette.requests import Request + + from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module + from litellm.proxy._experimental.mcp_server.oauth_utils import get_cimd_client_id + + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "1" if startup else "0") + await _mock_oauth_discovery(respx_mock, monkeypatch, server_url="https://up.example.com/mcp", scopes=["read"]) + metadata: Final = respx_mock.get("https://up.example.com/.well-known/oauth-authorization-server").respond( + json={ + "issuer": "https://up.example.com", + "authorization_endpoint": "https://up.example.com/authorize", + "token_endpoint": "https://up.example.com/token", + "client_id_metadata_document_supported": advertised, + } + ) + manager: Final = MCPServerManager() + monkeypatch.setattr(manager_module, "global_mcp_server_manager", manager) + configured: Final[MCPServer] + if source == "config": + await manager.load_servers_from_config({"manual": { + "url": "https://up.example.com/mcp", "transport": "http", "auth_type": "oauth2", + "oauth2_flow": "authorization_code", "authorization_url": "https://up.example.com/authorize", + "token_url": "https://up.example.com/token", "scopes": ["read"], + }}) + configured = next(iter(manager.config_mcp_servers.values())) + else: + row: Final = LiteLLM_MCPServerTable( + server_id="manual", alias="manual", url="https://up.example.com/mcp", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + authorization_url="https://up.example.com/authorize", token_url="https://up.example.com/token", + credentials={"scopes": ["read"]}, created_at=datetime.now(), updated_at=datetime.now(), + ) + built: Final = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) + manager.registry[built.server_id] = built + configured = built + async with manager.catalog.operation(): + response: Final = await discoverable_endpoints.register_client_with_server( + Request({"type": "http", "scheme": "https", "server": ("gateway.example.com", 443), + "path": "/register", "root_path": "", "headers": [], "query_string": b""}), + configured, "Gateway", None, None, None, + ) + body: Final = json.loads(response.body) if hasattr(response, "body") else response + assert body["client_id"] == ( + "https://gateway.example.com/oauth/client-metadata.json" if advertised else configured.server_name + ) + resolved: Final = await manager.ensure_oauth_metadata_discovered(configured) + again: Final = await manager.ensure_oauth_metadata_discovered(resolved) + assert metadata.call_count == 1 + assert resolved.client_id_metadata_document_supported is advertised + assert again == resolved + assert get_cimd_client_id(resolved) == ( + "https://gateway.example.com/oauth/client-metadata.json" if advertised else None + ) + assert resolved.effective_authorization_url == "https://up.example.com/authorize" + assert resolved.effective_token_url == "https://up.example.com/token" + + +@pytest.mark.asyncio +async def test_optional_client_metadata_discovery_failure_preserves_manual_endpoints( + respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "0") + await _mock_oauth_discovery(respx_mock, monkeypatch, server_url="https://up.example.com/mcp", scopes=["read"]) + respx_mock.get("https://up.example.com/.well-known/oauth-authorization-server").respond(503) + respx_mock.route().respond(404) + manager: Final = MCPServerManager() + await manager.load_servers_from_config({"manual": { + "url": "https://up.example.com/mcp", "transport": "http", "auth_type": "oauth2", + "oauth2_flow": "authorization_code", "authorization_url": "https://up.example.com/authorize", + "token_url": "https://up.example.com/token", "scopes": ["read"], + }}) + configured: Final = next(iter(manager.config_mcp_servers.values())) + resolved: Final = await manager.ensure_oauth_metadata_discovered(configured) + attempts: Final = len(respx_mock.calls) + retry: Final = await manager.ensure_oauth_metadata_discovered(resolved) + assert attempts > 0 + assert len(respx_mock.calls) == attempts + assert retry == resolved + assert resolved.client_id_metadata_document_supported is None + assert resolved.effective_authorization_url == "https://up.example.com/authorize" + assert resolved.effective_token_url == "https://up.example.com/token" + assert manager.oauth_discovery_slot(resolved.server_id) is not None + + +@pytest.mark.parametrize("capability", [True, False]) +@pytest.mark.parametrize("rebuild", ["same", "repointed", "fresh_discovery", "anchored"]) +def test_oauth_rebuild_retains_only_corroborated_cimd_capability(capability: bool, rebuild: str) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import carry_forward_resolved_oauth_endpoints + + previous: Final = MCPServer( + server_id="cimd-rebuild", name="cimd_rebuild", url="https://mcp.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + client_id_metadata_document_supported=capability, + ) + rebuilt: Final = previous.model_copy(update={ + "client_id_metadata_document_supported": not capability if rebuild == "fresh_discovery" else None, + "authorization_url": "https://changed.example.com/authorize" if rebuild == "repointed" else previous.authorization_url, + "token_url": None, + "issuer": "https://idp.example.com" if rebuild == "anchored" else None, + "issuer_is_anchored": rebuild == "anchored", + }) + carry_forward_resolved_oauth_endpoints(rebuilt, previous) + expected: Final = not capability if rebuild == "fresh_discovery" else None if rebuild in ("repointed", "anchored") else capability + assert rebuilt.client_id_metadata_document_supported is expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize("endpoint", ["authorization_url", "token_url"]) +async def test_repeated_stale_discovery_uses_current_callers_endpoint(endpoint: str) -> None: + manager: Final = MCPServerManager() + original: Final = MCPServer( + server_id="partial-replacement", name="replacement", url="https://old.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + ) + replacement: Final = original.model_copy(update={endpoint: "https://new.example.com/oauth"}) + manager.registry[original.server_id] = replacement + resolved: Final = await manager._rejoin_oauth_metadata_discovery( + original, needed_endpoint=lambda server: getattr(server, endpoint), retry_stale=False, + ) + assert resolved is replacement diff --git a/tests/unit/proxy/auth/test_auth_utils.py b/tests/unit/proxy/auth/test_auth_utils.py index 130d29a9560..e2cec1cc5a2 100644 --- a/tests/unit/proxy/auth/test_auth_utils.py +++ b/tests/unit/proxy/auth/test_auth_utils.py @@ -4,6 +4,7 @@ Unit tests for auth_utils functions related to rate limiting and customer ID ext import base64 import logging +from collections.abc import Callable from typing import Final, Optional from unittest.mock import MagicMock, patch @@ -2877,6 +2878,40 @@ class TestIsRequestBodySafeBlocksClaudePlatformWorkspaceOverride: is True ) + +class TestIsRequestBodySafeBlocksFireworksForwardUserId: + @pytest.mark.parametrize("value", [True, False]) + @pytest.mark.parametrize( + "body_for", + [ + pytest.param(lambda value: {"fireworks_forward_user_id": value}, id="root"), + pytest.param(lambda value: {"extra_body": {"fireworks_forward_user_id": value}}, id="extra_body"), + pytest.param(lambda value: {"metadata": {"fireworks_forward_user_id": value}}, id="metadata"), + ], + ) + def test_fireworks_forward_user_id_in_request_body_is_rejected( + self, body_for: Callable[[bool], dict[str, object]], value: bool + ) -> None: + with pytest.raises(ValueError, match="fireworks_forward_user_id"): + is_request_body_safe( + request_body={"model": "fireworks-model", "user": "someone-else", **body_for(value)}, + general_settings={}, + llm_router=None, + model="fireworks-model", + ) + + def test_admin_opt_in_proxy_wide_allows_fireworks_forward_user_id(self) -> None: + assert ( + is_request_body_safe( + request_body={"model": "fireworks-model", "fireworks_forward_user_id": False}, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="fireworks-model", + ) + is True + ) + + class TestIsRequestBodySafeBlocksRustOptIn: """``rust`` hands the whole call to the Rust core, which signs and sends with its own HTTP client rather than the one the deployment configured, and diff --git a/tests/unit/proxy/auth/test_user_api_key_auth.py b/tests/unit/proxy/auth/test_user_api_key_auth.py index a58599d28e9..9cc8552efd5 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth.py @@ -5,11 +5,12 @@ import litellm.proxy import litellm.proxy.proxy_server -from typing import Dict, List, Optional -from unittest.mock import MagicMock, patch, AsyncMock +from typing import Dict, Final, List, Optional +from unittest.mock import AsyncMock, MagicMock, patch import pytest from starlette.datastructures import URL +from starlette.types import Message from litellm._logging import verbose_proxy_logger import logging import litellm @@ -1811,3 +1812,91 @@ def test_mapped_key_jwt_falls_through_to_the_shared_user_budget_attach(): "the mapped-key branch returns before the shared virtual-key checks, so the " "user's per-model budget is never attached and never enforced" ) + + +def _rejected_websocket(sent: list[Message], bearer: str) -> WebSocket: + async def receive() -> Message: + return {"type": "websocket.connect"} + + async def send(message: Message) -> None: + sent.append(message) + + return WebSocket( + { + "type": "websocket", + "path": "/v1/responses", + "query_string": b"model=gpt-5.4", + "headers": [(b"authorization", f"Bearer {bearer}".encode())], + }, + receive, + send, + ) + + +def _serve_virtual_key(monkeypatch: pytest.MonkeyPatch, user_key: str, token: UserAPIKeyAuth) -> None: + from litellm.proxy.proxy_server import hash_token, user_api_key_cache + + user_api_key_cache.set_cache(key=hash_token(user_key), value=token) + monkeypatch.setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) + monkeypatch.setattr(litellm.proxy.proxy_server, "master_key", MASTER_KEY) + monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", "connected") + monkeypatch.setattr(litellm, "log_client_error_tracebacks", False) + monkeypatch.setattr(verbose_proxy_logger, "propagate", True) + + +@pytest.mark.asyncio +async def test_user_api_key_auth_websocket_logs_a_model_access_denial_as_one_warning_without_a_traceback( + caplog: pytest.LogCaptureFixture, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_websocket + from litellm.proxy.proxy_server import hash_token + + user_key: Final = "sk-websocket-key-limited-to-mini" + _serve_virtual_key( + monkeypatch, + user_key, + UserAPIKeyAuth(token=hash_token(user_key), models=["gpt-5.4-mini"]), + ) + sent: Final[list[Message]] = [] + with ( + caplog.at_level(logging.WARNING, logger=verbose_proxy_logger.name), + pytest.raises(HTTPException) as rejection, + ): + await user_api_key_auth_websocket(_rejected_websocket(sent, user_key)) + + assert rejection.value.status_code == 403 + assert sent == [{"type": "websocket.close", "code": status.WS_1008_POLICY_VIOLATION, "reason": ""}] + proxy_records: Final = [record for record in caplog.records if record.name == verbose_proxy_logger.name] + assert [record.getMessage() for record in proxy_records if record.exc_info is not None] == [] + assert [record.getMessage() for record in proxy_records if record.levelno == logging.WARNING] == [ + "key not allowed to access model. This key can only access models=['gpt-5.4-mini']. Tried to access gpt-5.4" + ] + + +@pytest.mark.asyncio +async def test_user_api_key_auth_websocket_rejection_adds_no_traceback_for_other_auth_errors( + caplog: pytest.LogCaptureFixture, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_websocket + from litellm.proxy.proxy_server import hash_token + + user_key: Final = "sk-websocket-key-expired-in-2020" + expired_in_2020: Final = "2020-01-01T00:00:00+00:00" + _serve_virtual_key( + monkeypatch, + user_key, + UserAPIKeyAuth(token=hash_token(user_key), expires=expired_in_2020), + ) + sent: Final[list[Message]] = [] + with ( + caplog.at_level(logging.DEBUG, logger=verbose_proxy_logger.name), + pytest.raises(HTTPException) as rejection, + ): + await user_api_key_auth_websocket(_rejected_websocket(sent, user_key)) + + assert rejection.value.status_code == 403 + assert "expired key" in str(rejection.value.detail).lower() + assert sent == [{"type": "websocket.close", "code": status.WS_1008_POLICY_VIOLATION, "reason": ""}] + proxy_records: Final = [record for record in caplog.records if record.name == verbose_proxy_logger.name] + assert [record.getMessage() for record in proxy_records if record.exc_info is not None] == [] + assert [record.levelno for record in proxy_records if record.levelno >= logging.WARNING] == [logging.ERROR] diff --git a/tests/unit/proxy/client/cli/test_pi.py b/tests/unit/proxy/client/cli/test_pi.py index 6bb49566580..ff4dc25cd0c 100644 --- a/tests/unit/proxy/client/cli/test_pi.py +++ b/tests/unit/proxy/client/cli/test_pi.py @@ -195,6 +195,7 @@ class TestProviderBlock: assert block == { "baseUrl": "http://localhost:4000/v1", "api": "openai-completions", + "compat": {"supportsStore": False, "supportsLongCacheRetention": False}, "apiKey": "$LITELLM_PROXY_API_KEY", "models": [{"id": "m-1"}, {"id": "m-2"}], } diff --git a/tests/unit/proxy/common_utils/test_check_batch_cost.py b/tests/unit/proxy/common_utils/test_check_batch_cost.py index 87df82c01d0..58d6a44e4a4 100644 --- a/tests/unit/proxy/common_utils/test_check_batch_cost.py +++ b/tests/unit/proxy/common_utils/test_check_batch_cost.py @@ -23,6 +23,32 @@ _CLAIM_UNIFIED_BATCH_ID = "dW5pZmllZF9iYXRjaF9pZA==" _CLAIM_OUTPUT_FILE_ID = "file-output-123" +def test_batch_output_file_object_derives_metadata(): + from litellm_enterprise.proxy.hooks.managed_files import _batch_output_file_object + + output_file = _batch_output_file_object( + unified_file_id="unified-output", + raw_file_id="s3://bucket/path/to/output.jsonl.out", + size_bytes=4321, + fallback=False, + ) + provider_file = _batch_output_file_object( + unified_file_id="unified-provider", + raw_file_id="file-abc", + size_bytes=0, + fallback=False, + ) + + assert output_file.id == "unified-output" + assert output_file.object == "file" + assert output_file.purpose == "batch_output" + assert output_file.filename == "output.jsonl.out" + assert output_file.bytes == 4321 + assert output_file.status == "processed" + assert provider_file.filename == "file-abc" + assert provider_file.bytes == 0 + + def _batch_cost_result( cost: float, usage: dict, @@ -1015,9 +1041,11 @@ class TestCheckBatchCost: ) mock_llm_router.aretrieve_batch = AsyncMock(return_value=response) - mock_hook = MagicMock() + from litellm.proxy.openai_files_endpoints.common_utils import ManagedBatchOutputFileWriter + + mock_hook = MagicMock(spec=ManagedBatchOutputFileWriter) mock_hook.get_unified_output_file_id.side_effect = [unified_error_file_id] - mock_hook.store_unified_file_id = AsyncMock() + mock_hook.store_batch_output_file = AsyncMock() check_batch_cost_instance.proxy_logging_obj.get_proxy_hook.return_value = mock_hook await check_batch_cost_instance.check_batch_cost() @@ -1027,14 +1055,19 @@ class TestCheckBatchCost: model_id="model-123", model_name="gpt-5-batch", ) - stored = { - next(iter(c.kwargs["model_mappings"].values())): c.kwargs["file_id"] - for c in mock_hook.store_unified_file_id.call_args_list - } - assert stored == {raw_error_file_id: unified_error_file_id} - for store_call in mock_hook.store_unified_file_id.call_args_list: - assert store_call.kwargs["user_api_key_dict"].user_id == "user-1" - assert store_call.kwargs["user_api_key_dict"].team_id == "team-1" + mock_hook.store_batch_output_file.assert_awaited_once_with( + unified_file_id=unified_error_file_id, + provider_file_id=raw_error_file_id, + model_id="model-123", + model_name="gpt-5-batch", + owner=mock_hook.store_batch_output_file.await_args.kwargs["owner"], + litellm_parent_otel_span=None, + size_bytes=None, + fetch_provider_details=True, + ) + owner = mock_hook.store_batch_output_file.await_args.kwargs["owner"] + assert owner.user_id == "user-1" + assert owner.team_id == "team-1" assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1 update_call = mock_prisma_client.db.litellm_managedobjecttable.update.call_args @@ -1506,6 +1539,8 @@ class TestCheckBatchCost: Without this, GET /batches/{id} returns a raw file ID that cannot be routed through the proxy, causing API_KEY errors when clients call GET /files/{id}/content. """ + from litellm.proxy.openai_files_endpoints.common_utils import ManagedBatchOutputFileWriter + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) @@ -1540,12 +1575,12 @@ class TestCheckBatchCost: mock_deployment.model_info.model_dump.return_value = {} mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment) - mock_hook = MagicMock() + mock_hook = MagicMock(spec=ManagedBatchOutputFileWriter) mock_hook.get_unified_output_file_id.side_effect = [ fake_managed_output_id, fake_managed_error_id, ] - mock_hook.store_unified_file_id = AsyncMock() + mock_hook.store_batch_output_file = AsyncMock() check_batch_cost_instance.proxy_logging_obj.get_proxy_hook.return_value = mock_hook mock_file_content = MagicMock() @@ -1608,16 +1643,25 @@ class TestCheckBatchCost: model_id="model-123", model_name="gpt-5-batch", ) - assert mock_hook.store_unified_file_id.await_count == 2 - # {raw_file_id: managed_file_id} for each store call + assert mock_hook.store_batch_output_file.await_count == 2 + assert { + call.kwargs["model_name"] + for call in mock_hook.store_batch_output_file.await_args_list + } == {"gpt-5-batch"} stored = { - next(iter(c[1]["model_mappings"].values())): c[1]["file_id"] - for c in mock_hook.store_unified_file_id.call_args_list + c.kwargs["provider_file_id"]: c.kwargs["unified_file_id"] + for c in mock_hook.store_batch_output_file.call_args_list } assert stored == { raw_output_file_id: fake_managed_output_id, raw_error_file_id: fake_managed_error_id, } + stored_sizes = { + c.kwargs["provider_file_id"]: c.kwargs["size_bytes"] + for c in mock_hook.store_batch_output_file.call_args_list + } + assert stored_sizes[raw_output_file_id] == len(mock_file_content.content) + assert stored_sizes[raw_error_file_id] is None assert mock_response.output_file_id == fake_managed_output_id assert mock_response.error_file_id == fake_managed_error_id @@ -2192,10 +2236,10 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: from litellm_enterprise.proxy.common_utils.check_batch_cost import ( CheckBatchCost, ) + + from enterprise.litellm_enterprise.proxy.hooks.managed_files import PROXY_LiteLLMManagedFiles + from litellm.proxy.openai_files_endpoints.common_utils import ManagedBatchOutputFileWriter from litellm.types.utils import LiteLLMBatch - from enterprise.litellm_enterprise.proxy.hooks.managed_files import ( - PROXY_LiteLLMManagedFiles, - ) router = MagicMock() router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"}) @@ -2206,13 +2250,13 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: deployment.model_info.model_dump.return_value = {} router.get_deployment = MagicMock(return_value=deployment) - hook = MagicMock() + hook = MagicMock(spec=ManagedBatchOutputFileWriter) hook.get_unified_output_file_id = lambda output_file_id, model_id, model_name: ( PROXY_LiteLLMManagedFiles.get_unified_output_file_id( None, output_file_id=output_file_id, model_id=model_id, model_name=model_name ) ) - hook.store_unified_file_id = AsyncMock() + hook.store_batch_output_file = AsyncMock() proxy_logging_obj = MagicMock() proxy_logging_obj.get_proxy_hook.return_value = hook diff --git a/tests/unit/proxy/common_utils/test_check_responses_cost.py b/tests/unit/proxy/common_utils/test_check_responses_cost.py index e806e9a3394..8ba3f0b000a 100644 --- a/tests/unit/proxy/common_utils/test_check_responses_cost.py +++ b/tests/unit/proxy/common_utils/test_check_responses_cost.py @@ -3,7 +3,9 @@ Unit tests for CheckResponsesCost class """ import asyncio +from collections.abc import Mapping from datetime import datetime +from typing import Final from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest @@ -12,6 +14,19 @@ from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse +class _RecordingManagedObjectTable: + def __init__(self, rows: tuple[object, ...]) -> None: + self.rows: Final = rows + self.updates: tuple[tuple[Mapping[str, object], Mapping[str, object]], ...] = () + + async def find_many(self, **query: object) -> tuple[object, ...]: + return self.rows + + async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: + self.updates = (*self.updates, (where, data)) + return len(self.rows) + + class TestCheckResponsesCost: """Test suite for CheckResponsesCost class""" @@ -364,6 +379,264 @@ class TestCheckResponsesCost: # Stale cleanup still ran via _expire_stale_rows check_responses_cost_instance._expire_stale_rows.assert_called_once() + @pytest.mark.asyncio + async def test_check_responses_cost_marks_404_response_stale_expired( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """A provider 404 on a response polled through its deployment marks the row stale_expired.""" + import litellm + from litellm.responses.utils import ResponsesAPIRequestUtils + + encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id( + custom_llm_provider="openai", + model_id="deployment-404", + response_id="resp_upstream_404", + ) + + mock_job = MagicMock() + mock_job.unified_object_id = encoded_response_id + mock_job.created_by = "test-user" + mock_job.id = "job-404" + mock_job.file_object = {"model": "gpt-4o", "id": encoded_response_id} + + mock_llm_router.get_deployment.return_value = {"model_id": "deployment-404"} + mock_llm_router.aget_responses = AsyncMock( + side_effect=litellm.NotFoundError( + message="Response with id 'resp_upstream_404' not found.", model="gpt-5", llm_provider="openai" + ) + ) + + table = _RecordingManagedObjectTable(rows=(mock_job,)) + mock_prisma_client.db.litellm_managedobjecttable = table + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget: + await check_responses_cost_instance.check_responses_cost() + + mock_sdk_aget.assert_not_called() + mock_llm_router.get_deployment.assert_called_once_with(model_id="deployment-404") + assert mock_llm_router.aget_responses.call_args.kwargs["response_id"] == encoded_response_id + assert table.updates == (({"id": {"in": ["job-404"]}}, {"status": "stale_expired"}),) + + @pytest.mark.asyncio + async def test_check_responses_cost_marks_mapped_provider_404_stale_expired( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """The provider's own 404 body, mapped the way the GET path maps it, still expires the row.""" + import json + + import openai + + import litellm + from litellm.llms.openai.common_utils import OpenAIError + from litellm.responses.utils import ResponsesAPIRequestUtils + + encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id( + custom_llm_provider="openai", + model_id="deployment-404-mapped", + response_id="resp_upstream_404_mapped", + ) + provider_body = json.dumps( + { + "error": { + "message": "Response with id 'resp_upstream_404_mapped' not found.", + "type": "invalid_request_error", + "param": None, + "code": None, + } + } + ) + with pytest.raises(openai.APIStatusError) as mapped: + raise litellm.exception_type( + model="gpt-5.5", + custom_llm_provider="openai", + original_exception=OpenAIError(message=provider_body, status_code=404), + completion_kwargs={}, + extra_kwargs={}, + ) + assert mapped.value.status_code == 404 + + mock_job = MagicMock() + mock_job.unified_object_id = encoded_response_id + mock_job.created_by = "test-user" + mock_job.id = "job-404-mapped" + mock_job.file_object = {"model": "gpt-5.5", "id": encoded_response_id} + + mock_llm_router.get_deployment.return_value = {"model_id": "deployment-404-mapped"} + mock_llm_router.aget_responses = AsyncMock(side_effect=mapped.value) + + table = _RecordingManagedObjectTable(rows=(mock_job,)) + mock_prisma_client.db.litellm_managedobjecttable = table + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget: + await check_responses_cost_instance.check_responses_cost() + + mock_sdk_aget.assert_not_called() + assert table.updates == (({"id": {"in": ["job-404-mapped"]}}, {"status": "stale_expired"}),) + + @pytest.mark.asyncio + async def test_check_responses_cost_404_without_router_deployment_keeps_row_for_retry( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """A 404 while polling without a resolved deployment skips the row for retry.""" + import litellm + from litellm.responses.utils import ResponsesAPIRequestUtils + + encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id( + custom_llm_provider="openai", + model_id="deployment-gone", + response_id="resp_upstream_gone", + ) + + mock_job = MagicMock() + mock_job.unified_object_id = encoded_response_id + mock_job.created_by = "test-user" + mock_job.id = "job-404-fallback" + mock_job.file_object = {"model": "gpt-4o", "id": encoded_response_id} + + mock_llm_router.get_deployment.return_value = None + + table = _RecordingManagedObjectTable(rows=(mock_job,)) + mock_prisma_client.db.litellm_managedobjecttable = table + + with patch( + "litellm.aget_responses", + new_callable=AsyncMock, + side_effect=litellm.NotFoundError( + message="Response not found", model="gpt-5", llm_provider="openai" + ), + ) as mock_sdk_aget: + await check_responses_cost_instance.check_responses_cost() + + mock_sdk_aget.assert_awaited_once() + mock_llm_router.aget_responses.assert_not_called() + assert table.updates == () + + @pytest.mark.asyncio + async def test_check_responses_cost_404_not_naming_the_response_keeps_row_for_retry( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """A 404 that does not name the response (a gateway or a reconfigured deployment) is retried, not expired.""" + import litellm + from litellm.responses.utils import ResponsesAPIRequestUtils + + encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id( + custom_llm_provider="azure", + model_id="deployment-renamed", + response_id="resp_upstream_still_there", + ) + + mock_job = MagicMock() + mock_job.unified_object_id = encoded_response_id + mock_job.created_by = "test-user" + mock_job.id = "job-404-other" + mock_job.file_object = {"model": "azure-gpt-5", "id": encoded_response_id} + + mock_llm_router.get_deployment.return_value = {"model_id": "deployment-renamed"} + mock_llm_router.aget_responses = AsyncMock( + side_effect=litellm.NotFoundError( + message="Resource not found", model="gpt-5", llm_provider="azure" + ) + ) + + table = _RecordingManagedObjectTable(rows=(mock_job,)) + mock_prisma_client.db.litellm_managedobjecttable = table + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget: + await check_responses_cost_instance.check_responses_cost() + + mock_sdk_aget.assert_not_called() + assert mock_llm_router.aget_responses.call_args.kwargs["response_id"] == encoded_response_id + assert table.updates == () + + @pytest.mark.asyncio + async def test_check_responses_cost_deployment_lookup_error_skips_only_that_job( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """A deployment lookup that raises skips its own job and the cycle still records the next job.""" + from litellm.responses.utils import ResponsesAPIRequestUtils + + broken_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id( + custom_llm_provider="openai", model_id="deployment-broken", response_id="resp_upstream_broken" + ) + healthy_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id( + custom_llm_provider="openai", model_id="deployment-healthy", response_id="resp_upstream_healthy" + ) + + broken_job = MagicMock() + broken_job.unified_object_id = broken_response_id + broken_job.created_by = "test-user" + broken_job.id = "job-broken" + broken_job.file_object = {"model": "gpt-5.5", "id": broken_response_id} + + healthy_job = MagicMock() + healthy_job.unified_object_id = healthy_response_id + healthy_job.created_by = "test-user" + healthy_job.id = "job-healthy" + healthy_job.file_object = {"model": "gpt-5.5", "id": healthy_response_id} + + mock_llm_router.get_deployment.side_effect = [ + Exception("Model invalid format - "), + {"model_id": "deployment-healthy"}, + ] + mock_llm_router.aget_responses = AsyncMock( + return_value=ResponsesAPIResponse( + id=healthy_response_id, + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150), + ) + ) + + table = _RecordingManagedObjectTable(rows=(broken_job, healthy_job)) + mock_prisma_client.db.litellm_managedobjecttable = table + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget: + await check_responses_cost_instance.check_responses_cost() + + mock_sdk_aget.assert_not_called() + mock_llm_router.aget_responses.assert_awaited_once() + assert mock_llm_router.aget_responses.call_args.kwargs["response_id"] == healthy_response_id + assert table.updates == (({"id": {"in": ["job-healthy"]}}, {"status": "completed"}),) + + @pytest.mark.asyncio + async def test_check_responses_cost_non_404_error_keeps_row_for_retry( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """A non-404 provider error through a resolved deployment skips the job so it is retried next cycle.""" + import litellm + from litellm.responses.utils import ResponsesAPIRequestUtils + + encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id( + custom_llm_provider="openai", + model_id="deployment-500", + response_id="resp_upstream_500", + ) + + mock_job = MagicMock() + mock_job.unified_object_id = encoded_response_id + mock_job.created_by = "test-user" + mock_job.id = "job-500" + mock_job.file_object = {"model": "gpt-4o", "id": encoded_response_id} + + mock_llm_router.get_deployment.return_value = {"model_id": "deployment-500"} + mock_llm_router.aget_responses = AsyncMock( + side_effect=litellm.InternalServerError( + message="boom", model="gpt-5", llm_provider="openai" + ) + ) + + table = _RecordingManagedObjectTable(rows=(mock_job,)) + mock_prisma_client.db.litellm_managedobjecttable = table + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget: + await check_responses_cost_instance.check_responses_cost() + + mock_sdk_aget.assert_not_called() + mock_llm_router.aget_responses.assert_awaited_once() + assert table.updates == () + @pytest.mark.asyncio async def test_check_responses_cost_multiple_jobs( self, check_responses_cost_instance, mock_prisma_client diff --git a/tests/unit/proxy/common_utils/test_http_parsing_utils.py b/tests/unit/proxy/common_utils/test_http_parsing_utils.py index 7677837aa0f..f6f148f02a8 100644 --- a/tests/unit/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/unit/proxy/common_utils/test_http_parsing_utils.py @@ -34,6 +34,7 @@ from fastapi import Request as Request_http_parsing from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.utils import _invalidate_model_cost_lowercase_map from starlette.types import Message +from starlette.websockets import WebSocket from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome @@ -1565,3 +1566,36 @@ async def test_read_request_body_unexpected_error(): result = await read_request_body(_request(receive)) assert result == {} + + +async def _never_receive() -> Message: + raise AssertionError("the OTLP route check never reads the body") + + +async def _never_send(message: Message) -> None: + raise AssertionError("the OTLP route check never sends") + + +@pytest.mark.parametrize( + ("scope", "expected"), + [ + ({"type": "http", "method": "POST", "path": "/v1/traces"}, True), + ({"type": "http", "method": "POST", "path": "/v1/logs"}, True), + ({"type": "http", "method": "GET", "path": "/v1/traces"}, False), + ({"type": "http", "method": "POST", "path": "/v1/responses"}, False), + ({"type": "http", "path": "/v1/traces"}, False), + ({"type": "websocket", "path": "/v1/traces"}, False), + ({"type": "websocket", "path": "/v1/responses"}, False), + ], +) +def test_is_otlp_trace_request_matches_only_http_posts_to_the_otlp_routes( + scope: dict[str, object], expected: bool +) -> None: + full_scope: Final[dict[str, object]] = {**scope, "headers": [], "query_string": b""} + connection: Final = ( + WebSocket(full_scope, _never_receive, _never_send) + if scope["type"] == "websocket" + else Request(full_scope, _never_receive) + ) + + assert http_parsing_utils.is_otlp_trace_request(connection) is expected diff --git a/tests/unit/proxy/db/test_db_spend_update_writer.py b/tests/unit/proxy/db/test_db_spend_update_writer.py index dc9b8fd395b..4a508fcd272 100644 --- a/tests/unit/proxy/db/test_db_spend_update_writer.py +++ b/tests/unit/proxy/db/test_db_spend_update_writer.py @@ -537,6 +537,54 @@ async def test_update_daily_spend_retries_connect_errors(monkeypatch): assert len(prisma_client.db.statements) == 2 +@pytest.mark.asyncio +async def test_update_daily_spend_retries_lock_timeout_errors(monkeypatch: pytest.MonkeyPatch) -> None: + def _lock_timeout_error() -> PrismaDataError: + return PrismaDataError( + data={ + "user_facing_error": { + "is_panic": False, + "message": "Error querying the database: canceling statement due to lock timeout", + "meta": {"code": "55P03", "message": "canceling statement due to lock timeout"}, + } + } + ) + + outcomes: Final = iter([_lock_timeout_error(), None]) + + def first_attempt_locks_out() -> int: + outcome: Final = next(outcomes) + if outcome is not None: + raise outcome + return 1 + + prisma_client: Final = _RecordingPrisma(execute_raw=first_attempt_locks_out) + proxy_logging: Final = MagicMock() + proxy_logging.failure_handler = AsyncMock() + + async def fake_sleep(seconds: float) -> None: + return None + + monkeypatch.setattr("litellm.proxy.db.db_spend_update_writer.asyncio.sleep", fake_sleep) + daily_spend_transactions: Final = {"k1": _daily_txn()} + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=3, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging, + daily_spend_transactions=daily_spend_transactions, + entity_type="user", + entity_id_field="user_id", + ) + + assert len(prisma_client.db.statements) == 2, ( + "a 55P03 lock_timeout cancels the upsert before it applies, so the writer must " + "resend it in place instead of only requeueing it for a next tick a shutdown " + "flush never gets" + ) + assert daily_spend_transactions == {}, "the retried batch must drain the transactions dict" + proxy_logging.failure_handler.assert_not_called() + + @pytest.mark.asyncio async def test_update_daily_spend_sorting(): """ diff --git a/tests/unit/proxy/db/test_log_db_metrics_service_spans.py b/tests/unit/proxy/db/test_log_db_metrics_service_spans.py new file mode 100644 index 00000000000..a7fd830b948 --- /dev/null +++ b/tests/unit/proxy/db/test_log_db_metrics_service_spans.py @@ -0,0 +1,260 @@ +import asyncio +import importlib +from collections.abc import Sequence +from datetime import datetime, timedelta +from types import SimpleNamespace +from typing import Final, NotRequired, TypedDict + +from typing_extensions import ReadOnly +from unittest.mock import AsyncMock, patch + +import httpx +import pytest +from opentelemetry.sdk.trace import ReadableSpan, TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor, SpanExportResult +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.trace import StatusCode +from prisma.errors import ClientNotConnectedError +from pydantic import TypeAdapter + +import litellm +from litellm._service_logger import ServiceTypes +from litellm.integrations.datadog.datadog import DataDogLogger +from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig +from litellm.proxy.db.log_db_metrics import log_db_metrics +from litellm.proxy.db.prisma_client import _PrismaDrainTracker, _TrackedPrismaEngine +from litellm.proxy.proxy_server import proxy_logging_obj +from litellm.integrations.prometheus_services import PrometheusServicesLogger +from prometheus_client import REGISTRY + + +class _ServiceEvent(TypedDict): + service: ReadOnly[str] + call_type: ReadOnly[str] + duration: ReadOnly[float] + is_error: ReadOnly[bool] + error: ReadOnly[str | None] + event_metadata: ReadOnly[dict[str, str] | None] + table_name: NotRequired[ReadOnly[str]] + + +class _ServiceSpanExporter(InMemorySpanExporter): + def __init__(self) -> None: + super().__init__() + self.service_span_exported: Final = asyncio.Event() + + def export(self, spans: Sequence[ReadableSpan]) -> SpanExportResult: + result: Final = super().export(spans) + self.service_span_exported.set() + return result + + +class _Rig: + def __init__(self, exporter: _ServiceSpanExporter, provider: TracerProvider, datadog: DataDogLogger) -> None: + self.exporter: Final = exporter + self.provider: Final = provider + self.datadog: Final = datadog + + def service_spans(self) -> tuple[ReadableSpan, ...]: + return tuple(span for span in self.exporter.get_finished_spans() if span.name != "request") + + def events(self) -> tuple[_ServiceEvent, ...]: + adapter: Final = TypeAdapter(_ServiceEvent) + return tuple(adapter.validate_json(entry["message"]) for entry in self.datadog.log_queue) + + +def _discard_periodic_flush(coroutine: object) -> None: + close: Final = getattr(coroutine, "close") + close() + + +@pytest.fixture +def rig(monkeypatch: pytest.MonkeyPatch) -> _Rig: + monkeypatch.setenv("DD_API_KEY", "test_api_key") + monkeypatch.setenv("DD_SITE", "test.datadoghq.com") + monkeypatch.setattr(litellm, "datadog_params", None) + exporter: Final = _ServiceSpanExporter() + provider: Final = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + otel: Final = OpenTelemetry(config=OpenTelemetryConfig(exporter=exporter), tracer_provider=provider) + with patch("asyncio.create_task", side_effect=_discard_periodic_flush): + datadog: Final = DataDogLogger() + monkeypatch.setattr(litellm, "service_callback", [otel, datadog, "prometheus_system"]) + monkeypatch.setattr(proxy_logging_obj.service_logging_obj, "dd_logger", datadog, raising=False) + monkeypatch.setattr( + proxy_logging_obj.service_logging_obj, "prometheusServicesLogger", PrometheusServicesLogger(), raising=False + ) + return _Rig(exporter, provider, datadog) + + +async def _run_prisma_query() -> None: + engine: Final = _TrackedPrismaEngine(SimpleNamespace(query=AsyncMock(return_value={})), _PrismaDrainTracker()) + await engine.query("{}", tx_id=None) + + +@log_db_metrics +async def read_spend_rows(**kwargs: object) -> str: + await _run_prisma_query() + return "success" + + +def _logged_db_latency() -> tuple[float, float]: + labels: Final = {ServiceTypes.DB.value: ServiceTypes.DB.value} + total: Final = REGISTRY.get_sample_value("litellm_postgres_latency_sum", labels) + count: Final = REGISTRY.get_sample_value("litellm_postgres_latency_count", labels) + return (total or 0.0, count or 0.0) + + +def _ns(moment: datetime) -> int: + return int(moment.timestamp() * 1e9) + + +_DB_CALL_START: Final = datetime(2026, 1, 1, 12, 0, 0) +_DB_CALL_DURATION: Final = timedelta(milliseconds=250) + + +class _ScriptedClock: + def __init__(self, *moments: datetime) -> None: + self._moments: Final = iter(moments) + + def now(self) -> datetime: + return next(self._moments) + + +@pytest.fixture +def scripted_db_clock(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + importlib.import_module("litellm.proxy.db.log_db_metrics"), + "datetime", + _ScriptedClock(_DB_CALL_START, _DB_CALL_START + _DB_CALL_DURATION), + ) + + +@pytest.mark.asyncio +async def test_a_db_success_is_reported_on_the_parent_span_with_its_duration_and_times( + rig: _Rig, scripted_db_clock: None +) -> None: + parent: Final = rig.provider.get_tracer("db-test").start_span("request") + latency_before: Final = _logged_db_latency() + + result: Final = await read_spend_rows(parent_otel_span=parent) + await asyncio.wait_for(rig.exporter.service_span_exported.wait(), timeout=10) + + assert result == "success" + spans: Final = rig.service_spans() + assert len(spans) == 1 + span: Final = spans[0] + assert span.parent is not None + assert span.parent.span_id == parent.get_span_context().span_id + assert span.attributes is not None + assert span.attributes["service"] == ServiceTypes.DB.value + assert span.attributes["call_type"] == "read_spend_rows" + assert span.status.status_code == StatusCode.OK + assert span.start_time == _ns(_DB_CALL_START) + assert span.end_time == _ns(_DB_CALL_START + _DB_CALL_DURATION) + latency_after: Final = _logged_db_latency() + assert latency_after[1] - latency_before[1] == 1 + assert latency_after[0] - latency_before[0] == pytest.approx(_DB_CALL_DURATION.total_seconds()) + assert rig.events() == () + + +@pytest.mark.asyncio +async def test_db_event_metadata_names_only_the_table_and_never_the_raw_kwargs(rig: _Rig) -> None: + parent: Final = rig.provider.get_tracer("db-test").start_span("request") + + await read_spend_rows( + parent_otel_span=parent, + table_name="LiteLLM_SpendLogs", + token="sk-secret-should-not-leak", + prisma_client=object(), + ) + await asyncio.wait_for(rig.exporter.service_span_exported.wait(), timeout=10) + + span_attributes: Final = rig.service_spans()[0].attributes + assert span_attributes is not None + assert span_attributes["table_name"] == "LiteLLM_SpendLogs" + assert not {"token", "prisma_client", "parent_otel_span"} & set(span_attributes) + assert all("sk-secret-should-not-leak" not in str(value) for value in span_attributes.values()) + + +@pytest.mark.asyncio +async def test_the_logged_db_duration_is_the_span_wall_clock_of_the_wrapped_call( + rig: _Rig, scripted_db_clock: None +) -> None: + parent: Final = rig.provider.get_tracer("db-test").start_span("request") + latency_before: Final = _logged_db_latency() + + await read_spend_rows(parent_otel_span=parent) + await asyncio.wait_for(rig.exporter.service_span_exported.wait(), timeout=10) + + span: Final = rig.service_spans()[0] + assert span.start_time is not None and span.end_time is not None + latency_after: Final = _logged_db_latency() + logged_duration: Final = latency_after[0] - latency_before[0] + assert latency_after[1] - latency_before[1] == 1 + assert logged_duration == pytest.approx((span.end_time - span.start_time) / 1e9, rel=1e-3, abs=2e-6) + assert logged_duration == pytest.approx(_DB_CALL_DURATION.total_seconds()) + + +@log_db_metrics +async def disconnected_read(**kwargs: object) -> str: + raise ClientNotConnectedError() + + +@pytest.mark.asyncio +async def test_a_prisma_error_is_reported_as_a_db_failure_and_reraised(rig: _Rig) -> None: + parent: Final = rig.provider.get_tracer("db-test").start_span("request") + + with pytest.raises(ClientNotConnectedError, match="Client is not connected to the query engine"): + await disconnected_read(parent_otel_span=parent) + + spans: Final = rig.service_spans() + assert len(spans) == 1 + assert spans[0].parent is not None + assert spans[0].parent.span_id == parent.get_span_context().span_id + assert spans[0].status.status_code == StatusCode.ERROR + assert spans[0].attributes is not None + assert spans[0].attributes["call_type"] == "disconnected_read" + assert spans[0].attributes["service"] == ServiceTypes.DB.value + assert "Client is not connected" in str(spans[0].attributes["error"]) + events: Final = rig.events() + assert len(events) == 1 + assert events[0]["is_error"] is True + assert events[0]["call_type"] == "disconnected_read" + assert isinstance(events[0]["duration"], float) + assert "Client is not connected" in (events[0]["error"] or "") + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("error", "is_db_error"), + [ + (ValueError("Generic error"), False), + (KeyError("Missing key"), False), + (TypeError("Type error"), False), + (httpx.ConnectError("Failed to connect"), True), + (httpx.TimeoutException("Request timed out"), True), + (ClientNotConnectedError(), True), + ], +) +async def test_only_db_errors_are_reported_as_db_failures(rig: _Rig, error: Exception, is_db_error: bool) -> None: + parent: Final = rig.provider.get_tracer("db-test").start_span("request") + + @log_db_metrics + async def failing_read(**kwargs: object) -> str: + raise error + + with pytest.raises(type(error)): + await failing_read(parent_otel_span=parent) + + spans: Final = rig.service_spans() + events: Final = rig.events() + if is_db_error: + assert [span.status.status_code for span in spans] == [StatusCode.ERROR] + assert [(event["service"], event["call_type"], event["is_error"]) for event in events] == [ + (ServiceTypes.DB.value, "failing_read", True) + ] + assert isinstance(events[0]["duration"], float) + else: + assert spans == () + assert events == () diff --git a/tests/unit/proxy/google_endpoints/test_google_api_endpoints.py b/tests/unit/proxy/google_endpoints/test_google_api_endpoints.py index e4cd7d9dfa8..451792109de 100644 --- a/tests/unit/proxy/google_endpoints/test_google_api_endpoints.py +++ b/tests/unit/proxy/google_endpoints/test_google_api_endpoints.py @@ -3,6 +3,7 @@ Test to verify the Google GenAI proxy API endpoints """ +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -179,3 +180,30 @@ def test_google_count_tokens_unchanged(): assert response.status_code == 200 body = response.json() assert body["totalTokens"] == 7 + + +def test_google_count_tokens_uses_normalized_total_when_provider_shape_differs(): + """A provider that counts in Anthropic shape (Bedrock, Anthropic) carries no totalTokens key, so the normalized total must win over the raw response's missing field.""" + try: + client: Final = _build_test_client() + except ImportError as e: + pytest.skip(f"Skipping test due to missing dependency: {e}") + + fake_response: Final = MagicMock() + fake_response.original_response = {"input_tokens": 2167} + fake_response.total_tokens = 2167 + + with patch( + "litellm.proxy.proxy_server.token_counter", + new_callable=AsyncMock, + return_value=fake_response, + ): + response: Final = client.post( + "/v1beta/models/claude-opus-4-8:countTokens", + json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]}, + ) + + assert response.status_code == 200 + body: Final = response.json() + assert body["totalTokens"] == 2167 + assert body["promptTokensDetails"] == [] diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py b/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py index 1e3c16ecf39..8fa544819e9 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py @@ -5,6 +5,8 @@ import os from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch +from typing import Final, Literal + import httpx import pytest from fastapi import HTTPException @@ -20,7 +22,8 @@ from litellm.proxy.guardrails.guardrail_registry import ( guardrail_class_registry, guardrail_initializer_registry, ) -from litellm.types.utils import GenericGuardrailAPIInputs +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse def test_akto_in_guardrail_initializer_registry(): @@ -1869,3 +1872,121 @@ async def test_a_modified_verdict_that_changed_no_text_blocks(akto_pre_call, sam inputs=sample_inputs, request_data=sample_request_data, input_type="request" ) assert exc_info.value.message == "Content masked by Akto guardrail policy could not be applied" + + +REQUEST_CHECK: Final = {"akto_connector": "litellm", "guardrails": "true", "ingest_data": "true"} +RESPONSE_CHECK: Final = {"akto_connector": "litellm", "response_guardrails": "true", "ingest_data": "true"} + + +def _logging_only_akto( + post: AsyncMock, unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed" +) -> AktoGuardrail: + handler: Final = MagicMock(spec=AsyncHTTPHandler) + handler.post = post + return AktoGuardrail( + async_handler=handler, + akto_base_url="http://localhost:9090", + akto_api_key="test-token", + guardrail_name="test-logging_only", + event_hook="logging_only", + unreachable_fallback=unreachable_fallback, + ) + + +def _logged_call(text: str = "Hello, how are you?") -> dict[str, object]: + return { + "model": "gpt-5.5", + "messages": [{"role": "user", "content": text}], + "litellm_call_id": "call-1", + "litellm_params": {"metadata": {"user_api_key_request_route": "/v1/chat/completions"}}, + "standard_logging_object": {"guardrail_information": []}, + } + + +def _logged_response(text: str = "Fine, thanks") -> ModelResponse: + return ModelResponse(id="resp-1", choices=[{"message": {"role": "assistant", "content": text}}]) + + +def _recorded_entries(logged_kwargs: dict[str, object]) -> list[dict[str, object]]: + standard_logging_object: Final = logged_kwargs["standard_logging_object"] + assert isinstance(standard_logging_object, dict), logged_kwargs + entries: Final = standard_logging_object["guardrail_information"] + assert isinstance(entries, list), standard_logging_object + return entries + + +def test_logging_only_is_a_supported_mode() -> None: + assert GuardrailEventHooks.logging_only in AktoGuardrail.get_supported_event_hooks() + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("input_type", "flags"), [("request", REQUEST_CHECK), ("response", RESPONSE_CHECK)]) +async def test_logging_only_handles_both_directions( + input_type: Literal["request", "response"], flags: dict[str, str] +) -> None: + guardrail: Final = _logging_only_akto(AsyncMock(return_value=_mock_allowed_response())) + request_data: Final = _with_complete_response({}) if input_type == "response" else {} + + await guardrail.apply_guardrail( + inputs=GenericGuardrailAPIInputs(texts=["hi"]), request_data=request_data, input_type=input_type + ) + + assert [params for params, _ in _calls(guardrail)] == [flags] + + +@pytest.mark.asyncio +async def test_logging_only_checks_and_records_the_logged_request_and_response() -> None: + guardrail: Final = _logging_only_akto(AsyncMock(return_value=_mock_allowed_response())) + + await guardrail.async_logging_hook(_logged_call(), _logged_response(), "acompletion") + + sent: Final = _calls(guardrail) + assert [params for params, _ in sent] == [REQUEST_CHECK, RESPONSE_CHECK] + assert "Hello, how are you?" in sent[0][1]["requestPayload"] + assert "Fine, thanks" in sent[1][1]["responsePayload"] + + +@pytest.mark.asyncio +async def test_logging_only_block_verdict_is_recorded_without_raising() -> None: + guardrail: Final = _logging_only_akto(AsyncMock(return_value=_mock_blocked_response("Rejected"))) + response: Final = _logged_response() + + out_kwargs, out_result = await guardrail.async_logging_hook(_logged_call(), response, "acompletion") + + assert out_result is response + assert [params for params, _ in _calls(guardrail)] == [REQUEST_CHECK] + [entry] = _recorded_entries(out_kwargs) + assert (entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"]) == ( + "test-logging_only", + "logging_only", + "guardrail_intervened", + ) + + +@pytest.mark.asyncio +async def test_logging_only_ignores_an_unreachable_akto_even_when_fail_closed() -> None: + guardrail: Final = _logging_only_akto(AsyncMock(side_effect=httpx.ConnectError("refused")), "fail_closed") + response: Final = _logged_response() + + out_kwargs, out_result = await guardrail.async_logging_hook(_logged_call(), response, "acompletion") + + assert out_result is response + [entry] = _recorded_entries(out_kwargs) + assert (entry["guardrail_mode"], entry["guardrail_response"]) == ( + "logging_only", + "Akto guardrail service unreachable", + ) + + +@pytest.mark.asyncio +async def test_logging_only_sends_each_attachment_once() -> None: + guardrail: Final = _logging_only_akto(_file_verdict({"Allowed": True})) + image: Final = {"type": "image_url", "image_url": {"url": "https://example.com/a.png"}} + call: Final = {**_logged_call(), "messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}, image]}]} + + await guardrail.async_logging_hook(call, _logged_response(), "acompletion") + + [file_call] = _file_calls(guardrail) + assert json.loads(file_call.kwargs["data"])["files"] == [ + {"filename": "a.png", "type": "image", "url": "https://example.com/a.png"} + ] diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto_attachments.py b/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto_attachments.py index 8be2f58f059..b3f61985c7e 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto_attachments.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto_attachments.py @@ -1,5 +1,6 @@ import base64 import json +from typing import Final import pytest @@ -471,3 +472,13 @@ def test_text_that_isnt_valid_utf8_is_still_sent(block): [attachment] = request_attachments(request_data).attachments assert base64.b64decode(attachment.content or "") == "a\ud800".encode(errors="surrogatepass") + + +def test_request_attachments_reads_a_list_shared_by_messages_and_input_once() -> None: + shared: Final = [ + {"role": "user", "content": [{"type": "image_url", "image_url": {"url": "https://example.com/a.png"}}]} + ] + + assert request_attachments({"messages": shared, "input": shared}).attachments == ( + Attachment("a.png", "image", url="https://example.com/a.png"), + ) diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 01be28470d1..f9795197f90 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -4,12 +4,14 @@ Unit tests for Bedrock Guardrails import json import asyncio +from typing import Final from datetime import datetime, timezone import sys from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest +import respx from fastapi import HTTPException @@ -7483,3 +7485,142 @@ async def test_should_raise_guardrail_blocked_exception_null_fields(): guardrail._should_raise_guardrail_blocked_exception(response_null_grounding) is False ) + + +_MASKING_GUARDRAIL_URL: Final = "https://bedrock-runtime.us-east-1.amazonaws.com/guardrail/wf0hkdb5x07f/version/DRAFT/apply" + + +def _masking_guardrail() -> BedrockGuardrail: + return BedrockGuardrail( + guardrailIdentifier="wf0hkdb5x07f", + guardrailVersion="DRAFT", + aws_access_key_id="fake-access-key", + aws_secret_access_key="fake-secret-key", + aws_region_name="us-east-1", + ) + + +def _anonymized_reply(masked_texts: tuple[str, ...], entity_types: tuple[str, ...]) -> httpx.Response: + return httpx.Response( + 200, + json={ + "action": "GUARDRAIL_INTERVENED", + "outputs": [{"text": text} for text in masked_texts], + "assessments": [ + { + "sensitiveInformationPolicy": { + "piiEntities": [ + {"type": entity_type, "match": "redacted", "action": "ANONYMIZED"} + for entity_type in entity_types + ] + } + } + ], + }, + ) + + +def _sent_texts(route: respx.Route) -> tuple[str, ...]: + body: Final = json.loads(route.calls[0].request.content) + assert body["source"] == "INPUT" + return tuple(item["text"]["text"] for item in body["content"]) + + +@pytest.mark.asyncio +async def test_during_call_masking_rewrites_pii_in_messages( + respx_mock: respx.MockRouter, httpx_transport: None +) -> None: + guardrail: Final = _masking_guardrail() + route: Final = respx_mock.post(_MASKING_GUARDRAIL_URL).mock( + return_value=_anonymized_reply( + ( + "Hello, my phone number is {PHONE}", + "Hello, how can I help you today?", + "I need to cancel my order", + "ok, my credit card number is {CREDIT_DEBIT_CARD_NUMBER}", + ), + ("PHONE", "CREDIT_DEBIT_CARD_NUMBER"), + ) + ) + + response: Final = await guardrail.async_moderation_hook( + data={ + "model": "gpt-5.5", + "messages": [ + {"role": "user", "content": "Hello, my phone number is +1 412 555 1212"}, + {"role": "assistant", "content": "Hello, how can I help you today?"}, + {"role": "user", "content": "I need to cancel my order"}, + {"role": "user", "content": "ok, my credit card number is 1234-5678-9012-3456"}, + ], + }, + user_api_key_dict=UserAPIKeyAuth(), + call_type="completion", + ) + + assert route.call_count == 1 + assert _sent_texts(route) == ( + "Hello, my phone number is +1 412 555 1212", + "Hello, how can I help you today?", + "I need to cancel my order", + "ok, my credit card number is 1234-5678-9012-3456", + ) + assert response is not None + assert [message["content"] for message in response["messages"]] == [ + "Hello, my phone number is {PHONE}", + "Hello, how can I help you today?", + "I need to cancel my order", + "ok, my credit card number is {CREDIT_DEBIT_CARD_NUMBER}", + ] + + +@pytest.mark.asyncio +async def test_during_call_masking_rewrites_only_pii_block_in_content_list( + respx_mock: respx.MockRouter, httpx_transport: None +) -> None: + guardrail: Final = _masking_guardrail() + route: Final = respx_mock.post(_MASKING_GUARDRAIL_URL).mock( + return_value=_anonymized_reply( + ( + "Hello, my phone number is {PHONE}", + "what time is it?", + "Hello, how can I help you today?", + "who is the president of the united states?", + ), + ("PHONE",), + ) + ) + + response: Final = await guardrail.async_moderation_hook( + data={ + "model": "gpt-5.5", + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Hello, my phone number is +1 412 555 1212"}, + {"type": "text", "text": "what time is it?"}, + ], + }, + {"role": "assistant", "content": "Hello, how can I help you today?"}, + {"role": "user", "content": "who is the president of the united states?"}, + ], + }, + user_api_key_dict=UserAPIKeyAuth(), + call_type="completion", + ) + + assert route.call_count == 1 + assert _sent_texts(route) == ( + "Hello, my phone number is +1 412 555 1212", + "what time is it?", + "Hello, how can I help you today?", + "who is the president of the united states?", + ) + assert response is not None + messages: Final = response["messages"] + assert messages[0]["content"] == [ + {"type": "text", "text": "Hello, my phone number is {PHONE}"}, + {"type": "text", "text": "what time is it?"}, + ] + assert messages[1]["content"] == "Hello, how can I help you today?" + assert messages[2]["content"] == "who is the president of the united states?" diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py index 0cac6228085..7d0c81af4c2 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -8,13 +8,19 @@ import copy import json import os import re +from collections.abc import Iterable, Sequence +from concurrent.futures import ThreadPoolExecutor from contextlib import asynccontextmanager from typing import Final, Literal from unittest.mock import MagicMock, patch +import aiohttp from aiohttp import web +from aiohttp.client_proto import ResponseHandler from aiohttp.test_utils import TestServer import pytest +from pydantic import TypeAdapter +from typing_extensions import ReadOnly, TypedDict import litellm @@ -27,6 +33,7 @@ from litellm.proxy.guardrails.guardrail_hooks.presidio import ( ) from litellm.exceptions import GuardrailRaisedException from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType +from litellm.types.proxy.guardrails.guardrail_hooks.presidio import PresidioAnalyzeRequest, PresidioAnalyzeResponseItem from litellm.types.utils import Choices, Delta, Message, ModelResponse, StreamingChoices from litellm.exceptions import BlockedPiiEntityError @@ -4575,3 +4582,251 @@ async def test_presidio_language_configuration_with_per_request_override(): assert analyze_request_default["language"] == "de" assert analyze_request_default["text"] == test_text + + +_CARD_NUMBER: Final = "4111-1111-1111-1111" +_EMAIL: Final = "test@example.com" +_BLOCK_CARD_MASK_EMAIL: Final = { + PiiEntityType.CREDIT_CARD: PiiAction.BLOCK, + PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK, +} + + +class _PresidioAnonymizeRequest(TypedDict): + text: ReadOnly[str] + + +class _PresidioAnonymizeReply(TypedDict): + text: ReadOnly[str] + items: ReadOnly[tuple[()]] + + +_ANALYZE_REQUEST: Final = TypeAdapter(PresidioAnalyzeRequest) +_ANONYMIZE_REQUEST: Final = TypeAdapter(_PresidioAnonymizeRequest) + + +def _card_and_email_spans(text: str) -> tuple[PresidioAnalyzeResponseItem, ...]: + return tuple( + PresidioAnalyzeResponseItem( + entity_type=entity_type, + start=text.index(value), + end=text.index(value) + len(value), + score=1.0, + analysis_explanation=None, + ) + for entity_type, value in (("CREDIT_CARD", _CARD_NUMBER), ("EMAIL_ADDRESS", _EMAIL)) + if value in text + ) + + +class _InMemoryPresidioTransport(asyncio.Transport): + def __init__(self, protocol: ResponseHandler, analyzed: asyncio.Queue[PresidioAnalyzeRequest]) -> None: + super().__init__() + self._protocol: Final = protocol + self._analyzed: Final = analyzed + self._received = b"" + self._closing = False + + def write(self, data: bytes | bytearray | memoryview) -> None: + self._received += bytes(data) + head, separator, body = self._received.partition(b"\r\n\r\n") + if not separator: + return + request_line, *header_lines = head.decode().split("\r\n") + headers: Final = {name.lower(): value.strip() for name, _, value in (line.partition(":") for line in header_lines)} + if len(body) < int(headers.get("content-length", "0")): + return + self._received = b"" + reply: Final = json.dumps(self._reply(request_line.split(" ")[1], body)).encode() + response: Final = ( + b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nConnection: close\r\n" + + f"Content-Length: {len(reply)}\r\n\r\n".encode() + + reply + ) + asyncio.get_running_loop().call_soon(self._protocol.data_received, response) + + def _reply(self, path: str, body: bytes) -> tuple[PresidioAnalyzeResponseItem, ...] | _PresidioAnonymizeReply: + if path == "/analyze": + analyze_request: Final = _ANALYZE_REQUEST.validate_json(body) + self._analyzed.put_nowait(analyze_request) + return _card_and_email_spans(analyze_request.get("text") or "") + assert path == "/anonymize", path + return _PresidioAnonymizeReply(text=_ANONYMIZE_REQUEST.validate_json(body)["text"], items=()) + + def writelines(self, list_of_data: Iterable[bytes | bytearray | memoryview]) -> None: + self.write(b"".join(bytes(chunk) for chunk in list_of_data)) + + def is_closing(self) -> bool: + return self._closing + + def close(self) -> None: + if not self._closing: + self._closing = True + asyncio.get_running_loop().call_soon(self._protocol.connection_lost, None) + + def abort(self) -> None: + self.close() + + def get_extra_info(self, name: str, default: object = None) -> object: + return default + + def can_write_eof(self) -> bool: + return False + + def get_write_buffer_size(self) -> int: + return 0 + + def pause_reading(self) -> None: + return None + + def resume_reading(self) -> None: + return None + + +class _InMemoryPresidioConnector(aiohttp.BaseConnector): + def __init__(self, analyzed: asyncio.Queue[PresidioAnalyzeRequest]) -> None: + super().__init__() + self._analyzed: Final = analyzed + + async def _create_connection( # pyright: ignore[reportImplicitOverride] # aiohttp's connector extension point + self, req: aiohttp.ClientRequest, traces: Sequence[object], timeout: aiohttp.ClientTimeout + ) -> ResponseHandler: + protocol: Final = ResponseHandler(asyncio.get_running_loop()) + protocol.connection_made(_InMemoryPresidioTransport(protocol, self._analyzed)) + return protocol + + +def _drain(analyzed: asyncio.Queue[PresidioAnalyzeRequest]) -> tuple[PresidioAnalyzeRequest, ...]: + return tuple(analyzed.get_nowait() for _ in range(analyzed.qsize())) + + +def _guardrail_with_in_memory_presidio( + analyzed: asyncio.Queue[PresidioAnalyzeRequest], +) -> OPTIONAL_PresidioPIIMasking: + guardrail: Final = OPTIONAL_PresidioPIIMasking( + pii_entities_config=_BLOCK_CARD_MASK_EMAIL, + presidio_analyzer_api_base="http://presidio-analyzer.test/", + presidio_anonymizer_api_base="http://presidio-anonymizer.test/", + ) + guardrail._http_session = aiohttp.ClientSession(connector=_InMemoryPresidioConnector(analyzed)) + return guardrail + + +@pytest.mark.asyncio +async def test_check_pii_raises_blocked_entity_for_card() -> None: + text: Final = f"My credit card number is {_CARD_NUMBER} and my email is {_EMAIL}" + analyzed: Final[asyncio.Queue[PresidioAnalyzeRequest]] = asyncio.Queue() + guardrail: Final = _guardrail_with_in_memory_presidio(analyzed) + try: + with pytest.raises(BlockedPiiEntityError) as excinfo: + await guardrail.check_pii(text=text, output_parse_pii=True, presidio_config=None, request_data={}) + finally: + await guardrail._close_http_session() + + analyze_requests: Final = _drain(analyzed) + assert len(analyze_requests) == 1 + assert analyze_requests[0].get("text") == text + assert set(analyze_requests[0].get("entities") or ()) == set(_BLOCK_CARD_MASK_EMAIL) + assert excinfo.value.entity_type == PiiEntityType.CREDIT_CARD + assert excinfo.value.guardrail_name == guardrail.guardrail_name + + +@pytest.mark.asyncio +async def test_pre_call_hook_raises_blocked_entity_for_card_message( + mock_user_api_key: UserAPIKeyAuth, mock_cache: DualCache +) -> None: + user_text: Final = f"My credit card is {_CARD_NUMBER} and my email is {_EMAIL}." + analyzed: Final[asyncio.Queue[PresidioAnalyzeRequest]] = asyncio.Queue() + guardrail: Final = _guardrail_with_in_memory_presidio(analyzed) + try: + with pytest.raises(BlockedPiiEntityError) as excinfo: + await guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key, + cache=mock_cache, + data={ + "messages": [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": user_text}, + ], + "model": "gpt-5-mini", + }, + call_type="completion", + ) + finally: + await guardrail._close_http_session() + + analyze_requests: Final = _drain(analyzed) + assert user_text in [payload.get("text") for payload in analyze_requests] + assert all(set(payload.get("entities") or ()) == set(_BLOCK_CARD_MASK_EMAIL) for payload in analyze_requests) + assert excinfo.value.entity_type == PiiEntityType.CREDIT_CARD + assert excinfo.value.guardrail_name == guardrail.guardrail_name + + +@pytest.mark.asyncio +async def test_legacy_pii_masking_config_registers_logging_only_guardrail(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("PRESIDIO_ANALYZER_API_BASE", "http://localhost:5002") + monkeypatch.setenv("PRESIDIO_ANONYMIZER_API_BASE", "http://localhost:5001") + monkeypatch.setattr(litellm, "guardrail_name_config_map", {}) + monkeypatch.setattr(litellm, "callbacks", []) + + from litellm.proxy.guardrails.init_guardrails import initialize_guardrails + from litellm.types.guardrails import GuardrailEventHooks + + guardrails_config: Final = [ + { + "pii_masking": { + "callbacks": ["presidio"], + "default_on": True, + "logging_only": True, + } + } + ] + assert len(litellm.guardrail_name_config_map) == 0 + initialize_guardrails( + guardrails_config=guardrails_config, + premium_user=True, + config_file_path="", + litellm_settings={"guardrails": guardrails_config}, + ) + assert len(litellm.guardrail_name_config_map) == 1 + + pii_masking_obj: Final = next( + (c for c in litellm.callbacks if isinstance(c, OPTIONAL_PresidioPIIMasking)), + None, + ) + assert pii_masking_obj is not None + assert hasattr(pii_masking_obj, "logging_only") + assert pii_masking_obj.event_hook == GuardrailEventHooks.logging_only + assert pii_masking_obj.should_run_guardrail( + data={}, event_type=GuardrailEventHooks.logging_only + ) + + +async def _one_session(guardrail: OPTIONAL_PresidioPIIMasking) -> aiohttp.ClientSession: + async with guardrail._get_session_iterator() as session: + return session + + +@pytest.mark.asyncio +async def test_get_session_iterator_reuses_one_session_on_main_thread( + presidio_guardrail: OPTIONAL_PresidioPIIMasking, +) -> None: + sessions: Final = tuple([await _one_session(presidio_guardrail) for _ in range(10)]) + assert all(session is sessions[0] for session in sessions) + assert sessions[0] is presidio_guardrail._http_session + await presidio_guardrail._close_http_session() + + +def test_get_session_iterator_reuses_one_session_per_background_loop( + presidio_guardrail: OPTIONAL_PresidioPIIMasking, +) -> None: + async def collect_and_close() -> tuple[aiohttp.ClientSession, ...]: + collected: Final = tuple([await _one_session(presidio_guardrail) for _ in range(10)]) + await collected[0].close() + return collected + + with ThreadPoolExecutor(max_workers=1) as pool: + sessions: Final = pool.submit(asyncio.run, collect_and_close()).result() + assert len(sessions) == 10 + assert all(session is sessions[0] for session in sessions) + assert presidio_guardrail._http_session is None diff --git a/tests/unit/proxy/guardrails/test_guardrail_endpoints.py b/tests/unit/proxy/guardrails/test_guardrail_endpoints.py index fa33c2d462b..b7bc5f3f8ec 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/unit/proxy/guardrails/test_guardrail_endpoints.py @@ -575,7 +575,7 @@ async def test_get_guardrail_info_from_db(mocker, mock_prisma_client): """Test getting guardrail info from DB""" mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - response = await get_guardrail_info("test-db-guardrail") + response: Final = await get_guardrail_info("test-db-guardrail") assert response.guardrail_id == "test-db-guardrail" assert response.guardrail_name == "Test DB Guardrail" @@ -584,6 +584,21 @@ async def test_get_guardrail_info_from_db(mocker, mock_prisma_client): @pytest.mark.asyncio +async def test_get_guardrail_info_tolerates_invalid_stored_stream_scope(mocker, mock_prisma_client): + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mock_prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock( + return_value={ + **MOCK_DB_GUARDRAIL, + "litellm_params": { + **MOCK_DB_GUARDRAIL["litellm_params"], + "stream_scope": "sometimes", + }, + } + ) + + response = await get_guardrail_info("test-db-guardrail") + + assert response.litellm_params.stream_scope is None async def test_get_guardrail_info_normalizes_invalid_scope_from_db( mocker, mock_guardrail_registry, mock_in_memory_handler ): @@ -745,6 +760,40 @@ def test_get_guardrails_list_response_includes_guardrail_id(): assert response.guardrails[0].guardrail_id == "stable-config-id" +def test_get_guardrails_list_response_tolerates_invalid_config_stream_scope(): + from litellm.proxy.guardrails.guardrail_endpoints import ( + _get_guardrails_list_response, + ) + + response = _get_guardrails_list_response( + [ + { + "guardrail_id": "invalid-scope", + "guardrail_name": "invalid-scope", + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "stream_scope": "sometimes", + }, + }, + { + "guardrail_id": "valid-scope", + "guardrail_name": "valid-scope", + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "stream_scope": "STREAMING", + }, + }, + ] + ) + + assert response.guardrails[0].litellm_params is not None + assert response.guardrails[0].litellm_params.stream_scope is None + assert response.guardrails[1].litellm_params is not None + assert response.guardrails[1].litellm_params.stream_scope == "streaming" + + def test_get_provider_specific_params(): """Test getting provider-specific parameters""" from litellm.proxy.guardrails.guardrail_endpoints import _get_fields_from_model diff --git a/tests/unit/proxy/guardrails/test_guardrail_registry.py b/tests/unit/proxy/guardrails/test_guardrail_registry.py index c7b13f0bd00..3bc6d163e5d 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_registry.py +++ b/tests/unit/proxy/guardrails/test_guardrail_registry.py @@ -1,7 +1,7 @@ import json from collections.abc import Iterable, Iterator from typing import ClassVar, Final -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest from pydantic import ValidationError @@ -160,6 +160,43 @@ def test_duplicate_config_guardrail_names_get_distinct_stable_ids(): registry_module.guardrail_initializer_registry.pop("dup_name_test", None) +def test_initialize_guardrail_treats_invalid_stored_scope_as_both(): + from litellm.proxy.guardrails import guardrail_registry as registry_module + + guardrail_type: Final = "invalid_stored_scope_test" + + def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail: + return CustomGuardrail( + guardrail_name=guardrail["guardrail_name"], + event_hook=GuardrailEventHooks(litellm_params.mode), + default_on=True, + ) + + registry_module.guardrail_initializer_registry[guardrail_type] = _initializer + try: + handler: Final = InMemoryGuardrailHandler() + guardrail: Final = Guardrail( + guardrail_id="invalid-stored-scope", + guardrail_name="invalid-stored-scope", + litellm_params={ + "guardrail": guardrail_type, + "mode": "pre_call", + "default_on": True, + "stream_scope": "sometimes", + }, + ) + + parsed_guardrail: Final = handler.initialize_guardrail(guardrail=guardrail, source="db") + callback: Final = handler.guardrail_id_to_custom_guardrail["invalid-stored-scope"] + + assert parsed_guardrail["litellm_params"].stream_scope is None + assert callback is not None + assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is True + assert callback.should_run_guardrail(data={"stream": True}, event_type=GuardrailEventHooks.pre_call) is True + finally: + registry_module.guardrail_initializer_registry.pop(guardrail_type, None) + + def _register_mode_following_initializer(guardrail_type: str): """Registers like the shipped initializers do: construct, then add the instance to litellm's callbacks.""" import litellm @@ -1605,6 +1642,35 @@ def test_sync_guardrail_from_db_applies_db_dict_params_to_live_instance(): cb_list[:] = snapshot +def test_configure_callback_scoping_copies_stream_scope_when_constructor_omits_it(): + from litellm.proxy.guardrails.guardrail_registry import _configure_callback_scoping + + class _CtorWithoutStreamScope(CustomGuardrail): + def __init__(self) -> None: + super().__init__( + guardrail_name="scoped", + event_hook=GuardrailEventHooks.post_call, + default_on=True, + ) + + instance = _CtorWithoutStreamScope() + params = LitellmParams(guardrail="bedrock", mode="post_call", stream_scope="streaming") + _configure_callback_scoping(instance, "scoped", params) + + assert instance.stream_scope_default == "streaming" + assert instance.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True + assert instance.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False + + +def test_configure_callback_scoping_tolerates_a_custom_logger_callback(): + from litellm.integrations.custom_logger import CustomLogger + from litellm.proxy.guardrails.guardrail_registry import _configure_callback_scoping + + callback: Final = CustomLogger() + _configure_callback_scoping(callback, "logger-backed", LitellmParams(guardrail="custom", mode="pre_call")) # pyright: ignore[reportArgumentType] # module-path guardrails may be plain CustomLogger + assert "stream_scope_by_hook" not in vars(callback) + + _ENCRYPTED_PREFIX = "litellm_enc::" diff --git a/tests/unit/proxy/hooks/test_batch_rate_limiter.py b/tests/unit/proxy/hooks/test_batch_rate_limiter.py index a3e60a89c9f..2f71b177705 100644 --- a/tests/unit/proxy/hooks/test_batch_rate_limiter.py +++ b/tests/unit/proxy/hooks/test_batch_rate_limiter.py @@ -6,14 +6,19 @@ batch under a per-minute RPM/TPM budget. Scopes that configure `tpd_limit` are charged against a 24h token window instead of their minute counters. """ +import json import time -from collections.abc import Iterator +from collections.abc import Iterator, Sequence from datetime import datetime, timezone -from typing import Final +from typing import Final, Literal +import httpx import pytest +import respx from fastapi import HTTPException +from typing_extensions import ReadOnly, TypedDict +import litellm from litellm import DualCache from litellm.constants import BATCH_TPD_WINDOW_SECONDS from litellm.proxy._types import UserAPIKeyAuth @@ -22,6 +27,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.utils import InternalUsageCache, hash_token +from litellm.types.llms.openai import ChatCompletionUserMessage, LiteLLMBatchCreateRequest class _Clock: @@ -296,3 +302,242 @@ async def test_batch_rate_limit_error_reports_reset_time_in_utc_on_a_non_utc_pro assert exc.value.headers["retry-after"] == str(BATCH_TPD_WINDOW_SECONDS - 3 * 3600) assert exc.value.headers["reset_at"] == "2026-09-14 08:00:00 UTC" assert str(exc.value.detail).endswith("Limit resets at: 2026-09-14 08:00:00 UTC") + + +_BATCH_MODEL: Final = "gpt-3.5-turbo" +_MANAGED_FILE_ID: Final = ( + "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9vY3RldC1zdHJlYW07dW5pZmllZF9pZCxyZWdyZXNzaW9uLXRlc3QtZmlsZQ==" +) + + +class _BatchBody(TypedDict): + model: ReadOnly[str] + messages: ReadOnly[Sequence[ChatCompletionUserMessage]] + + +class _BatchLine(TypedDict): + custom_id: ReadOnly[str] + method: ReadOnly[Literal["POST"]] + url: ReadOnly[Literal["/v1/chat/completions"]] + body: ReadOnly[_BatchBody] + + +@pytest.fixture +def openai_files(monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter) -> respx.MockRouter: + monkeypatch.setenv("OPENAI_API_KEY", "sk-test") + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + return respx_mock + + +def _batch_rows(messages: Sequence[str]) -> tuple[_BatchLine, ...]: + return tuple( + _BatchLine( + custom_id=f"request-{i}", + method="POST", + url="/v1/chat/completions", + body=_BatchBody(model=_BATCH_MODEL, messages=(ChatCompletionUserMessage(role="user", content=message),)), + ) + for i, message in enumerate(messages, start=1) + ) + + +def _serve_file(router: respx.MockRouter, file_id: str, rows: Sequence[_BatchLine]) -> respx.Route: + jsonl: Final = "\n".join(json.dumps(row) for row in rows) + return router.get(f"https://api.openai.com/v1/files/{file_id}/content").mock( + return_value=httpx.Response(200, content=jsonl.encode()) + ) + + +def _token_counter_total(rows: Sequence[_BatchLine]) -> int: + return sum(litellm.token_counter(model=row["body"]["model"], messages=row["body"]["messages"]) for row in rows) + + +def _create_batch_data(input_file_id: str) -> LiteLLMBatchCreateRequest: + return LiteLLMBatchCreateRequest(model=_BATCH_MODEL, input_file_id=input_file_id) + + +@pytest.mark.asyncio +async def test_count_input_file_usage_matches_token_counter(openai_files: respx.MockRouter): + _, _, batch_limiter = _make_limiters() + rows: Final = _batch_rows(("Hello", "Hi there", "Hey")) + content_route: Final = _serve_file(openai_files, "file-abc123", rows) + + usage: Final = await batch_limiter.count_input_file_usage(file_id="file-abc123", custom_llm_provider="openai") + + assert content_route.call_count == 1 + assert usage.request_count == 3 + assert usage.total_tokens == _token_counter_total(rows) + + +@pytest.mark.asyncio +async def test_batch_rate_limit_single_file_under_and_over_tpm(openai_files: respx.MockRouter): + small_rows: Final = _batch_rows(("Hello", "Hi", "Hey")) + big_rows: Final = _batch_rows( + ("This is a longer message that will consume more tokens from the rate limit. " * 100,) * 3 + ) + _serve_file(openai_files, "file-small", small_rows) + _serve_file(openai_files, "file-big", big_rows) + user_api_key_dict: Final = UserAPIKeyAuth(api_key="test-key-123", tpm_limit=200, rpm_limit=10) + _, _, small_limiter = _make_limiters() + + data_small: Final = dict(_create_batch_data("file-small")) + result: Final = await small_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data_small, + call_type="acreate_batch", + ) + + assert result is data_small + assert data_small["_batch_token_count"] == _token_counter_total(small_rows) + assert data_small["_batch_request_count"] == 3 + + _, _, big_limiter = _make_limiters() + with pytest.raises(HTTPException) as exc_info: + await big_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=dict(_create_batch_data("file-big")), + call_type="acreate_batch", + ) + assert exc_info.value.status_code == 429 + assert "tokens" in str(exc_info.value.detail).lower() + + +@pytest.mark.asyncio +async def test_batch_rate_limit_cumulative_tpm_rejects_second_request(openai_files: respx.MockRouter): + _, _, batch_limiter = _make_limiters() + user_api_key_dict: Final = UserAPIKeyAuth(api_key="test-key-456", tpm_limit=200, rpm_limit=10) + first_rows: Final = _batch_rows(("This message has some content to reach about 100 tokens total. " * 4,) * 2) + second_rows: Final = _batch_rows( + ("This is another message with more content to exceed the remaining limit. " * 11,) * 2 + ) + _serve_file(openai_files, "file-1", first_rows) + _serve_file(openai_files, "file-2", second_rows) + first_tokens: Final = _token_counter_total(first_rows) + assert first_tokens <= 200 < first_tokens + _token_counter_total(second_rows) + + first_data: Final = dict(_create_batch_data("file-1")) + first_result: Final = await batch_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=first_data, + call_type="acreate_batch", + ) + assert first_result is first_data + assert first_data["_batch_token_count"] == first_tokens + + with pytest.raises(HTTPException) as exc_info: + await batch_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=dict(_create_batch_data("file-2")), + call_type="acreate_batch", + ) + assert exc_info.value.status_code == 429 + assert "tokens" in str(exc_info.value.detail).lower() + + +@pytest.mark.asyncio +async def test_batch_rate_limiter_reads_a_provider_file_with_user_context(openai_files: respx.MockRouter): + _, _, batch_limiter = _make_limiters() + user_api_key_dict: Final = UserAPIKeyAuth( + api_key="test-key-managed-files", user_id="test-user-abc123", tpm_limit=500, rpm_limit=10 + ) + rows: Final = _batch_rows(("This is a test message for batch rate limiting with managed files. " * 5,) * 3) + content_route: Final = _serve_file(openai_files, "file-abc123", rows) + + data: Final = dict(_create_batch_data("file-abc123")) + result: Final = await batch_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="acreate_batch", + ) + + assert content_route.call_count == 1 + assert result is data + assert data["_batch_token_count"] == _token_counter_total(rows) + assert data["_batch_request_count"] == 3 + + +@pytest.mark.asyncio +async def test_batch_rate_limiter_without_user_context(openai_files: respx.MockRouter): + _, _, batch_limiter = _make_limiters() + rows: Final = _batch_rows(("Hello",)) + content_route: Final = _serve_file(openai_files, "file-abc123", rows) + + usage_without_context: Final = await batch_limiter.count_input_file_usage( + file_id="file-abc123", custom_llm_provider="openai", user_api_key_dict=None + ) + usage_with_context: Final = await batch_limiter.count_input_file_usage( + file_id="file-abc123", + custom_llm_provider="openai", + user_api_key_dict=UserAPIKeyAuth(api_key="test-key", user_id="test-user-123"), + ) + + assert content_route.call_count == 2 + assert usage_without_context.request_count == usage_with_context.request_count == 1 + assert usage_without_context.total_tokens == usage_with_context.total_tokens == _token_counter_total(rows) + + +@pytest.mark.asyncio +async def test_managed_file_is_read_through_the_managed_files_hook_with_user_context( + openai_files: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + + from litellm import Router + from litellm.models.managed_files import LiteLLM_ManagedFileTable + from litellm.proxy import proxy_server + from litellm.proxy.openai_files_endpoints.common_utils import is_base64_encoded_unified_file_id + from litellm.proxy.utils import ProxyLogging + + assert is_base64_encoded_unified_file_id(_MANAGED_FILE_ID) + rows: Final = _batch_rows(("Test message for regression",)) + provider_route: Final = _serve_file(openai_files, "file-provider-1", rows) + standard_route: Final = _serve_file(openai_files, "file-abc123", rows) + unrouted_managed_read: Final = _serve_file(openai_files, _MANAGED_FILE_ID, rows) + file_cache: Final = InternalUsageCache(dual_cache=DualCache()) + await file_cache.async_set_cache( + key=_MANAGED_FILE_ID, + value=LiteLLM_ManagedFileTable( + unified_file_id=_MANAGED_FILE_ID, + model_mappings={"deployment-1": "file-provider-1"}, + flat_model_file_ids=["file-provider-1"], + created_by="test-user-regression", + ).model_dump(), + litellm_parent_otel_span=None, + ) + proxy_logging: Final = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging.proxy_hook_mapping["managed_files"] = _PROXY_LiteLLMManagedFiles( + internal_usage_cache=file_cache, prisma_client=None + ) + router: Final = Router( + model_list=[ + { + "model_name": _BATCH_MODEL, + "litellm_params": {"model": f"openai/{_BATCH_MODEL}", "api_key": "sk-test"}, + "model_info": {"id": "deployment-1"}, + } + ] + ) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", proxy_logging) + monkeypatch.setattr(proxy_server, "llm_router", router) + _, _, batch_limiter = _make_limiters() + user_api_key_dict: Final = UserAPIKeyAuth( + api_key="test-key-regression", user_id="test-user-regression", tpm_limit=1000, rpm_limit=10 + ) + + managed_usage: Final = await batch_limiter.count_input_file_usage( + file_id=_MANAGED_FILE_ID, custom_llm_provider="openai", user_api_key_dict=user_api_key_dict + ) + standard_usage: Final = await batch_limiter.count_input_file_usage( + file_id="file-abc123", custom_llm_provider="openai", user_api_key_dict=user_api_key_dict + ) + + assert provider_route.call_count == 1 + assert standard_route.call_count == 1 + assert not unrouted_managed_read.called + assert managed_usage.request_count == standard_usage.request_count == 1 + assert managed_usage.total_tokens == standard_usage.total_tokens == _token_counter_total(rows) diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index fee53320802..e9ad63a701e 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -1,5 +1,6 @@ import asyncio -from collections.abc import Callable, Mapping +from collections.abc import AsyncGenerator, Callable, Mapping +from contextlib import asynccontextmanager from datetime import datetime, timedelta, timezone from types import SimpleNamespace from typing import Final @@ -49,7 +50,7 @@ from litellm.proxy.lens.models import ( TraceIdentity, Worker, ) -from litellm.proxy.lens.repository import DueLens, Row +from litellm.proxy.lens.repository import Database, DueLens, Row from litellm.proxy.lens.signals import SignalConfig, StoredTraceSignal from litellm.proxy.lens.state import claim_job, queue_job, replace_job from litellm.rust_bridge.trace.generated.models import ExecutionRow, LensSampleParams @@ -69,6 +70,10 @@ class ResultDatabase: self.stored = stored self.completed: tuple[ReviewVersion, ...] = () + @asynccontextmanager + async def transaction(self) -> AsyncGenerator[Database]: + yield self + async def query_raw(self, query: str, *args: object) -> tuple[Row, ...]: if query.startswith("SELECT data FROM"): return (Row(data=self.stored.model_dump(mode="json")),) diff --git a/tests/unit/proxy/lens/test_feedback_endpoints.py b/tests/unit/proxy/lens/test_feedback_endpoints.py index 3ee7c79f169..71d60e22eea 100644 --- a/tests/unit/proxy/lens/test_feedback_endpoints.py +++ b/tests/unit/proxy/lens/test_feedback_endpoints.py @@ -94,34 +94,31 @@ class FakeClickHouse(ClickHouseStorage): and self._visible(str(r["TeamId"]), str(r["ApiKeyHash"]), parameters) ) case LensFeedbackSummaryParams(): - live = tuple( + live: Final = tuple( r for r in self._latest() if r["TraceId"] in parameters.trace_ids and self._visible(str(r["TeamId"]), str(r["ApiKeyHash"]), parameters) ) - keys = sorted({(str(r["TeamId"]), str(r["ApiKeyHash"]), str(r["TraceId"])) for r in live}) - return tuple( - FeedbackSummaryRow( - trace_id=trace, - trace_ref=ref(team, key, trace), - count=len(scores), - average=sum(scores) / len(scores), - lowest=min(scores), - ) - for team, key, trace in keys - for scores in [ - [ - int(str(r["Score"])) - for r in live - if (r["TeamId"], r["ApiKeyHash"], r["TraceId"]) == (team, key, trace) - ] - ] - ) + keys: Final = sorted({(str(r["TeamId"]), str(r["ApiKeyHash"]), str(r["TraceId"])) for r in live}) + return tuple(_summary_row(live, team, key, trace) for team, key, trace in keys) case _: raise AssertionError(f"unexpected query {query.name}") +def _summary_row(live: tuple[Mapping[str, object], ...], team: str, key: str, trace: str) -> FeedbackSummaryRow: + scores: Final = tuple( + int(str(r["Score"])) for r in live if (r["TeamId"], r["ApiKeyHash"], r["TraceId"]) == (team, key, trace) + ) + return FeedbackSummaryRow( + trace_id=trace, + trace_ref=ref(team, key, trace), + count=len(scores), + average=sum(scores) / len(scores), + lowest=min(scores), + ) + + def store(**traces: tuple[tuple[str, str], ...]) -> ClickHouseFeedbackStore: return ClickHouseFeedbackStore(FakeClickHouse(traces or {"t1": (("team-a", "key-a"),)})) @@ -207,6 +204,20 @@ async def test_viewers_can_read_but_not_write_and_non_admins_cannot_read_in_lens assert (write.value.status_code, read.value.status_code) == (403, 403) +@pytest.mark.asyncio +async def test_a_caller_without_a_team_or_key_cannot_write_on_a_teamless_trace() -> None: + feedback: Final = store(t1=(("", "key-a"),)) + await submit_feedback(submission(9, "mine", user="customer-1"), ADMIN, feedback, T0) + + with pytest.raises(HTTPException) as write: + await submit_feedback(submission(1, "overwrite", user="customer-1"), INTERNAL, feedback, T0) + with pytest.raises(HTTPException) as delete: + await delete_feedback(FeedbackDeletion(trace_id="t1", user="customer-1"), INTERNAL, feedback, T0) + + assert (write.value.status_code, delete.value.status_code) == (403, 403) + assert [f.score for f in (await read_feedback(FeedbackTarget(trace_id="t1"), ADMIN, feedback)).feedback] == [9] + + @pytest.mark.asyncio async def test_tenant_comes_from_the_trace_and_author_defaults_to_the_caller() -> None: feedback: Final = store() diff --git a/tests/unit/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/unit/proxy/openai_files_endpoint/test_files_common_utils.py index 50768e48d43..180ece77a6e 100644 --- a/tests/unit/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/unit/proxy/openai_files_endpoint/test_files_common_utils.py @@ -3,7 +3,9 @@ from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest +from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles +from litellm.caching.caching import DualCache from litellm.proxy.openai_files_endpoints.common_utils import ( apply_unified_file_ids, get_credentials_for_model, @@ -11,7 +13,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( map_raw_file_ids_to_unified, ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError -from litellm.proxy.utils import handle_exception_on_proxy +from litellm.proxy.utils import InternalUsageCache, handle_exception_on_proxy from litellm.types.utils import LiteLLMBatch _RAW_MODEL_WITH_PROMPT: Final = "opus-4.6 Please summarize my medical records\nPatient has diabetes" @@ -514,3 +516,139 @@ class TestCompletedBatchSafeToRetire: ) def test_is_litellm_executed_batch_reads_the_llm_batch_id_prefix(decoded_unified_batch_id: str, executed: bool): assert is_litellm_executed_batch(decoded_unified_batch_id) is executed + + +def _managed_files_hook(prisma_client: MagicMock) -> _PROXY_LiteLLMManagedFiles: + return _PROXY_LiteLLMManagedFiles(internal_usage_cache=InternalUsageCache(DualCache()), prisma_client=prisma_client) + + +def _managed_object_prisma(db_batch_object: MagicMock | None) -> MagicMock: + prisma_client: Final = MagicMock() + prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=db_batch_object) + prisma_client.db.litellm_managedobjecttable.update = AsyncMock() + return prisma_client + + +@pytest.mark.asyncio +async def test_batch_status_sync_from_provider_to_database(caplog: pytest.LogCaptureFixture): + import json + import logging + + from litellm.proxy.openai_files_endpoints.common_utils import ( + get_batch_from_database, + update_batch_in_database, + ) + + batch_id: Final = "batch_test123" + unified_batch_id: Final = "litellm_proxy:test_unified_batch" + stored_row: Final = MagicMock( + unified_object_id=batch_id, + status="validating", + file_object=json.dumps( + { + "id": batch_id, + "object": "batch", + "status": "validating", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-test123", + "completion_window": "24h", + "created_at": 1234567890, + } + ), + ) + prisma_client: Final = _managed_object_prisma(stored_row) + managed_files: Final = _managed_files_hook(prisma_client) + logger: Final = logging.getLogger("test_batch_status_sync") + + db_batch_object, response_batch = await get_batch_from_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + managed_files_obj=managed_files, + prisma_client=prisma_client, + verbose_proxy_logger=logger, + ) + + prisma_client.db.litellm_managedobjecttable.find_first.assert_awaited_once_with( + where={"unified_object_id": batch_id} + ) + assert db_batch_object is stored_row + assert isinstance(response_batch, LiteLLMBatch) + assert response_batch.id == batch_id + assert response_batch.status == "validating" + assert response_batch.input_file_id == "file-test123" + + completed: Final = LiteLLMBatch( + id=batch_id, + object="batch", + status="completed", + endpoint="/v1/chat/completions", + input_file_id="file-test123", + completion_window="24h", + created_at=1234567890, + output_file_id="file-output123", + ) + with caplog.at_level(logging.INFO, logger=logger.name): + await update_batch_in_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + response=completed, + managed_files_obj=managed_files, + prisma_client=prisma_client, + verbose_proxy_logger=logger, + db_batch_object=db_batch_object, + operation="retrieve", + poller_owns_accounting=False, + ) + + update: Final = prisma_client.db.litellm_managedobjecttable.update + update.assert_awaited_once() + assert update.await_args.kwargs["where"] == {"unified_object_id": batch_id} + written: Final = update.await_args.kwargs["data"] + assert written["status"] == "complete" + assert written["batch_processed"] is True + assert written["updated_at"] is not None + assert json.loads(written["file_object"])["status"] == "completed" + assert json.loads(written["file_object"])["output_file_id"] == "file-output123" + assert f"Updating batch {batch_id} status from validating to completed" in caplog.messages + + +@pytest.mark.asyncio +async def test_batch_cancel_updates_database(caplog: pytest.LogCaptureFixture): + import json + import logging + + from litellm.proxy.openai_files_endpoints.common_utils import update_batch_in_database + + batch_id: Final = "batch_cancel_test" + cancelled: Final = LiteLLMBatch( + id=batch_id, + object="batch", + status="cancelled", + endpoint="/v1/chat/completions", + input_file_id="file-test123", + completion_window="24h", + created_at=1234567890, + cancelled_at=1234567999, + ) + prisma_client: Final = _managed_object_prisma(None) + logger: Final = logging.getLogger("test_batch_cancel_updates_database") + + with caplog.at_level(logging.INFO, logger=logger.name): + await update_batch_in_database( + batch_id=batch_id, + unified_batch_id="litellm_proxy:cancel_test", + response=cancelled, + managed_files_obj=_managed_files_hook(prisma_client), + prisma_client=prisma_client, + verbose_proxy_logger=logger, + operation="cancel", + ) + + update: Final = prisma_client.db.litellm_managedobjecttable.update + update.assert_awaited_once() + assert update.await_args.kwargs["where"] == {"unified_object_id": batch_id} + written: Final = update.await_args.kwargs["data"] + assert written["status"] == "cancelled" + assert "batch_processed" not in written + assert json.loads(written["file_object"])["cancelled_at"] == 1234567999 + assert f"Updating batch {batch_id} status to cancelled after cancel" in caplog.messages diff --git a/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py index 4ce31cbb449..793921255ec 100644 --- a/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/unit/proxy/openai_files_endpoint/test_files_endpoint.py @@ -2145,6 +2145,54 @@ def test_get_file_content_streams_openai_direct_path( proxy_logging_obj.post_call_failure_hook.assert_not_called() +@respx.mock +def test_get_file_content_forwards_upstream_download_headers(monkeypatch, llm_router: Router): + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + setup_proxy_logging_object(monkeypatch, llm_router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setenv("OPENAI_API_KEY", "sk-test") + + body: Final = b'{"prompt": "Hello", "completion": "Hi"}' + upstream_filename: Final = "mydata.jsonl" + upstream_request_id: Final = "req_upstream_123" + upstream_route: Final = respx.get("https://api.openai.com/v1/files/file-abc123/content").mock( + return_value=httpx.Response( + status_code=200, + content=body, + headers={ + "content-type": "application/octet-stream", + "content-length": str(len(body)), + "content-disposition": f'attachment; filename="{upstream_filename}"', + "x-request-id": upstream_request_id, + }, + ) + ) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="test-user", + ) + try: + response: Final = client.get( + "/v1/files/file-abc123/content", + headers={"Authorization": "Bearer test-key"}, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert upstream_route.call_count == 1 + assert response.content == body + assert response.headers["content-type"].startswith("application/octet-stream") + assert int(response.headers["content-length"]) == len(response.content) + assert upstream_filename in response.headers["content-disposition"] + assert response.headers["x-request-id"] == upstream_request_id + + def test_get_file_content_routed_provider_skips_streaming_when_resolved_provider_is_not_supported( mocker: MockerFixture, monkeypatch, llm_router: Router ): diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_assembly_passthrough_route.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_assembly_passthrough_route.py new file mode 100644 index 00000000000..a862d0fa108 --- /dev/null +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_assembly_passthrough_route.py @@ -0,0 +1,172 @@ +import asyncio +import json +from collections.abc import Mapping +from datetime import datetime +from typing import Final, cast + +import httpx +import pytest +import respx +from fastapi import Request, Response +from starlette.types import Message + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import assemblyai_proxy_route +from litellm.types.utils import StandardLoggingPayload + +_KEY: Final = "synthetic-assemblyai-key" +_UPSTREAM: Final = "https://api.assemblyai.com/v2/transcript" +_AUDIO_URL: Final = "https://assembly.ai/wildfires.mp3" + + +def _is_transcript_log(payload: StandardLoggingPayload, transcript_id: str) -> bool: + response: Final = payload["response"] + return isinstance(response, dict) and response.get("id") == transcript_id + + +class _SuccessRecorder(CustomLogger): + def __init__(self, transcript_id: str, loop: asyncio.AbstractEventLoop) -> None: + super().__init__() + self.transcript_id: Final = transcript_id + self.loop: Final = loop + self.logged: Final = asyncio.Event() + self.payloads: tuple[StandardLoggingPayload, ...] = () + + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + payload: Final = cast(StandardLoggingPayload, kwargs["standard_logging_object"]) + self.payloads = (*self.payloads, payload) + if _is_transcript_log(payload, self.transcript_id): + self.loop.call_soon_threadsafe(self.logged.set) + + +def _request(method: str, path: str, body: bytes = b"") -> Request: + scope: Final = { + "type": "http", + "http_version": "1.1", + "method": method, + "scheme": "http", + "path": path, + "raw_path": path.encode(), + "root_path": "", + "query_string": b"", + "headers": [(b"content-type", b"application/json")], + "client": ("127.0.0.1", 51234), + "server": ("proxy.local", 4000), + "state": {}, + } + + async def receive() -> Message: + return {"type": "http.request", "body": body, "more_body": False} + + return Request(scope, receive) + + +def _install_recorder(monkeypatch: pytest.MonkeyPatch, transcript_id: str) -> _SuccessRecorder: + recorder: Final = _SuccessRecorder(transcript_id, asyncio.get_running_loop()) + monkeypatch.setenv("ASSEMBLYAI_API_KEY", _KEY) + monkeypatch.setattr(litellm, "_async_success_callback", [recorder]) + monkeypatch.setattr(litellm, "success_callback", []) + return recorder + + +async def _wait_for_success_log(recorder: _SuccessRecorder) -> StandardLoggingPayload: + await asyncio.wait_for(recorder.logged.wait(), 30) + logged: Final = tuple( + payload for payload in recorder.payloads if _is_transcript_log(payload, recorder.transcript_id) + ) + assert len(logged) == 1, f"AssemblyAI success log for {recorder.transcript_id} emitted {len(logged)} times" + return logged[0] + + +@pytest.mark.asyncio +async def test_assemblyai_transcribe_create_poll_delete( + respx_mock: respx.MockRouter, httpx_transport: None, monkeypatch: pytest.MonkeyPatch +) -> None: + recorder: Final = _install_recorder(monkeypatch, "tr_1") + create_route: Final = respx_mock.post(_UPSTREAM).mock( + return_value=httpx.Response(200, json={"id": "tr_1", "status": "queued", "audio_url": _AUDIO_URL}) + ) + poll_route: Final = respx_mock.get(f"{_UPSTREAM}/tr_1").mock( + return_value=httpx.Response(200, json={"id": "tr_1", "status": "completed", "text": "fires near town"}) + ) + delete_route: Final = respx_mock.delete(f"{_UPSTREAM}/tr_1").mock( + return_value=httpx.Response(200, json={"id": "tr_1", "status": "deleted"}) + ) + admin: Final = UserAPIKeyAuth(api_key="sk-master", user_role=LitellmUserRoles.PROXY_ADMIN) + create_body: Final = {"audio_url": _AUDIO_URL, "speech_models": ["universal-2"]} + + create: Final = await assemblyai_proxy_route( + endpoint="v2/transcript", + request=_request("POST", "/assemblyai/v2/transcript", json.dumps(create_body).encode()), + fastapi_response=Response(), + user_api_key_dict=admin, + ) + await _wait_for_success_log(recorder) + poll: Final = await assemblyai_proxy_route( + endpoint="v2/transcript/tr_1", + request=_request("GET", "/assemblyai/v2/transcript/tr_1"), + fastapi_response=Response(), + user_api_key_dict=admin, + ) + delete: Final = await assemblyai_proxy_route( + endpoint="v2/transcript/tr_1", + request=_request("DELETE", "/assemblyai/v2/transcript/tr_1"), + fastapi_response=Response(), + user_api_key_dict=admin, + ) + + assert isinstance(create, Response) and isinstance(poll, Response) and isinstance(delete, Response) + assert create.status_code == 200 + assert json.loads(create.body)["id"] == "tr_1" + assert poll.status_code == 200 + assert json.loads(poll.body)["status"] == "completed" + assert delete.status_code == 200 + assert json.loads(delete.body)["status"] == "deleted" + assert create_route.call_count == 1 + assert json.loads(create_route.calls[0].request.content) == create_body + assert delete_route.call_count == 1 + client_poll: Final = poll_route.calls[-1].request + assert client_poll.headers["authorization"] == _KEY + assert create_route.calls[0].request.headers["authorization"] == _KEY + assert delete_route.calls[0].request.headers["authorization"] == _KEY + + +@pytest.mark.asyncio +async def test_assemblyai_transcribe_with_non_admin_key_logs_key_identity( + respx_mock: respx.MockRouter, httpx_transport: None, monkeypatch: pytest.MonkeyPatch +) -> None: + recorder: Final = _install_recorder(monkeypatch, "tr_9") + create_route: Final = respx_mock.post(_UPSTREAM).mock( + return_value=httpx.Response(200, json={"id": "tr_9", "status": "queued", "audio_url": _AUDIO_URL}) + ) + respx_mock.get(f"{_UPSTREAM}/tr_9").mock( + return_value=httpx.Response(200, json={"id": "tr_9", "status": "completed", "text": "fires near town"}) + ) + non_admin: Final = UserAPIKeyAuth( + api_key="hashed-non-admin-key", + token="hashed-non-admin-key", + user_id="non-admin-user", + team_id="non-admin-team", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + response: Final = await assemblyai_proxy_route( + endpoint="v2/transcript", + request=_request("POST", "/assemblyai/v2/transcript", json.dumps({"audio_url": _AUDIO_URL}).encode()), + fastapi_response=Response(), + user_api_key_dict=non_admin, + ) + logged: Final = await _wait_for_success_log(recorder) + + assert isinstance(response, Response) + assert response.status_code == 200, response.body + assert create_route.call_count == 1 + assert create_route.calls[0].request.headers["authorization"] == _KEY + assert logged["metadata"]["user_api_key_hash"] == "hashed-non-admin-key" + assert logged["metadata"]["user_api_key_user_id"] == "non-admin-user" + assert logged["metadata"]["user_api_key_team_id"] == "non-admin-team" + assert logged["custom_llm_provider"] == "assemblyai" diff --git a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 1887021a53a..d329db75fb4 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -23,6 +23,7 @@ from starlette.datastructures import FormData import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail from tests._master_key import MASTER_KEY as SHARED_MASTER_KEY from litellm.caching.caching import DualCache from litellm.types.utils import CallTypesLiteral @@ -47,6 +48,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( get_vertex_base_url, is_azure_ai_search_service_level_index_create, gigachat_proxy_route, + handle_bedrock_passthrough_router_model, llm_passthrough_factory_proxy_route, milvus_proxy_route, mistral_proxy_route, @@ -61,9 +63,32 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( from litellm.proxy._types import LitellmUserRoles, SpecialHeaders, UserAPIKeyAuth from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials +def _assert_bedrock_processing_data_classification(data: dict[str, object], is_streaming: bool) -> None: + streaming_guardrail: Final = CustomGuardrail( + guardrail_name="streaming-only", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="streaming", + ) + non_streaming_guardrail: Final = CustomGuardrail( + guardrail_name="non-streaming-only", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="non_streaming", + ) + provider_body: Final = data["data"] + + assert streaming_guardrail.should_run_guardrail(data, GuardrailEventHooks.pre_call) is is_streaming + assert non_streaming_guardrail.should_run_guardrail(data, GuardrailEventHooks.pre_call) is not is_streaming + assert isinstance(provider_body, dict) + assert "is_streaming_request" not in provider_body + assert "litellm_server_streaming_classification" not in provider_body + + class TestVertexPassthroughGetVertexBaseUrl: """Module-local get_vertex_base_url (trailing slash); rules match common_utils.""" @@ -413,9 +438,7 @@ class TestVertexAIPassThroughHandler: # Mock the vertex handler for global location mock_handler = Mock() - mock_handler.get_default_base_target_url.return_value = ( - "https://aiplatform.googleapis.com/" - ) + mock_handler.get_default_base_target_url.return_value = "https://aiplatform.googleapis.com/" mock_get_handler.return_value = mock_handler # Mock create_pass_through_route to return a function that returns a mock response @@ -1237,9 +1260,7 @@ class TestVertexAIDiscoveryPassThroughHandler: # Mock the discovery handler mock_handler = Mock() - mock_handler.get_default_base_target_url.return_value = ( - "https://discoveryengine.googleapis.com" - ) + mock_handler.get_default_base_target_url.return_value = "https://discoveryengine.googleapis.com" mock_get_handler.return_value = mock_handler # Mock create_pass_through_route to return a function that returns a mock response @@ -1459,6 +1480,170 @@ class TestBedrockLLMProxyRoute: assert call_kwargs["model"] == "anthropic.claude-3-sonnet-20240229-v1:0" assert result == "success" + @pytest.mark.asyncio + @pytest.mark.parametrize( + "action, is_streaming", + [ + ("converse-stream", True), + ("invoke-with-response-stream", True), + ("converse", False), + ("invoke", False), + ], + ) + async def test_bedrock_direct_actions_classify_guardrail_stream_scope( + self, action: str, is_streaming: bool + ) -> None: + mock_request: Final = Mock() + mock_request.method = "POST" + mock_processor: Final = Mock() + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") + request_body: Final = {"messages": [{"role": "user", "content": "test"}]} + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._read_request_body", + return_value=request_body, + ), + patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing", + return_value=mock_processor, + ) as processor_constructor, + ): + result: Final = await bedrock_llm_proxy_route( + endpoint=f"model/test-model/{action}", + request=mock_request, + fastapi_response=Mock(), + user_api_key_dict=Mock(), + ) + + assert result == "success" + processing_data: Final = processor_constructor.call_args.kwargs["data"] + _assert_bedrock_processing_data_classification(processing_data, is_streaming) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "action, is_streaming", + [ + ("converse-stream", True), + ("invoke-with-response-stream", True), + ("converse", False), + ("invoke", False), + ], + ) + async def test_bedrock_router_actions_classify_guardrail_stream_scope( + self, action: str, is_streaming: bool + ) -> None: + mock_request: Final = Mock() + mock_request.method = "POST" + mock_processor: Final = Mock() + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") + request_body: Final = {"messages": [{"role": "user", "content": "test"}]} + + with patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing", + return_value=mock_processor, + ) as processor_constructor: + result: Final = await handle_bedrock_passthrough_router_model( + model="test-model", + endpoint=f"model/test-model/{action}", + request=mock_request, + request_body=request_body, + llm_router=Mock(), + user_api_key_dict=Mock(), + proxy_logging_obj=Mock(), + general_settings={}, + proxy_config=None, + select_data_generator=None, + user_model=None, + user_temperature=None, + user_request_timeout=None, + user_max_tokens=None, + user_api_base=None, + version=None, + ) + + assert result == "success" + processing_data: Final = processor_constructor.call_args.kwargs["data"] + _assert_bedrock_processing_data_classification(processing_data, is_streaming) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("model_id", "action"), + [ + ("my-converse-stream-model", "converse"), + ("my-invoke-with-response-stream-model", "invoke"), + ], + ) + async def test_bedrock_direct_model_id_does_not_imply_streaming(self, model_id: str, action: str) -> None: + mock_request: Final = Mock() + mock_request.method = "POST" + mock_processor: Final = Mock() + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") + request_body: Final = {"messages": [{"role": "user", "content": "test"}]} + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._read_request_body", + return_value=request_body, + ), + patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing", + return_value=mock_processor, + ) as processor_constructor, + ): + result: Final = await bedrock_llm_proxy_route( + endpoint=f"/model/{model_id}/{action}", + request=mock_request, + fastapi_response=Mock(), + user_api_key_dict=Mock(), + ) + + assert result == "success" + processing_data: Final = processor_constructor.call_args.kwargs["data"] + _assert_bedrock_processing_data_classification(processing_data, is_streaming=False) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("model_id", "action"), + [ + ("my-converse-stream-model", "converse"), + ("my-invoke-with-response-stream-model", "invoke"), + ], + ) + async def test_bedrock_router_model_id_does_not_imply_streaming(self, model_id: str, action: str) -> None: + mock_request: Final = Mock() + mock_request.method = "POST" + mock_processor: Final = Mock() + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") + request_body: Final = {"messages": [{"role": "user", "content": "test"}]} + + with patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing", + return_value=mock_processor, + ) as processor_constructor: + result: Final = await handle_bedrock_passthrough_router_model( + model=model_id, + endpoint=f"/model/{model_id}/{action}", + request=mock_request, + request_body=request_body, + llm_router=Mock(), + user_api_key_dict=Mock(), + proxy_logging_obj=Mock(), + general_settings={}, + proxy_config=None, + select_data_generator=None, + user_model=None, + user_temperature=None, + user_request_timeout=None, + user_max_tokens=None, + user_api_base=None, + version=None, + ) + + assert result == "success" + processing_data: Final = processor_constructor.call_args.kwargs["data"] + _assert_bedrock_processing_data_classification(processing_data, is_streaming=False) + @pytest.mark.asyncio async def test_bedrock_error_handling_returns_actual_error(self): """ @@ -1879,7 +2064,6 @@ class TestBedrockAgentRuntimePassthroughToggle: class TestBedrockAgentRuntimePassthroughVirtualKeyLeak: - VKEY: Final = "sk-litellm-victim-key" MASTER_KEY: Final = SHARED_MASTER_KEY ENDPOINT: Final = "knowledgebases/KB1234567/retrieve" @@ -2075,10 +2259,10 @@ class TestVLLMProxyRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model", return_value=True, ) - @patch("litellm.proxy.proxy_server.llm_router") # test-quality-ok: patching litellm internal for unit test isolation - async def test_vllm_proxy_route_with_router_model( - self, mock_llm_router, mock_is_router, mock_get_body - ): + @patch( + "litellm.proxy.proxy_server.llm_router" + ) # test-quality-ok: patching litellm internal for unit test isolation + async def test_vllm_proxy_route_with_router_model(self, mock_llm_router, mock_is_router, mock_get_body): mock_request = MagicMock(spec=Request) mock_request.method = "POST" mock_request.headers = {"content-type": "application/json"} @@ -2111,9 +2295,7 @@ class TestVLLMProxyRoute: @patch( # test-quality-ok: patching litellm internal for unit test isolation "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.llm_passthrough_factory_proxy_route" ) - async def test_vllm_proxy_route_fallback_to_factory( - self, mock_factory_route, mock_is_router, mock_get_body - ): + async def test_vllm_proxy_route_fallback_to_factory(self, mock_factory_route, mock_is_router, mock_get_body): mock_request = MagicMock(spec=Request) mock_fastapi_response = MagicMock(spec=Response) mock_user_api_key_dict = MagicMock() @@ -2140,10 +2322,10 @@ class TestGigachatProxyRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model", return_value=True, ) - @patch("litellm.proxy.proxy_server.llm_router") # test-quality-ok: patching litellm internal for unit test isolation - async def test_gigachat_proxy_route_with_router_model( - self, mock_llm_router, mock_is_router, mock_get_body - ): + @patch( + "litellm.proxy.proxy_server.llm_router" + ) # test-quality-ok: patching litellm internal for unit test isolation + async def test_gigachat_proxy_route_with_router_model(self, mock_llm_router, mock_is_router, mock_get_body): mock_request = MagicMock(spec=Request) mock_request.method = "POST" mock_request.headers = {"content-type": "application/json"} @@ -2396,21 +2578,25 @@ class TestGigachatProxyRoute: return _inner() - with patch.object( - processor, - "common_processing_pre_call_logic", - new=AsyncMock( - return_value=( - processor.data, - processor.data["litellm_logging_obj"], - ) + with ( + patch.object( + processor, + "common_processing_pre_call_logic", + new=AsyncMock( + return_value=( + processor.data, + processor.data["litellm_logging_obj"], + ) + ), + ), + patch( # test-quality-ok: patching litellm internal for unit test isolation + "litellm.proxy.common_request_processing.route_request", + new=_fake_route_request, + ), + patch( # test-quality-ok: patching litellm internal for unit test isolation + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.get_custom_headers", + return_value={"x-litellm-call-id": "call-123"}, ), - ), patch( # test-quality-ok: patching litellm internal for unit test isolation - "litellm.proxy.common_request_processing.route_request", - new=_fake_route_request, - ), patch( # test-quality-ok: patching litellm internal for unit test isolation - "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.get_custom_headers", - return_value={"x-litellm-call-id": "call-123"}, ): result = await processor.base_passthrough_process_llm_request( request=mock_request, @@ -4223,7 +4409,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: @pytest.mark.parametrize( ("credential", "authenticated"), [ - pytest.param("modified_key", UserAPIKeyAuth(api_key="modified_key"), id="custom-auth-echoing-opaque-credential"), + pytest.param( + "modified_key", UserAPIKeyAuth(api_key="modified_key"), id="custom-auth-echoing-opaque-credential" + ), pytest.param( LITELLM_JWT, UserAPIKeyAuth(api_key=LITELLM_JWT, user_id="jwt-subject"), @@ -4298,7 +4486,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: (b"x-goog-api-key", b"AIza-real-google-api-key"), (b"content-type", b"application/json"), ], - authenticated=UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN), + authenticated=UserAPIKeyAuth( + api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN + ), ) assert raised is None assert forwarded is not None @@ -4440,10 +4630,14 @@ class TestAnthropicPassthroughVirtualKeyLeak: raised, forwarded = await self._run( monkeypatch, [(header, value), (b"anthropic-version", b"2023-06-01"), (b"content-type", b"application/json")], - authenticated=UserAPIKeyAuth(api_key="sk-ant-api03-callers-own-key", user_role=LitellmUserRoles.INTERNAL_USER), + authenticated=UserAPIKeyAuth( + api_key="sk-ant-api03-callers-own-key", user_role=LitellmUserRoles.INTERNAL_USER + ), master_key=None, ) - assert raised is None, "with no master key the proxy authenticated nothing, so nothing of the caller's is a LiteLLM secret" + assert raised is None, ( + "with no master key the proxy authenticated nothing, so nothing of the caller's is a LiteLLM secret" + ) assert forwarded is not None assert forwarded.get(header.decode()) == value.decode() @@ -5231,7 +5425,9 @@ class TestTranscribeProxyRoute: ) -> None: with respx.mock(assert_all_called=False) as upstream: route = upstream.post(TRANSCRIBE_UPSTREAM) - response = transcribe_client.post("/transcribe/StartTranscriptionJob", json={**dict(self.START_JOB_BODY), **body}) + response = transcribe_client.post( + "/transcribe/StartTranscriptionJob", json={**dict(self.START_JOB_BODY), **body} + ) assert response.status_code == 403 assert member in response.json()["detail"] @@ -6873,9 +7069,7 @@ class TestAzureBodyModelGroupRelay: AZURE_SPEECH_SHORT_AUDIO_ENDPOINT: Final = "/speech/recognition/conversation/cognitiveservices/v1" AZURE_SPEECH_BATCH_ENDPOINT: Final = "/speechtotext/v3.2/transcriptions" AZURE_SPEECH_FAST_ENDPOINT: Final = "/speechtotext/transcriptions:transcribe" -AZURE_SPEECH_PCM16_HEADER: Final = ( - b"RIFF\x24\x0c\x00\x00WAVEfmt \x10\x00\x00\x00\x01\x00\x01\x00\x80\x3e\x00\x00\x00\x7d\x00\x00\x02\x00\x10\x00data\x00\x0c\x00\x00" -) +AZURE_SPEECH_PCM16_HEADER: Final = b"RIFF\x24\x0c\x00\x00WAVEfmt \x10\x00\x00\x00\x01\x00\x01\x00\x80\x3e\x00\x00\x00\x7d\x00\x00\x02\x00\x10\x00data\x00\x0c\x00\x00" AZURE_SPEECH_WAV_BYTES: Final = AZURE_SPEECH_PCM16_HEADER + b"\x00" * 3072 AZURE_SPEECH_WAV_SECONDS: Final = 3072 / (16000 * 2) AZURE_SPEECH_NON_UTF8_WAV_BYTES: Final = AZURE_SPEECH_PCM16_HEADER + bytes(range(256)) * 12 @@ -6912,9 +7106,9 @@ class TestAzureSpeechProxyRoute: def test_short_audio_forwards_raw_wav_bytes_with_server_key(self, azure_speech_client: TestClient) -> None: with respx.mock(assert_all_called=True) as upstream: - route = upstream.post( - f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}" - ).mock(return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT)) + route = upstream.post(f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}").mock( + return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT) + ) response = azure_speech_client.post( f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", @@ -7390,9 +7584,9 @@ class TestAzureSpeechRawBodyThroughRealAuth: self, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, body: bytes ) -> None: with respx.mock(assert_all_called=True) as upstream, caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"): - route = upstream.post( - f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}" - ).mock(return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT)) + route = upstream.post(f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}").mock( + return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT) + ) response = self._post_wav( monkeypatch, f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", "sk-master-key", body=body @@ -7416,9 +7610,9 @@ class TestAzureSpeechRawBodyThroughRealAuth: ) -> None: boundary: Final = "lit7939boundary" multipart_body: Final = ( - f"--{boundary}\r\nContent-Disposition: form-data; name=\"definition\"\r\n\r\n".encode() + f'--{boundary}\r\nContent-Disposition: form-data; name="definition"\r\n\r\n'.encode() + json.dumps({"locales": ["en-US"]}).encode() - + f"\r\n--{boundary}\r\nContent-Disposition: form-data; name=\"audio\"; filename=\"eagle.wav\"\r\n" + + f'\r\n--{boundary}\r\nContent-Disposition: form-data; name="audio"; filename="eagle.wav"\r\n' "Content-Type: audio/wav\r\n\r\n".encode() + AZURE_SPEECH_NON_UTF8_WAV_BYTES + f"\r\n--{boundary}--\r\n".encode() @@ -8170,9 +8364,7 @@ class TestTinyFishProxyRoute: assert response.status_code == 200 assert json.loads(route.calls.last.request.content)["use_vault"] is True - def test_returns_401_on_missing_api_key( - self, tinyfish_client: TestClient, monkeypatch: pytest.MonkeyPatch - ) -> None: + def test_returns_401_on_missing_api_key(self, tinyfish_client: TestClient, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.delenv("TINYFISH_API_KEY") with respx.mock: diff --git a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py index b604de5ea97..7d0bdb74279 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -5,7 +5,7 @@ import logging import os import sys import zlib -from collections.abc import Callable, Mapping +from collections.abc import AsyncIterator, Callable, Mapping from contextlib import ExitStack, contextmanager from dataclasses import dataclass from io import BytesIO @@ -27,6 +27,7 @@ from starlette.datastructures import UploadFile as StarletteUploadFile import litellm from litellm._logging import verbose_proxy_logger from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._lazy_features import LazyFeature, attach_lazy_features @@ -55,6 +56,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import ( ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.types import utils as types_utils +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, @@ -1638,6 +1640,223 @@ async def test_pass_through_request_streaming_marks_logging_obj_as_stream(): assert logging_obj.model_call_details["stream"] is True +@pytest.mark.asyncio +@pytest.mark.parametrize("body_stream", [None, True], ids=["stream-absent", "stream-true"]) +async def test_pass_through_request_preserves_caller_streaming_request_field(body_stream): + captured_hook_data: dict[str, object] = {} + + async def capture_pre_call(user_api_key_dict, data, call_type, endpoint_type: EndpointType): + captured_hook_data.update(data) + return data + + with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client" + ) as mock_get_client: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor" + ) as mock_chunk_processor: + mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=capture_pre_call) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) + + upstream_response = MagicMock() + upstream_response.status_code = 200 + upstream_response.headers = {} + upstream_response.raise_for_status = MagicMock() + + async_client = MagicMock() + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) + mock_get_client.return_value = MagicMock(client=async_client) + + async def _empty_chunks(*args, **kwargs): + return + yield # pragma: no cover + + mock_chunk_processor.return_value = _empty_chunks() + + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = httpx.URL( + "http://test-proxy.com/gemini/v1beta/models/gemini-pro:streamGenerateContent" + ) + mock_request.scope = {"path": "/gemini/v1beta/models/gemini-pro:streamGenerateContent"} + request_body: Final = { + "contents": [{"parts": [{"text": "hi"}]}], + "is_streaming_request": "caller-value", + **({"stream": True} if body_stream is True else {}), + } + mock_request.body = AsyncMock(return_value=json.dumps(request_body).encode()) + mock_request.headers = Headers({"content-type": "application/json"}) + mock_request.query_params = QueryParams({}) + + await pass_through_request( + request=mock_request, + target="https://generativelanguage.googleapis.com/v1beta/models/gemini-pro:streamGenerateContent", + custom_headers={}, + user_api_key_dict=MagicMock(), + stream=True, + ) + + assert captured_hook_data.get("is_streaming_request") == "caller-value" + assert captured_hook_data.get("stream") is body_stream + + upstream_json = async_client.build_request.call_args.kwargs["json"] + assert upstream_json["is_streaming_request"] == "caller-value" + assert "litellm_server_streaming_classification" not in upstream_json + assert upstream_json["contents"] == request_body["contents"] + + +@pytest.mark.asyncio +async def test_streaming_pass_through_drops_marker_after_hook_rebuilds_body_from_json(): + async def json_rebuilding_pre_call(user_api_key_dict, data, call_type, endpoint_type: EndpointType): + rebuilt = json.loads(json.dumps({k: v for k, v in data.items() if k != "litellm_logging_obj"})) + return {**rebuilt, "litellm_logging_obj": data["litellm_logging_obj"]} + + with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client" + ) as mock_get_client: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor" + ) as mock_chunk_processor: + mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=json_rebuilding_pre_call) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) + + upstream_response = MagicMock() + upstream_response.status_code = 200 + upstream_response.headers = {} + upstream_response.raise_for_status = MagicMock() + + async_client = MagicMock() + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) + mock_get_client.return_value = MagicMock(client=async_client) + + async def _empty_chunks(*args, **kwargs): + return + yield # pragma: no cover + + mock_chunk_processor.return_value = _empty_chunks() + + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = httpx.URL("http://test-proxy.com/openai/v1/chat/completions") + mock_request.scope = {"path": "/openai/v1/chat/completions"} + request_body: Final = { + "model": "gpt-5-mini", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + } + mock_request.body = AsyncMock(return_value=json.dumps(request_body).encode()) + mock_request.headers = Headers({"content-type": "application/json"}) + mock_request.query_params = QueryParams({}) + + await pass_through_request( + request=mock_request, + target="https://api.openai.com/v1/chat/completions", + custom_headers={}, + user_api_key_dict=MagicMock(), + ) + + upstream_json = async_client.build_request.call_args.kwargs["json"] + assert upstream_json == request_body, upstream_json + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("route_stream", "body_stream", "expected_streaming"), + [ + (True, False, False), + (False, True, True), + (True, None, True), + (None, None, False), + ], + ids=["body-disables-route-stream", "body-enables-streaming", "route-enables-absent-body", "both-absent"], +) +async def test_passthrough_guardrails_follow_effective_relay_stream_decision( + route_stream: bool | None, + body_stream: bool | None, + expected_streaming: bool, +): + async def return_pre_call_data( + user_api_key_dict: UserAPIKeyAuth, + data: dict[str, object], + call_type: str, + endpoint_type: EndpointType, + ) -> dict[str, object]: + return dict(data) + + with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client" + ) as mock_get_client: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor" + ) as mock_chunk_processor: + mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=return_pre_call_data) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) + + upstream_response = MagicMock() + upstream_response.status_code = 200 + upstream_response.headers = {} + upstream_response.aread = AsyncMock(return_value=b"{}") + upstream_response.text = "{}" + upstream_response.raise_for_status = MagicMock() + + async_client = MagicMock() + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) + async_client.request = AsyncMock(return_value=upstream_response) + mock_get_client.return_value = MagicMock(client=async_client) + + async def empty_chunks() -> AsyncIterator[bytes]: + yield b"" + + mock_chunk_processor.return_value = empty_chunks() + + request_body: Final = { + "message": "hello", + **({"stream": body_stream} if body_stream is not None else {}), + } + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = httpx.URL("http://test-proxy.com/guardrail-stream-scope") + mock_request.scope = {"path": "/guardrail-stream-scope"} + mock_request.body = AsyncMock(return_value=json.dumps(request_body).encode()) + mock_request.headers = Headers({"content-type": "application/json"}) + mock_request.query_params = QueryParams({}) + + await pass_through_request( + request=mock_request, + target="http://upstream.test/guardrail-stream-scope", + custom_headers={}, + user_api_key_dict=MagicMock(), + stream=route_stream, + ) + hook_data: Final[dict[str, object]] = mock_proxy_logging.pre_call_hook.call_args.kwargs["data"] + + streaming_guardrail: Final = CustomGuardrail( + guardrail_name="streaming-only", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="streaming", + ) + non_streaming_guardrail: Final = CustomGuardrail( + guardrail_name="non-streaming-only", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="non_streaming", + ) + assert streaming_guardrail.should_run_guardrail(hook_data, GuardrailEventHooks.pre_call) is expected_streaming + assert ( + non_streaming_guardrail.should_run_guardrail(hook_data, GuardrailEventHooks.pre_call) is not expected_streaming + ) + + @pytest.mark.asyncio async def test_pass_through_request_sse_response_marks_logging_obj_as_stream(): """ diff --git a/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_live_passthrough.py b/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_live_passthrough.py index f6f63aca71d..9c372380749 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_live_passthrough.py +++ b/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_live_passthrough.py @@ -1,17 +1,28 @@ -from collections.abc import Sequence +import asyncio +import json +import uuid +from collections.abc import Mapping, Sequence from datetime import datetime +from typing import Final, cast from unittest.mock import MagicMock, patch +import httpx import litellm import pytest +import respx +from fastapi import Request, Response +from starlette.types import Message from typing_extensions import NotRequired, ReadOnly, TypedDict +from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler import ( VertexAILivePassthroughLoggingHandler, ) +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import vertex_proxy_route from litellm.proxy.pass_through_endpoints.success_handler import PassThroughEndpointLogging -from litellm.types.utils import CostBreakdown, LlmProviders, Usage +from litellm.types.utils import CostBreakdown, LlmProviders, StandardLoggingPayload, Usage class _LiveTurn(TypedDict): @@ -887,3 +898,95 @@ class TestVertexAILivePassthroughErrorHandling: assert "result" in result assert "kwargs" in result + + +class _SuccessRecorder(CustomLogger): + def __init__(self, api_base: str, loop: asyncio.AbstractEventLoop) -> None: + super().__init__() + self.api_base: Final = api_base + self.loop: Final = loop + self.logged: Final = asyncio.Event() + self.payloads: tuple[StandardLoggingPayload, ...] = () + + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + payload: Final = cast(StandardLoggingPayload, kwargs["standard_logging_object"]) + self.payloads = (*self.payloads, payload) + if payload["api_base"] == self.api_base: + self.loop.call_soon_threadsafe(self.logged.set) + + +def _generate_content_endpoint(project: str) -> str: + return f"v1/projects/{project}/locations/us-central1/publishers/google/models/gemini-2.0-flash:generateContent" + + +def _vertex_request(endpoint: str, body: bytes) -> Request: + path: Final = f"/vertex_ai/{endpoint}" + scope: Final = { + "type": "http", + "http_version": "1.1", + "method": "POST", + "scheme": "http", + "path": path, + "raw_path": path.encode(), + "root_path": "", + "query_string": b"", + "headers": [(b"content-type", b"application/json"), (b"authorization", b"Bearer client-google-token")], + "client": ("127.0.0.1", 51234), + "server": ("proxy.local", 4000), + "state": {}, + } + + async def receive() -> Message: + return {"type": "http.request", "body": body, "more_body": False} + + return Request(scope, receive) + + +@pytest.mark.asyncio +async def test_vertex_ai_generate_content_spendlog( + respx_mock: respx.MockRouter, httpx_transport: None, monkeypatch: pytest.MonkeyPatch +) -> None: + endpoint: Final = _generate_content_endpoint(f"p-{uuid.uuid4().hex}") + upstream_url: Final = f"https://us-central1-aiplatform.googleapis.com/{endpoint}" + recorder: Final = _SuccessRecorder(upstream_url, asyncio.get_running_loop()) + monkeypatch.setattr(litellm, "_async_success_callback", [recorder]) + monkeypatch.setattr(litellm, "success_callback", []) + contents: Final = [{"role": "user", "parts": [{"text": "hi"}]}] + route: Final = respx_mock.post(upstream_url).mock( + return_value=httpx.Response( + 200, + json={ + "candidates": [ + {"content": {"role": "model", "parts": [{"text": "hello vertex"}]}, "finishReason": "STOP"} + ], + "usageMetadata": {"promptTokenCount": 9, "candidatesTokenCount": 6, "totalTokenCount": 15}, + }, + ) + ) + + response: Final = await vertex_proxy_route( + endpoint=endpoint, + request=_vertex_request(endpoint, json.dumps({"contents": contents}).encode()), + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token", user_role=LitellmUserRoles.PROXY_ADMIN), + ) + assert isinstance(response, Response) + assert response.status_code == 200 + call_id: Final = response.headers.get("x-litellm-call-id") + assert call_id + await asyncio.wait_for(recorder.logged.wait(), 30) + assert route.call_count == 1 + outbound: Final = route.calls[0].request + assert outbound.headers["authorization"] == "Bearer client-google-token" + assert json.loads(outbound.content)["contents"] == contents + matching: Final = tuple(payload for payload in recorder.payloads if payload["api_base"] == upstream_url) + assert len(matching) == 1, recorder.payloads + logged: Final = matching[0] + assert logged["id"] == call_id + assert logged["response_cost"] > 0 + assert "gemini" in logged["model"] + assert logged["custom_llm_provider"] == "vertex_ai" + assert logged["prompt_tokens"] == 9 + assert logged["completion_tokens"] == 6 diff --git a/tests/unit/proxy/policy_engine/test_pipeline_executor.py b/tests/unit/proxy/policy_engine/test_pipeline_executor.py index 6aa7eca0f15..d13abbd379b 100644 --- a/tests/unit/proxy/policy_engine/test_pipeline_executor.py +++ b/tests/unit/proxy/policy_engine/test_pipeline_executor.py @@ -96,6 +96,20 @@ class AlwaysPassGuardrail(CustomGuardrail): return None +class StreamScopedPassGuardrail(CustomGuardrail): + def __init__(self, guardrail_name: str, stream_scope: object): + super().__init__( + guardrail_name=guardrail_name, + event_hook="pre_call", + default_on=True, + stream_scope=stream_scope, + ) + self.calls = 0 + + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + self.calls += 1 + + class PassthroughBlockGuardrail(CustomGuardrail): """Mock guardrail that blocks using the legacy passthrough contract.""" @@ -378,6 +392,46 @@ async def test_block_carries_original_guardrail_exception(monkeypatch): assert result.original_exception.detail == "Content policy violation" +@pytest.mark.asyncio +async def test_pipeline_step_honors_stream_scope(monkeypatch): + stream_only = StreamScopedPassGuardrail(guardrail_name="stream-only", stream_scope="streaming") + later = AlwaysFailGuardrail(guardrail_name="later-block") + monkeypatch.setattr(litellm, "callbacks", [stream_only, later]) + steps = [ + PipelineStep(guardrail="stream-only", on_fail="block", on_pass="allow"), + PipelineStep(guardrail="later-block", on_fail="block", on_pass="allow"), + ] + + skipped = await PipelineExecutor.execute_steps( + steps=steps, + mode="pre_call", + data={"messages": [{"role": "user", "content": "hi"}]}, + user_api_key_dict=MagicMock(), + call_type="completion", + policy_name="stream-scope", + ) + assert stream_only.calls == 0 + assert later.calls == 1 + assert skipped.step_results[0].outcome == "skip" + assert skipped.step_results[0].action_taken == "next" + assert skipped.terminal_action == "block" + + stream_only.calls = 0 + later.calls = 0 + ran = await PipelineExecutor.execute_steps( + steps=steps, + mode="pre_call", + data={"messages": [{"role": "user", "content": "hi"}], "stream": True}, + user_api_key_dict=MagicMock(), + call_type="completion", + policy_name="stream-scope", + ) + assert stream_only.calls == 1 + assert later.calls == 0 + assert ran.step_results[0].outcome == "pass" + assert ran.terminal_action == "allow" + + @pytest.mark.asyncio async def test_unsupported_mode_yields_error_outcome_without_exception(monkeypatch): """An unexpected hook mode must surface as an error outcome (carrying no diff --git a/tests/unit/proxy/proxy_server/test_exception_handlers.py b/tests/unit/proxy/proxy_server/test_exception_handlers.py index 0aff43057f9..f6efdecc884 100644 --- a/tests/unit/proxy/proxy_server/test_exception_handlers.py +++ b/tests/unit/proxy/proxy_server/test_exception_handlers.py @@ -5,6 +5,7 @@ Pins covered: - ``_close_dangling_otel_server_span`` - ``otel_request_validation_exception_handler`` - ``otel_unhandled_exception_handler`` +- ``otlp_http_exception_handler`` """ from __future__ import annotations @@ -16,8 +17,12 @@ from unittest.mock import MagicMock import httpx import pytest -from fastapi import HTTPException, Request +from fastapi import Depends, FastAPI, HTTPException, Request from fastapi.exceptions import RequestValidationError +from fastapi.testclient import TestClient +from starlette.testclient import WebSocketDenialResponse +from starlette.types import Message +from starlette.websockets import WebSocket, WebSocketDisconnect from litellm.proxy._types import ProxyException from litellm.proxy.proxy_server import ( @@ -25,6 +30,7 @@ from litellm.proxy.proxy_server import ( openai_exception_handler, otel_request_validation_exception_handler, otel_unhandled_exception_handler, + otlp_http_exception_handler, ) from .conftest import normalize @@ -507,6 +513,7 @@ async def test_otlp_auth_errors_hide_internal_details_and_survive_missing_native if isinstance(error, ProxyException) else await otlp_http_exception_handler(request, error) ) + assert response is not None assert response.status_code == (401 if isinstance(error, ProxyException) else 403) assert response.headers["content-type"].startswith(media_type) message: Final = ( @@ -516,3 +523,98 @@ async def test_otlp_auth_errors_hide_internal_details_and_survive_missing_native ) expected: Final = "Unauthorized" if isinstance(error, ProxyException) else "Forbidden" assert message == (expected if native_available or media_type == "application/json" else "") + + +def _websocket(sent: list[Message]) -> WebSocket: + async def receive() -> Message: + return {"type": "websocket.connect"} + + async def send(message: Message) -> None: + sent.append(message) + + return WebSocket({"type": "websocket", "path": "/v1/responses", "headers": [], "query_string": b""}, receive, send) + + +@pytest.mark.asyncio +async def test_http_exception_on_a_plain_http_request_keeps_the_default_json_body() -> None: + request: Final = _make_request(path="/v1/models") + + response: Final = await otlp_http_exception_handler(request, HTTPException(404, "not found")) + + assert response is not None + assert response.status_code == 404 + assert json.loads(bytes(response.body)) == {"detail": "not found"} + + +@pytest.mark.asyncio +async def test_http_exception_on_a_websocket_closed_before_accept_sends_nothing_more() -> None: + sent: Final[list[Message]] = [] + websocket: Final = _websocket(sent) + await websocket.close(code=1008) + + response: Final = await otlp_http_exception_handler(websocket, HTTPException(403, "No API key provided")) + + assert response is None + assert sent == [{"type": "websocket.close", "code": 1008, "reason": ""}] + + +@pytest.mark.asyncio +async def test_http_exception_on_a_connecting_websocket_denies_the_upgrade_with_its_status() -> None: + websocket: Final = _websocket([]) + + response: Final = await otlp_http_exception_handler(websocket, HTTPException(403, "No API key provided")) + + assert response is not None + assert response.status_code == 403 + assert json.loads(bytes(response.body)) == {"detail": "No API key provided"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("status_code", "close_code"), [(403, 1008), (429, 1008), (500, 1011), (503, 1011)]) +async def test_http_exception_on_an_accepted_websocket_closes_it_with_a_matching_code( + status_code: int, close_code: int +) -> None: + sent: Final[list[Message]] = [] + websocket: Final = _websocket(sent) + await websocket.accept() + + response: Final = await otlp_http_exception_handler(websocket, HTTPException(status_code, "late failure")) + + assert response is None + assert sent[-1] == {"type": "websocket.close", "code": close_code, "reason": ""} + + +def test_websocket_auth_rejection_reaches_the_client_as_a_denial_instead_of_a_server_error() -> None: + async def reject_like_user_api_key_auth_websocket(websocket: WebSocket) -> None: + await websocket.close(code=1008) + raise HTTPException(status_code=403, detail="No API key provided") + + async def responses(websocket: WebSocket, _: None = Depends(reject_like_user_api_key_auth_websocket)) -> None: + await websocket.accept() + + app: Final = FastAPI() + app.exception_handler(HTTPException)(otlp_http_exception_handler) + app.add_api_websocket_route("/v1/responses", responses) + + with pytest.raises(WebSocketDisconnect) as disconnect, TestClient(app).websocket_connect("/v1/responses"): + pass + + assert disconnect.value.code == 1008 + + +def test_websocket_rejected_before_any_close_reaches_the_client_as_an_http_denial_with_the_status() -> None: + async def reject_without_closing(websocket: WebSocket) -> None: + raise HTTPException(status_code=403, detail="No API key provided") + + async def responses(websocket: WebSocket, _: None = Depends(reject_without_closing)) -> None: + await websocket.accept() + + app: Final = FastAPI() + app.exception_handler(HTTPException)(otlp_http_exception_handler) + app.add_api_websocket_route("/v1/responses", responses) + + with pytest.raises(WebSocketDenialResponse) as denial, TestClient(app).websocket_connect("/v1/responses"): + pass + + assert denial.value.status_code == 403 + assert denial.value.json() == {"detail": "No API key provided"} diff --git a/tests/unit/proxy/public_endpoints/test_public_endpoints.py b/tests/unit/proxy/public_endpoints/test_public_endpoints.py index 30538c1167f..b720ba24c3e 100644 --- a/tests/unit/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/unit/proxy/public_endpoints/test_public_endpoints.py @@ -9,6 +9,7 @@ import pytest from fastapi import FastAPI from fastapi.testclient import TestClient +from pydantic import TypeAdapter from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.public_endpoints import router @@ -16,6 +17,7 @@ from litellm.router_strategy.complexity_router.fuse_presets import get_fuse_pres from litellm.types.proxy.management_endpoints.model_management_endpoints import ( ModelGroupInfoProxy, ) +from litellm.types.proxy.public_endpoints.public_endpoints import ProviderCreateInfo from litellm.types.utils import LlmProviders @@ -74,6 +76,32 @@ def test_get_provider_create_fields(): ), "Expected at least one provider to have detailed credential fields" +def test_get_litellm_model_cost_map_catalog_only_excludes_runtime_registered_entries( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm + from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap + + monkeypatch.setattr(litellm, "model_cost", dict(litellm.model_cost)) + runtime_key: Final = "lit-9263-deployment-alias" + litellm.register_model({runtime_key: {"litellm_provider": "openai", "mode": "chat"}}, persist_across_reloads=False) + app: Final = FastAPI() + app.include_router(router) + client: Final = TestClient(app) + + live_response: Final = client.get("/public/litellm_model_cost_map") + catalog_response: Final = client.get("/public/litellm_model_cost_map", params={"catalog_only": "true"}) + + assert live_response.status_code == 200 + assert live_response.json()[runtime_key]["litellm_provider"] == "openai" + assert catalog_response.status_code == 200 + catalog_payload: Final = catalog_response.json() + assert runtime_key not in catalog_payload + assert catalog_payload == json.loads( + json.dumps({key: dict(entry) for key, entry in GetModelCostMap.loaded_model_cost_map().items()}) + ) + + def test_get_litellm_model_cost_map_returns_cost_map(): app = FastAPI() app.include_router(router) @@ -402,6 +430,53 @@ def test_tencent_provider_fields(): assert fields_by_key["api_base"]["required"] is False +def _decisions_provider_entry(provider: str) -> ProviderCreateInfo: + app_instance: Final = FastAPI() + app_instance.include_router(router) + test_client: Final = TestClient(app_instance) + + response: Final = test_client.get("/public/providers/fields") + assert response.status_code == 200 + providers: Final = TypeAdapter(list[ProviderCreateInfo]).validate_python(response.json()) + entry: Final = next((p for p in providers if p.provider == provider), None) + assert entry is not None, f"{provider} provider entry not found" + return entry + + +def test_typesafe_provider_fields(): + typesafe: Final = _decisions_provider_entry("TypeSafe") + + assert typesafe.provider_display_name == "TypeSafe" + assert typesafe.litellm_provider == LlmProviders.TYPESAFE.value + assert typesafe.default_model_placeholder is not None + assert typesafe.default_model_placeholder.startswith("typesafe/") + + fields_by_key: Final = {f.key: f for f in typesafe.credential_fields} + + assert fields_by_key["api_key"].required is True + assert fields_by_key["api_key"].field_type == "password" + + assert fields_by_key["api_base"].required is False + assert fields_by_key["api_base"].field_type == "text" + + +def test_strands_decider_provider_fields(): + strands: Final = _decisions_provider_entry("StrandsDecider") + + assert strands.provider_display_name == "Strands Decider" + assert strands.litellm_provider == LlmProviders.STRANDS_DECIDER.value + assert strands.default_model_placeholder is not None + assert strands.default_model_placeholder.startswith("strands_decider/") + + fields_by_key: Final = {f.key: f for f in strands.credential_fields} + + assert fields_by_key["api_base"].required is True + assert fields_by_key["api_base"].field_type == "text" + + assert fields_by_key["api_key"].required is False + assert fields_by_key["api_key"].field_type == "password" + + ADD_MODEL_UNLISTED_PROVIDERS: Final = frozenset( { "a2a", @@ -436,12 +511,10 @@ ADD_MODEL_UNLISTED_PROVIDERS: Final = frozenset( "sagemaker_nova", "scaleway", "stability", - "strands_decider", "synthetic", "tensormesh", "text-completion-inception", "transcribe", - "typesafe", "valkey", "xiaomi_mimo", "zai", diff --git a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py index 0df6eecd4b9..9aa9110f7a4 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py @@ -8,18 +8,21 @@ from typing import Any, Final, cast from unittest.mock import AsyncMock, MagicMock, patch import pytest -from typing_extensions import ReadOnly, TypedDict +from pydantic import TypeAdapter +from typing_extensions import NotRequired, ReadOnly, TypedDict import litellm import litellm.constants as litellm_constants import litellm.proxy.spend_tracking.spend_tracking_utils as spend_tracking_utils from litellm.constants import LITELLM_TRUNCATED_PAYLOAD_FIELD, LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE, LITTELM_CLI_SERVICE_ACCOUNT_NAME, LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME, MAX_SPEND_LOG_MODEL_NAME_LENGTH, REDACTED_BY_LITELM_STRING, SESSION_ID_OMITTED_METADATA_KEY, UNKNOWN_MODEL_SPEND_LOG_MODEL from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup +from litellm.llms.base_llm.ocr.transformation import OCRResponse, OCRUsageInfo from litellm.litellm_core_utils.safe_json_dumps import safe_dumps -from litellm.proxy._types import SpendLogsPayload, UserAPIKeyAuth +from litellm.proxy._types import SpendLogsMetadataFields, SpendLogsPayload, UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.spend_tracking.spend_tracking_utils import ( + _extract_usage_for_ocr_call, _get_messages_for_spend_logs_payload, _get_proxy_server_request_for_spend_logs_payload, _get_request_duration_ms, @@ -34,9 +37,11 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import ( _sanitize_guardrail_information_for_spend_logs, _sanitize_request_body_for_spend_logs_payload, _scrub_raw_model_from_error_information, + configured_spend_logs_metadata_fields, get_logging_payload, get_spend_logs_id, should_store_prompts_and_responses_in_spend_logs, + spend_log_row_with_retained_metadata, ) from litellm.proxy.utils import hash_token from litellm.types.router import GenericLiteLLMParams @@ -76,10 +81,16 @@ def test_classifier_audit_spend_storage_obeys_privacy_and_truncation(monkeypatch "classifier_input": {"system": "rubric" * 1000, "messages": [{"role": "user", "content": "ask"}]}, "originating_request_masked": {"input": "source-only", "api_key": "REDACTED"}, } - stored: Final = json.loads(_get_proxy_server_request_for_spend_logs_payload( - metadata={}, litellm_params={"proxy_server_request": {"body": {"model": "classifier"}}}, - kwargs={"standard_logging_object": audit, "standard_callback_dynamic_params": {"turn_off_message_logging": redact}}, - )) + stored: Final = json.loads( + _get_proxy_server_request_for_spend_logs_payload( + metadata={}, + litellm_params={"proxy_server_request": {"body": {"model": "classifier"}}}, + kwargs={ + "standard_logging_object": audit, + "standard_callback_dynamic_params": {"turn_off_message_logging": redact}, + }, + ) + ) if not store_prompts or redact: assert "classifier_input" not in stored assert "originating_request_masked" not in stored @@ -205,9 +216,7 @@ def test_batch_lifecycle_rows_derive_the_same_session_from_the_batch_id(): from litellm.proxy.spend_tracking.spend_tracking_utils import _get_batch_trace_session_id create_session: Final = _get_batch_trace_session_id(call_type="acreate_batch", request_id="batch-uid-1") - cost_session: Final = _get_batch_trace_session_id( - call_type="aretrieve_batch", request_id="batch-uid-1_batch_cost" - ) + cost_session: Final = _get_batch_trace_session_id(call_type="aretrieve_batch", request_id="batch-uid-1_batch_cost") assert create_session == cost_session == "batch-uid-1" @@ -2908,7 +2917,11 @@ def test_proxy_server_request_payload_redacts_provider_credentials(mock_should_s "extra_headers": {"Authorization": "Bearer canary-extra-header"}, "tools": [ {"type": "function", "function": {"name": "f", "parameters": tool_parameters}}, - {"type": "mcp", "server_url": "https://mcp.example.com", "headers": {"Authorization": "canary-mcp"}}, + { + "type": "mcp", + "server_url": "https://mcp.example.com", + "headers": {"Authorization": "canary-mcp"}, + }, ], "fallbacks": [{"model": "azure-b", **credentials}], "metadata": metadata, @@ -5462,7 +5475,7 @@ ANTHROPIC_MESSAGES_SSE_CHUNKS: Final = ( 'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n', 'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn"},' '"usage":{"output_tokens":4}}\n\n', - "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n", + 'event: message_stop\ndata: {"type":"message_stop"}\n\n', ) @@ -5500,9 +5513,7 @@ def test_spend_log_request_id_is_the_message_id_a_non_streaming_messages_caller_ """ logging_obj = _anthropic_messages_logging_obj(stream=False) - logged_response = logging_obj._handle_anthropic_messages_response_logging( - result=ANTHROPIC_MESSAGES_RESPONSE - ) + logged_response = logging_obj._handle_anthropic_messages_response_logging(result=ANTHROPIC_MESSAGES_RESPONSE) assert logged_response.id == "msg_01Lit6806NonStreaming" assert ( @@ -5578,9 +5589,7 @@ def test_spend_log_request_id_still_falls_back_to_litellm_call_id_without_a_prov end_time=datetime.datetime.now(timezone.utc), logging_obj=logging_obj, ) - assert logging_obj.model_call_details["complete_streaming_response"].id == ( - "6806cafe-0000-4000-8000-000000000001" - ) + assert logging_obj.model_call_details["complete_streaming_response"].id == ("6806cafe-0000-4000-8000-000000000001") def test_spend_log_request_id_for_chat_completions_is_untouched(): @@ -5662,6 +5671,7 @@ def test_failed_agent_request_keeps_registered_display_name(): assert payload["status"] == "failure" assert payload["model_id"] == "registered-agent" + _CLI_SESSION_ALIAS: Final = "cli-session-alice" _CLI_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa" @@ -5793,15 +5803,337 @@ def test_baseline_estimate_metadata_comes_from_the_logging_stamp() -> None: def test_untrusted_agent_label_cannot_replace_verified_billing_identity(billing_agent: str | None) -> None: kwargs = { "model": "gpt-4", - "litellm_params": {"metadata": { - "user_api_key": "test-key", - "agent_id": "header-selected-agent", - "billing_agent_id": billing_agent, - }}, + "litellm_params": { + "metadata": { + "user_api_key": "test-key", + "agent_id": "header-selected-agent", + "billing_agent_id": billing_agent, + } + }, } payload = get_logging_payload( - kwargs=kwargs, response_obj={"id": "request"}, - start_time=datetime.datetime.now(timezone.utc), end_time=datetime.datetime.now(timezone.utc), + kwargs=kwargs, + response_obj={"id": "request"}, + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), ) assert payload["agent_id"] == "header-selected-agent" assert payload["billing_agent_id"] == billing_agent + + +def _spend_log_row(metadata: Mapping[str, object]) -> Mapping[str, object]: + return MappingProxyType({"request_id": "req-1", "spend": 0.5, "metadata": json.dumps(dict(metadata))}) + + +_STORED_METADATA: Final = MappingProxyType( + { + "status": "success", + "cold_storage_object_key": "logs/req-1.json", + "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"max_tokens": 10}}, + "usage_object": {"prompt_tokens": 3}, + "user_api_key_alias": "alias", + } +) + + +def test_spend_log_row_keeps_every_metadata_field_when_unconfigured() -> None: + row: Final = _spend_log_row(_STORED_METADATA) + + assert spend_log_row_with_retained_metadata(row, None) is row + + +def test_spend_log_row_drops_excluded_metadata_fields_only() -> None: + row: Final = _spend_log_row(_STORED_METADATA) + + stored: Final = spend_log_row_with_retained_metadata( + row, SpendLogsMetadataFields(exclude=("model_map_information", "user_api_key_alias")) + ) + + assert json.loads(cast(str, stored["metadata"])) == { + "status": "success", + "cold_storage_object_key": "logs/req-1.json", + "usage_object": {"prompt_tokens": 3}, + } + assert {name: value for name, value in stored.items() if name != "metadata"} == { + "request_id": "req-1", + "spend": 0.5, + } + + +def test_spend_log_row_include_keeps_listed_and_always_kept_fields() -> None: + stored: Final = spend_log_row_with_retained_metadata( + _spend_log_row(_STORED_METADATA), SpendLogsMetadataFields(include=("usage_object",)) + ) + + assert json.loads(cast(str, stored["metadata"])) == { + "status": "success", + "cold_storage_object_key": "logs/req-1.json", + "usage_object": {"prompt_tokens": 3}, + } + + +@pytest.mark.parametrize( + "configured", + [ + {"include": ["usage_object"], "exclude": ["model_map_information"]}, + {}, + {"exclude": ["model_map_informaton"]}, + {"include": ["usage_object", "not_a_field"]}, + {"exclude": ["status"]}, + {"exclude": ["cold_storage_object_key"]}, + {"exclude": ["model_map_information"], "drop": ["usage_object"]}, + ], +) +def test_spend_logs_metadata_fields_rejects_ambiguous_or_lossy_config(configured: dict[str, list[str]]) -> None: + from pydantic import ValidationError + + from litellm.proxy._types import ConfigGeneralSettings + + with pytest.raises(ValidationError): + ConfigGeneralSettings.model_validate({"spend_logs_metadata_fields": configured}) + + +def test_configured_spend_logs_metadata_fields_ignores_invalid_runtime_value() -> None: + with patch( + "litellm.proxy.proxy_server.general_settings", + {"spend_logs_metadata_fields": {"include": ["usage_object"], "exclude": ["status"]}}, + ): + assert configured_spend_logs_metadata_fields() is None + with patch( + "litellm.proxy.proxy_server.general_settings", + {"spend_logs_metadata_fields": {"exclude": ["model_map_information"]}}, + ): + assert configured_spend_logs_metadata_fields() == SpendLogsMetadataFields(exclude=("model_map_information",)) + + +class _OcrUsageInfoDict(TypedDict, total=False): + pages_processed: ReadOnly[int] + doc_size_bytes: ReadOnly[int] + + +class _OcrResponseDict(TypedDict): + id: ReadOnly[str] + object: ReadOnly[str] + model: ReadOnly[str] + usage_info: ReadOnly[NotRequired[_OcrUsageInfoDict]] + + +class _TokenUsageDict(TypedDict): + prompt_tokens: ReadOnly[int] + completion_tokens: ReadOnly[int] + total_tokens: ReadOnly[int] + + +class _CompletionResponseDict(TypedDict): + id: ReadOnly[str] + object: ReadOnly[str] + model: ReadOnly[str] + usage: ReadOnly[_TokenUsageDict] + + +class _LoggingMetadata(TypedDict, total=False): + user_api_key_user_id: ReadOnly[str] + user_api_key_team_id: ReadOnly[str] + + +class _LoggingLitellmParams(TypedDict, total=False): + metadata: ReadOnly[_LoggingMetadata] + + +class _LoggingKwargs(TypedDict): + model: ReadOnly[str] + call_type: ReadOnly[str] + litellm_params: ReadOnly[_LoggingLitellmParams] + response_cost: ReadOnly[float] + + +class _AdditionalUsageValues(TypedDict, total=False): + pages_processed: ReadOnly[int | None] + doc_size_bytes: ReadOnly[int | None] + + +class _SpendLogMetadata(TypedDict): + additional_usage_values: ReadOnly[_AdditionalUsageValues] + + +_SPEND_LOG_METADATA: Final = TypeAdapter(_SpendLogMetadata) +_OCR_LOGGED_AT: Final = datetime.datetime(2026, 1, 1, tzinfo=timezone.utc) +_OCR_RESPONSE_COST: Final = 0.05 + + +def _ocr_logging_kwargs(call_type: str = "ocr", metadata: _LoggingMetadata | None = None) -> _LoggingKwargs: + return _LoggingKwargs( + model="test-ocr-model", + call_type=call_type, + litellm_params=_LoggingLitellmParams() if metadata is None else _LoggingLitellmParams(metadata=metadata), + response_cost=_OCR_RESPONSE_COST, + ) + + +def _ocr_payload( + kwargs: _LoggingKwargs, response_obj: _OcrResponseDict | _CompletionResponseDict | OCRResponse +) -> SpendLogsPayload: + return get_logging_payload( + kwargs=dict(kwargs), + response_obj=response_obj, + start_time=_OCR_LOGGED_AT, + end_time=_OCR_LOGGED_AT, + ) + + +def _additional_usage_values(payload: SpendLogsPayload) -> _AdditionalUsageValues: + return _SPEND_LOG_METADATA.validate_json(payload["metadata"])["additional_usage_values"] + + +class TestExtractUsageForOCRCall: + def test_extract_usage_from_dict(self) -> None: + response_obj_dict: Final = {"usage_info": _OcrUsageInfoDict(pages_processed=5)} + + usage: Final = _extract_usage_for_ocr_call(response_obj_dict, response_obj_dict) + + assert usage == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0, "pages_processed": 5} + + def test_extract_usage_from_pydantic_model(self) -> None: + response_obj: Final = OCRResponse( + pages=[], + model="test-ocr-model", + usage_info=OCRUsageInfo(pages_processed=10, doc_size_bytes=1024), + ) + + usage: Final = _extract_usage_for_ocr_call(response_obj, response_obj.model_dump()) + + assert usage["prompt_tokens"] == 0 + assert usage["completion_tokens"] == 0 + assert usage["total_tokens"] == 0 + assert usage["pages_processed"] == 10 + assert usage["doc_size_bytes"] == 1024 + + def test_extract_usage_with_object_attributes(self) -> None: + class _SimpleUsageInfo: + def __init__(self, pages_processed: int) -> None: + self.pages_processed = pages_processed + + class _SimpleOCRResponse: + def __init__(self) -> None: + self.usage_info = _SimpleUsageInfo(pages_processed=3) + + usage: Final = _extract_usage_for_ocr_call(_SimpleOCRResponse(), {}) + + assert usage == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0, "pages_processed": 3} + + def test_extract_usage_missing_usage_info(self) -> None: + assert _extract_usage_for_ocr_call({}, {}) == {} + + def test_extract_usage_empty_usage_info(self) -> None: + usage: Final = _extract_usage_for_ocr_call({"usage_info": {}}, {"usage_info": {}}) + + assert usage == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0, "pages_processed": 0} + + +class TestGetLoggingPayloadOCR: + def test_ocr_call_with_dict_response(self) -> None: + payload: Final = _ocr_payload( + _ocr_logging_kwargs(), + { + "id": "ocr-test-123", + "object": "ocr", + "model": "test-ocr-model", + "usage_info": {"pages_processed": 7, "doc_size_bytes": 2048}, + }, + ) + + assert payload["call_type"] == "ocr" + assert payload["request_id"] == "ocr-test-123" + assert payload["prompt_tokens"] == 0 + assert payload["completion_tokens"] == 0 + assert payload["total_tokens"] == 0 + assert payload["spend"] == _OCR_RESPONSE_COST + assert _additional_usage_values(payload)["pages_processed"] == 7 + assert _additional_usage_values(payload)["doc_size_bytes"] == 2048 + + def test_aocr_call_with_pydantic_response(self) -> None: + payload: Final = _ocr_payload( + _ocr_logging_kwargs(call_type="aocr"), + OCRResponse(pages=[], model="test-ocr-model", usage_info=OCRUsageInfo(pages_processed=12)), + ) + + assert payload["call_type"] == "aocr" + assert payload["prompt_tokens"] == 0 + assert payload["completion_tokens"] == 0 + assert payload["total_tokens"] == 0 + assert payload["spend"] == _OCR_RESPONSE_COST + assert _additional_usage_values(payload)["pages_processed"] == 12 + + def test_ocr_call_missing_usage_info(self) -> None: + payload: Final = _ocr_payload( + _ocr_logging_kwargs(), + {"id": "ocr-test-789", "object": "ocr", "model": "test-ocr-model"}, + ) + + assert payload["call_type"] == "ocr" + assert payload["prompt_tokens"] == 0 + assert payload["completion_tokens"] == 0 + assert payload["total_tokens"] == 0 + assert payload["spend"] == _OCR_RESPONSE_COST + assert "pages_processed" not in _additional_usage_values(payload) + + def test_ocr_call_with_zero_pages(self) -> None: + payload: Final = _ocr_payload( + _ocr_logging_kwargs(), + { + "id": "ocr-test-000", + "object": "ocr", + "model": "test-ocr-model", + "usage_info": {"pages_processed": 0}, + }, + ) + + assert payload["call_type"] == "ocr" + assert payload["prompt_tokens"] == 0 + assert payload["completion_tokens"] == 0 + assert payload["total_tokens"] == 0 + assert payload["spend"] == _OCR_RESPONSE_COST + assert _additional_usage_values(payload)["pages_processed"] == 0 + + def test_non_ocr_call_uses_token_based_usage(self) -> None: + payload: Final = _ocr_payload( + _LoggingKwargs( + model="gpt-5.5", call_type="completion", litellm_params=_LoggingLitellmParams(), response_cost=0.02 + ), + { + "id": "completion-test-123", + "object": "chat.completion", + "model": "gpt-5.5", + "usage": {"prompt_tokens": 50, "completion_tokens": 100, "total_tokens": 150}, + }, + ) + + assert payload["call_type"] == "completion" + assert payload["prompt_tokens"] == 50 + assert payload["completion_tokens"] == 100 + assert payload["total_tokens"] == 150 + assert payload["spend"] == 0.02 + assert "pages_processed" not in _additional_usage_values(payload) + + def test_ocr_with_metadata(self) -> None: + payload: Final = _ocr_payload( + _ocr_logging_kwargs( + metadata=_LoggingMetadata(user_api_key_user_id="test-user", user_api_key_team_id="test-team") + ), + { + "id": "ocr-metadata-test", + "object": "ocr", + "model": "test-ocr-model", + "usage_info": {"pages_processed": 5, "doc_size_bytes": 1024}, + }, + ) + + assert payload["call_type"] == "ocr" + assert payload["user"] == "test-user" + assert payload["team_id"] == "test-team" + assert payload["prompt_tokens"] == 0 + assert payload["completion_tokens"] == 0 + assert payload["total_tokens"] == 0 + assert payload["spend"] == _OCR_RESPONSE_COST + assert _additional_usage_values(payload)["pages_processed"] == 5 + assert _additional_usage_values(payload)["doc_size_bytes"] == 1024 diff --git a/tests/unit/proxy/test__lazy_features.py b/tests/unit/proxy/test__lazy_features.py index c3c5068b890..c8e34ad750c 100644 --- a/tests/unit/proxy/test__lazy_features.py +++ b/tests/unit/proxy/test__lazy_features.py @@ -21,6 +21,20 @@ FLAG: Final = "LITELLM_DISABLE_LAZY_ROUTES" WARMUP_PATH: Final = "/lazy/warm/{name}" +def test_cimd_metadata_is_available_before_any_oauth_request(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv(FLAG, "false") + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + app: Final = FastAPI() + attach_lazy_features(app) + + with TestClient(app) as client: + response: Final = client.get("/oauth/client-metadata.json") + + assert response.status_code == 200 + assert response.json()["client_id"] == "https://gateway.example.com/oauth/client-metadata.json" + assert response.json()["redirect_uris"] == ["https://gateway.example.com/callback"] + + class _Operation(BaseModel): tags: tuple[str, ...] diff --git a/tests/unit/proxy/test_fallback_management_endpoints.py b/tests/unit/proxy/test_fallback_management_endpoints.py index 054dafbf5a7..3ee1b64580d 100644 --- a/tests/unit/proxy/test_fallback_management_endpoints.py +++ b/tests/unit/proxy/test_fallback_management_endpoints.py @@ -8,17 +8,27 @@ Tests: 4. Validation tests (invalid models, duplicate fallbacks, etc.) """ +import copy +import json +from collections.abc import AsyncIterator +from types import SimpleNamespace +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException +from litellm import Router +from litellm.proxy.utils import evict_config_param, get_config_param from litellm.proxy.management_endpoints.fallback_management_endpoints import ( FallbackCreateRequest, + FallbackDeleteResponse, + FallbackResponse, create_fallback, delete_fallback, get_fallback, ) +from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict class TestFallbackCreateRequest: @@ -104,6 +114,268 @@ class TestFallbackCreateRequest: assert request.fallback_type == "content_policy" +TEAM_ID: Final = "team-a" +PRIMARY_INTERNAL_NAME: Final = f"model_name_{TEAM_ID}_1a6437cb-4cab-432c-8099-1d7411731a8b" +FALLBACK_INTERNAL_NAME: Final = f"model_name_{TEAM_ID}_83151607-5556-4bbf-ac65-c474dbc64eba" +PRIMARY_DEPLOYMENT_ID: Final = "team-primary-id" +FALLBACK_DEPLOYMENT_ID: Final = "team-fallback-id" + +FallbackRules = list[dict[str, list[str]] | dict[str, object] | str] +RouterSettings = dict[str, FallbackRules | None] + + +def _litellm_params(mock_response: str | None) -> LiteLLMParamsTypedDict: + if mock_response is None: + return {"model": "openai/gpt-5.4-mini", "api_key": "fake"} + return {"model": "openai/gpt-5.4-mini", "api_key": "fake", "mock_response": mock_response} + + +def _team_scoped_deployment( + internal_name: str, public_name: str, deployment_id: str, mock_response: str | None = None +) -> DeploymentTypedDict: + return { + "model_name": internal_name, + "litellm_params": _litellm_params(mock_response), + "model_info": {"id": deployment_id, "team_id": TEAM_ID, "team_public_model_name": public_name}, + } + + +def _team_router(primary_mock_response: str | None = None, fallback_mock_response: str | None = None) -> Router: + return Router( + model_list=[ + {"model_name": "gpt-5.4-mini", "litellm_params": _litellm_params(None)}, + _team_scoped_deployment( + PRIMARY_INTERNAL_NAME, "team-primary", PRIMARY_DEPLOYMENT_ID, primary_mock_response + ), + _team_scoped_deployment( + FALLBACK_INTERNAL_NAME, "team-fallback", FALLBACK_DEPLOYMENT_ID, fallback_mock_response + ), + ], + num_retries=0, + ) + + +class TestCreateFallbackForTeamScopedModels: + """POST /fallback takes the public name a caller invokes a team-scoped model by, not only the stored internal one""" + + @pytest.fixture + def router(self) -> Router: + return _team_router() + + @pytest.fixture + def prisma_client(self) -> MagicMock: + client: Final = MagicMock() + client.db.litellm_config.upsert = AsyncMock() + return client + + @pytest.fixture + def proxy_config(self) -> MagicMock: + config: Final = MagicMock() + config.get_config = AsyncMock(return_value={"router_settings": {}}) + return config + + async def _create( + self, request: FallbackCreateRequest, router: Router, prisma_client: MagicMock, proxy_config: MagicMock + ) -> FallbackResponse: + with ( + patch("litellm.proxy.proxy_server.llm_router", router), + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + ): + return await create_fallback(request, MagicMock()) + + async def test_public_names_create_a_rule_keyed_on_the_public_name( + self, router: Router, prisma_client: MagicMock, proxy_config: MagicMock + ) -> None: + request: Final = FallbackCreateRequest(model="team-primary", fallback_models=["team-fallback"]) + + response: Final = await self._create(request, router, prisma_client, proxy_config) + + assert response.model == "team-primary" + assert response.fallback_models == ["team-fallback"] + assert router.fallbacks == [{"team-primary": ["team-fallback"]}] + persisted: Final = json.loads( + prisma_client.db.litellm_config.upsert.call_args.kwargs["data"]["create"]["param_value"] + ) + assert persisted["fallbacks"] == [{"team-primary": ["team-fallback"]}] + + async def test_internal_names_keep_working( + self, router: Router, prisma_client: MagicMock, proxy_config: MagicMock + ) -> None: + request: Final = FallbackCreateRequest(model=PRIMARY_INTERNAL_NAME, fallback_models=[FALLBACK_INTERNAL_NAME]) + + response: Final = await self._create(request, router, prisma_client, proxy_config) + + assert router.fallbacks == [{PRIMARY_INTERNAL_NAME: [FALLBACK_INTERNAL_NAME]}] + assert response.model == PRIMARY_INTERNAL_NAME + + async def test_public_fallback_target_behind_a_gateway_primary( + self, router: Router, prisma_client: MagicMock, proxy_config: MagicMock + ) -> None: + request: Final = FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"]) + + await self._create(request, router, prisma_client, proxy_config) + + assert router.fallbacks == [{"gpt-5.4-mini": ["team-fallback"]}] + + async def test_unknown_name_is_rejected_and_the_error_names_the_public_names( + self, router: Router, prisma_client: MagicMock, proxy_config: MagicMock + ) -> None: + request: Final = FallbackCreateRequest(model="team-missing", fallback_models=["team-fallback"]) + + with pytest.raises(HTTPException) as exc_info: + await self._create(request, router, prisma_client, proxy_config) + + assert exc_info.value.status_code == 404 + assert {"team-primary", "team-fallback", "gpt-5.4-mini"} <= set(exc_info.value.detail["available_models"]) + + async def test_a_team_request_fails_over_to_the_public_name_fallback( + self, prisma_client: MagicMock, proxy_config: MagicMock + ) -> None: + router: Final = _team_router(primary_mock_response="litellm.RateLimitError", fallback_mock_response="pong") + request: Final = FallbackCreateRequest(model="team-primary", fallback_models=["team-fallback"]) + await self._create(request, router, prisma_client, proxy_config) + + response: Final = await router.acompletion( + model="team-primary", + messages=[{"role": "user", "content": "ping"}], + metadata={"user_api_key_team_id": TEAM_ID}, + ) + + assert response.choices[0].message.content == "pong" + assert response._hidden_params["model_id"] == FALLBACK_DEPLOYMENT_ID + + +def _config_row(router_settings: RouterSettings) -> SimpleNamespace: + return SimpleNamespace(param_name="router_settings", param_value=router_settings) + + +class _StoredRouterSettings: + """The LiteLLM_Config router_settings row, with the proxy objects that read it through the config cache""" + + def __init__(self, router_settings: RouterSettings | None) -> None: + self.row: SimpleNamespace | None = None if router_settings is None else _config_row(router_settings) + self.prisma_client: Final = MagicMock() + self.prisma_client.get_generic_data = AsyncMock(side_effect=lambda **_: self.row) + self.prisma_client.db.litellm_config.upsert = AsyncMock(side_effect=self._upsert) + self.proxy_config: Final = MagicMock() + self.proxy_config.get_config = AsyncMock(side_effect=self._get_config) + + async def _upsert(self, where: dict[str, str], data: dict[str, dict[str, str]]) -> None: + self.row = _config_row(json.loads(data["update"]["param_value"])) + + async def _get_config(self) -> dict[str, RouterSettings]: + row: Final = await get_config_param(self.prisma_client, "router_settings") + return {"router_settings": copy.deepcopy(row.param_value) if row is not None else {}} + + def written_by_another_instance(self, router_settings: RouterSettings) -> None: + self.row = _config_row(router_settings) + + def stored_fallbacks(self) -> FallbackRules: + assert self.row is not None + return self.row.param_value["fallbacks"] + + async def cached_fallbacks(self) -> FallbackRules: + row: Final = await get_config_param(self.prisma_client, "router_settings") + return row.param_value["fallbacks"] + + +TEAM_RULE: Final = {"team-primary": ["team-fallback"]} +GATEWAY_RULE: Final = {"gpt-5.4-mini": ["team-fallback"]} +MALFORMED_GATEWAY_RULE: Final = {"gpt-5.4-mini": "team-fallback"} +NON_STANDARD_RULES: Final = [ + "claude-3-haiku", + {"model": "gpt-5.4-mini", "messages": [{"role": "user", "content": "retry"}]}, +] + + +class TestFallbackWritesSeeTheLatestStoredRules: + """A write reads the rules the database holds now and leaves no stale copy in the config cache behind""" + + @pytest.fixture(autouse=True) + async def clean_config_cache(self) -> AsyncIterator[None]: + await evict_config_param("router_settings") + yield + await evict_config_param("router_settings") + + async def _create(self, request: FallbackCreateRequest, stored: _StoredRouterSettings) -> FallbackResponse: + with ( + patch("litellm.proxy.proxy_server.llm_router", _team_router()), + patch("litellm.proxy.proxy_server.prisma_client", stored.prisma_client), + patch("litellm.proxy.proxy_server.proxy_config", stored.proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + ): + return await create_fallback(request, MagicMock()) + + async def _delete(self, model: str, stored: _StoredRouterSettings) -> FallbackDeleteResponse: + with ( + patch("litellm.proxy.proxy_server.llm_router", _team_router()), + patch("litellm.proxy.proxy_server.prisma_client", stored.prisma_client), + patch("litellm.proxy.proxy_server.proxy_config", stored.proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + ): + return await delete_fallback(model, "general", MagicMock()) + + async def test_a_second_create_keeps_the_first_rule(self) -> None: + stored: Final = _StoredRouterSettings(None) + + await self._create(FallbackCreateRequest(model="team-primary", fallback_models=["team-fallback"]), stored) + await self._create(FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"]), stored) + + assert stored.stored_fallbacks() == [TEAM_RULE, GATEWAY_RULE] + assert await stored.cached_fallbacks() == [TEAM_RULE, GATEWAY_RULE] + + async def test_a_create_keeps_a_rule_another_instance_stored_since_this_one_last_read(self) -> None: + stored: Final = _StoredRouterSettings({}) + await get_config_param(stored.prisma_client, "router_settings") + stored.written_by_another_instance({"fallbacks": [TEAM_RULE]}) + + await self._create(FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"]), stored) + + assert stored.stored_fallbacks() == [TEAM_RULE, GATEWAY_RULE] + + async def test_a_delete_keeps_a_rule_another_instance_stored_since_this_one_last_read(self) -> None: + stored: Final = _StoredRouterSettings({"fallbacks": [TEAM_RULE]}) + await get_config_param(stored.prisma_client, "router_settings") + stored.written_by_another_instance({"fallbacks": [TEAM_RULE, GATEWAY_RULE]}) + + await self._delete("team-primary", stored) + + assert stored.stored_fallbacks() == [GATEWAY_RULE] + assert await stored.cached_fallbacks() == [GATEWAY_RULE] + + async def test_writes_keep_the_non_standard_rules_the_router_accepts(self) -> None: + stored: Final = _StoredRouterSettings({"fallbacks": [*NON_STANDARD_RULES, TEAM_RULE]}) + + await self._create(FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"]), stored) + assert stored.stored_fallbacks() == [*NON_STANDARD_RULES, TEAM_RULE, GATEWAY_RULE] + + await self._delete("team-primary", stored) + assert stored.stored_fallbacks() == [*NON_STANDARD_RULES, GATEWAY_RULE] + + async def test_a_create_replaces_a_same_key_rule_whatever_its_targets_shape(self) -> None: + stored: Final = _StoredRouterSettings({"fallbacks": [MALFORMED_GATEWAY_RULE, TEAM_RULE]}) + + await self._create(FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"]), stored) + + assert stored.stored_fallbacks() == [GATEWAY_RULE, TEAM_RULE] + + async def test_a_delete_removes_a_same_key_rule_whatever_its_targets_shape(self) -> None: + stored: Final = _StoredRouterSettings({"fallbacks": [MALFORMED_GATEWAY_RULE, TEAM_RULE]}) + + await self._delete("gpt-5.4-mini", stored) + + assert stored.stored_fallbacks() == [TEAM_RULE] + + async def test_a_create_treats_a_null_rule_list_as_empty(self) -> None: + stored: Final = _StoredRouterSettings({"fallbacks": None}) + + await self._create(FallbackCreateRequest(model="gpt-5.4-mini", fallback_models=["team-fallback"]), stored) + + assert stored.stored_fallbacks() == [GATEWAY_RULE] + + @pytest.mark.asyncio class TestCreateFallback: """Test the create_fallback endpoint""" @@ -113,6 +385,7 @@ class TestCreateFallback: """Create a mock router""" router = MagicMock() router.model_names = {"gpt-3.5-turbo", "gpt-4", "claude-3-haiku"} + router.team_public_model_names = frozenset() router.fallbacks = [] router.context_window_fallbacks = [] router.content_policy_fallbacks = [] diff --git a/tests/unit/proxy/test_health_check_max_tokens.py b/tests/unit/proxy/test_health_check_max_tokens.py index 091de1e24b3..cfdecdba69d 100644 --- a/tests/unit/proxy/test_health_check_max_tokens.py +++ b/tests/unit/proxy/test_health_check_max_tokens.py @@ -1,7 +1,10 @@ +import asyncio import json import logging +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest import respx @@ -1131,3 +1134,166 @@ def test_transitive_probe_expansion_terminates_on_a_router_cycle(): probes = hc_module._dependency_deployments_to_probe(a_only, router.model_list, router) assert {d["model_info"]["id"] for d in probes} == {"b-1"} + + +@pytest.mark.asyncio +async def test_background_audio_speech_health_check_uses_model_info_voice( + httpx_transport: None, respx_mock: respx.MockRouter +) -> None: + upstream: Final = respx_mock.post("https://speech.example/v1/audio/speech").respond( + content=b"audio", + headers={"content-type": "audio/mpeg"}, + ) + + healthy, unhealthy, _ = await hc_module.perform_health_check( + [ + { + "litellm_params": { + "model": "openai/tts-1", + "api_key": "fake-key", + "api_base": "https://speech.example/v1", + }, + "model_info": {"id": "speech", "mode": "audio_speech", "health_check_voice": "nova"}, + } + ], + max_concurrency=1, + ) + + assert len(healthy) == 1 + assert unhealthy == [] + assert upstream.called + assert json.loads(upstream.calls.last.request.content)["voice"] == "nova" + + +@pytest.mark.asyncio +async def test_background_health_check_observes_the_concurrency_limit_and_queue( + httpx_transport: None, respx_mock: respx.MockRouter +) -> None: + request_started: Final = asyncio.Queue[None]() + release: Final = asyncio.Event() + + async def complete_request(_: httpx.Request) -> httpx.Response: + request_started.put_nowait(None) + await release.wait() + return httpx.Response( + 200, + content=b"audio", + headers={"content-type": "audio/mpeg"}, + ) + + upstream: Final = respx_mock.post("https://health.example/v1/audio/speech").mock(side_effect=complete_request) + model_list: Final = [ + { + "litellm_params": { + "model": "openai/tts-1", + "api_key": "fake-key", + "api_base": "https://health.example/v1", + }, + "model_info": {"id": f"audio-{index}", "mode": "audio_speech"}, + } + for index in range(10) + ] + tasks_before: Final = len(asyncio.all_tasks()) + perform_task: Final = asyncio.create_task(hc_module.perform_health_check(model_list, max_concurrency=2)) + + try: + await asyncio.wait_for(request_started.get(), timeout=1) + await asyncio.wait_for(request_started.get(), timeout=1) + for _ in range(20): + await asyncio.sleep(0) + extra_requests_started: Final = request_started.qsize() + tasks_while_blocked: Final = len(asyncio.all_tasks()) - tasks_before + finally: + release.set() + healthy, unhealthy, _ = await perform_task + + assert extra_requests_started == 0 + assert tasks_while_blocked <= 5 + assert upstream.call_count == 10 + assert len(healthy) == 10 + assert unhealthy == [] + + +@pytest.mark.asyncio +async def test_background_health_check_timeout_marks_a_blocked_provider_unhealthy( + httpx_transport: None, respx_mock: respx.MockRouter +) -> None: + never_release: Final = asyncio.Event() + request_started: Final = asyncio.Event() + + async def blocked_response(_: httpx.Request) -> httpx.Response: + request_started.set() + await never_release.wait() + return httpx.Response(200, content=b"audio", headers={"content-type": "audio/mpeg"}) + + respx_mock.post("https://health.example/v1/audio/speech").mock(side_effect=blocked_response) + model_list: Final = [ + { + "litellm_params": { + "model": "openai/tts-1", + "api_key": "fake-key", + "api_base": "https://health.example/v1", + }, + "model_info": {"id": "blocked", "mode": "audio_speech", "health_check_timeout": 2}, + } + ] + + healthy, unhealthy, _ = await asyncio.wait_for( + hc_module.perform_health_check(model_list), + timeout=4, + ) + + assert request_started.is_set() + assert unhealthy[0]["error"] == "Timeout exceeded" + assert healthy == [] + assert len(unhealthy) == 1 + assert unhealthy[0]["model"] == "openai/tts-1" + + +@pytest.mark.asyncio +async def test_background_health_check_timeout_does_not_cancel_a_sibling( + httpx_transport: None, respx_mock: respx.MockRouter +) -> None: + never_release: Final = asyncio.Event() + slow_request_started: Final = asyncio.Event() + + async def blocked_response(_: httpx.Request) -> httpx.Response: + slow_request_started.set() + await never_release.wait() + return httpx.Response(200, content=b"audio", headers={"content-type": "audio/mpeg"}) + + respx_mock.post("https://slow.example/v1/audio/speech").mock(side_effect=blocked_response) + fast_upstream: Final = respx_mock.post("https://fast.example/v1/audio/speech").respond( + content=b"audio", + headers={"content-type": "audio/mpeg"}, + ) + model_list: Final = [ + { + "litellm_params": { + "model": "openai/tts-1", + "api_key": "fake-key", + "api_base": "https://slow.example/v1", + }, + "model_info": {"id": "slow", "mode": "audio_speech", "health_check_timeout": 1}, + }, + { + "litellm_params": { + "model": "openai/tts-1", + "api_key": "fake-key", + "api_base": "https://fast.example/v1", + }, + "model_info": {"id": "fast", "mode": "audio_speech", "health_check_timeout": 2}, + }, + ] + + healthy, unhealthy, _ = await asyncio.wait_for( + hc_module.perform_health_check(model_list, max_concurrency=1), + timeout=4, + ) + healthy_model_ids: Final = {endpoint["model_id"] for endpoint in healthy} + unhealthy_model_ids: Final = {endpoint["model_id"] for endpoint in unhealthy} + + assert slow_request_started.is_set() + assert fast_upstream.called + assert healthy_model_ids == {"fast"} + assert unhealthy_model_ids == {"slow"} diff --git a/tests/unit/proxy/test_litellm_pre_call_utils.py b/tests/unit/proxy/test_litellm_pre_call_utils.py index c6c473ae71e..cbf06319d13 100644 --- a/tests/unit/proxy/test_litellm_pre_call_utils.py +++ b/tests/unit/proxy/test_litellm_pre_call_utils.py @@ -16,6 +16,7 @@ from pydantic import ValidationError as PydanticValidationError from starlette.datastructures import Headers import litellm +from litellm.constants import SERVER_STREAMING_CLASSIFICATION_KEY, SERVER_STREAMING_CLASSIFICATION_MARKER from litellm.proxy._types import AddTeamCallback, ProxyException, TeamCallbackMetadata, UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import ( KeyAndTeamLoggingSettings, @@ -850,8 +851,13 @@ def test_initial_snapshot_refresh_clears_a_previous_guardrail_checkpoint() -> No from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot logging_obj: Final = Logging( - model="test-model", messages=[], stream=False, call_type="acompletion", - start_time=datetime.now(), litellm_call_id="new-request", function_id="new-request", + model="test-model", + messages=[], + stream=False, + call_type="acompletion", + start_time=datetime.now(), + litellm_call_id="new-request", + function_id="new-request", ) logging_obj.shadow_eval_request_snapshot = GuardrailRequestSnapshot.capture( {"messages": [{"role": "user", "content": "previous request"}]}, @@ -913,8 +919,14 @@ async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_l } data: Final = {"messages": messages, "metadata": metadata, "proxy_server_request": {}} logging_obj: Final = Logging( - model="test-model", messages=messages, stream=False, call_type="acompletion", - start_time=datetime.now(), litellm_call_id="mask-spend", function_id="mask-spend", kwargs=data, + model="test-model", + messages=messages, + stream=False, + call_type="acompletion", + start_time=datetime.now(), + litellm_call_id="mask-spend", + function_id="mask-spend", + kwargs=data, ) data["litellm_logging_obj"] = logging_obj refresh_proxy_server_request_body_snapshot(data, guardrails_applied=True) @@ -928,9 +940,13 @@ async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_l kwargs, _ = await guardrail.async_logging_hook( kwargs=logging_obj.model_call_details, result=None, call_type="acompletion" ) - stored: Final = json.loads(_get_proxy_server_request_for_spend_logs_payload( - metadata={}, litellm_params=kwargs["litellm_params"], kwargs=kwargs, - )) + stored: Final = json.loads( + _get_proxy_server_request_for_spend_logs_payload( + metadata={}, + litellm_params=kwargs["litellm_params"], + kwargs=kwargs, + ) + ) assert kwargs["messages"] == [{"role": "user", "content": "email [EMAIL]"}] assert stored["messages"] == kwargs["messages"] @@ -1051,6 +1067,66 @@ async def test_add_litellm_data_to_request_strips_string_encoded_admin_injection assert "_pipeline_managed_guardrails" not in other +@pytest.mark.asyncio +@pytest.mark.parametrize( + "forged_field,forged_value", + [ + ("fallback_depth", 1), + ("fallback_depth", True), + ("_target_order", 2), + ("attempted_targets", ["forged"]), + ], +) +async def test_add_litellm_data_to_request_strips_forged_fallback_hop_state( + forged_field: str, forged_value: object +) -> None: + request_mock = MagicMock(spec=Request) + request_mock.url = MagicMock() + request_mock.url.path = "/v1/responses" + request_mock.url.__str__.return_value = "http://localhost/v1/responses" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + updated = await add_litellm_data_to_request( + data={"model": "hop", "input": "hello", forged_field: forged_value}, + request=request_mock, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert forged_field not in updated + assert forged_field not in updated["proxy_server_request"]["body"] + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_keeps_the_request_max_fallbacks_cap() -> None: + request_mock = MagicMock(spec=Request) + request_mock.url = MagicMock() + request_mock.url.path = "/v1/responses" + request_mock.url.__str__.return_value = "http://localhost/v1/responses" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + updated = await add_litellm_data_to_request( + data={"model": "hop", "input": "hello", "max_fallbacks": 0}, + request=request_mock, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["max_fallbacks"] == 0 + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_strips_user_control_fields(): """Strip untrusted proxy-control fields before guardrails, logging, and headers read metadata.""" @@ -1089,6 +1165,7 @@ async def test_add_litellm_data_to_request_strips_user_control_fields(): "messages": [{"role": "user", "content": "hello"}], "mock_response": "free response", "mock_tool_calls": [{"id": "call_1"}], + "is_streaming_request": "caller-value", "disable_global_guardrails": True, "enable_prompt_caching": True, "routing_decision": {"cause": "forged", "routed_model": "spoofed"}, @@ -1114,6 +1191,7 @@ async def test_add_litellm_data_to_request_strips_user_control_fields(): assert "enable_prompt_caching" not in updated assert "routing_decision" not in updated assert "litellm_gateway_injected_cache" not in updated + assert updated["is_streaming_request"] == "caller-value" assert "weights" not in updated assert "_router_weights" not in updated assert "weights" not in updated["proxy_server_request"]["body"] @@ -8803,20 +8881,27 @@ async def test_mcp_credentials_only_removed_from_logging_copies(path: str, custo request.headers = Headers(request.headers) settings: Final = {"mcp_client_side_auth_header_name": custom_auth, "user_header_name": "x-user-id"} server: Final = MCPServer( - server_id="header-test", name="header-test", transport="http", url="https://example.com/mcp", + server_id="header-test", + name="header-test", + transport="http", + url="https://example.com/mcp", extra_headers=["x-service-token", "x-user-id"], ) with ( patch("litellm.proxy.proxy_server.general_settings", settings), patch.dict( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.config_mcp_servers", - {"header-test": server}, clear=True, + {"header-test": server}, + clear=True, ), ): updated: Final = await add_litellm_data_to_request( data={"model": "test-model", "messages": [{"role": "user", "content": "hello"}]}, - request=request, user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), - proxy_config=MagicMock(), general_settings=settings, version="test", + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings=settings, + version="test", ) for header_dict in _all_header_dicts(updated, metadata_name): assert not any(value in json.dumps(header_dict) for value in secrets.values()) @@ -8836,7 +8921,10 @@ def test_signoz_callback_vars_are_scoped_to_the_signoz_callback(): data=AddTeamCallback( callback_name="signoz", callback_type="success", - callback_vars={"signoz_ingestion_key": "team-key", "signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443"}, + callback_vars={ + "signoz_ingestion_key": "team-key", + "signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443", + }, ), team_callback_settings_obj=None, ) @@ -8856,6 +8944,54 @@ def test_signoz_callback_vars_are_scoped_to_the_signoz_callback(): assert under_other.callback_vars == {"langfuse_host": "https://cloud.langfuse.com"} +def test_body_snapshot_excludes_the_server_streaming_marker() -> None: + from litellm.constants import SERVER_STREAMING_CLASSIFICATION_KEY, SERVER_STREAMING_CLASSIFICATION_MARKER + from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot + + proxy_request: Final = {"body": {}} + data: Final = { + "messages": [{"role": "user", "content": "hi"}], + SERVER_STREAMING_CLASSIFICATION_KEY: SERVER_STREAMING_CLASSIFICATION_MARKER, + "proxy_server_request": proxy_request, + } + + refresh_proxy_server_request_body_snapshot(data) + + assert proxy_request == {"body": {"messages": [{"role": "user", "content": "hi"}]}} + + +@pytest.mark.parametrize( + "marker", + [ + SERVER_STREAMING_CLASSIFICATION_MARKER, + json.loads(json.dumps(SERVER_STREAMING_CLASSIFICATION_MARKER)), + ], + ids=["enum", "json-string"], +) +def test_body_snapshot_drops_only_the_marker_and_keeps_caller_value(marker: str) -> None: + from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot + + marker_request: Final = {"body": {}} + marker_data: Final = { + "messages": [{"role": "user", "content": "hi"}], + SERVER_STREAMING_CLASSIFICATION_KEY: marker, + "proxy_server_request": marker_request, + } + + refresh_proxy_server_request_body_snapshot(marker_data) + + assert SERVER_STREAMING_CLASSIFICATION_KEY not in marker_request["body"], marker_request + + caller_request: Final = {"body": {}} + caller_data: Final = { + "messages": [{"role": "user", "content": "hi"}], + SERVER_STREAMING_CLASSIFICATION_KEY: "caller-value", + "proxy_server_request": caller_request, + } + + refresh_proxy_server_request_body_snapshot(caller_data) + + assert caller_request["body"][SERVER_STREAMING_CLASSIFICATION_KEY] == "caller-value", caller_request def test_arize_otlp_protocol_on_a_key_logging_entry_reaches_the_destination(monkeypatch): from litellm.integrations.otel.model.config import is_otel_v2_enabled from litellm.proxy.litellm_pre_call_utils import resolve_tenant_otel_destinations diff --git a/tests/unit/proxy/test_pricing_field_strip.py b/tests/unit/proxy/test_pricing_field_strip.py index a0e25e91f37..395fdc3a16e 100644 --- a/tests/unit/proxy/test_pricing_field_strip.py +++ b/tests/unit/proxy/test_pricing_field_strip.py @@ -26,7 +26,6 @@ from litellm.proxy.litellm_pre_call_utils import ( from litellm.types.utils import CustomPricingLiteLLMParams - def _make_request_mock() -> Request: request_mock = MagicMock(spec=Request) request_mock.url.path = "/v1/chat/completions" @@ -58,9 +57,7 @@ class TestStripClientPricingOverrides: # The strip set is built from the model so additions are picked up # automatically — this test guards against the model and the strip # set drifting apart if someone replaces the auto-derivation later. - assert _CLIENT_PRICING_CONTROL_FIELDS == frozenset( - CustomPricingLiteLLMParams.model_fields.keys() - ) + assert _CLIENT_PRICING_CONTROL_FIELDS == frozenset(CustomPricingLiteLLMParams.model_fields.keys()) # Sanity: the obvious top-level pricing fields are in the set. for field in ( "input_cost_per_token", @@ -184,9 +181,7 @@ class TestStripClientPricingOverrides: verbose_proxy_logger.setLevel(logging.DEBUG) with caplog.at_level(logging.DEBUG, logger=verbose_proxy_logger.name): _strip_client_pricing_overrides({"model": "gpt-4", "temperature": 0.7}) - assert not any( - "pricing" in record.getMessage().lower() for record in caplog.records - ) + assert not any("pricing" in record.getMessage().lower() for record in caplog.records) @pytest.mark.asyncio @@ -211,6 +206,26 @@ async def test_add_litellm_data_to_request_strips_root_pricing_fields(): assert "output_cost_per_token" not in updated +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_preserves_caller_streaming_request(): + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "hi"}], + "is_streaming_request": True, + } + + updated = await add_litellm_data_to_request( + data=data, + request=_make_request_mock(), + user_api_key_dict=_user_api_key_auth(), + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["is_streaming_request"] is True + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_strips_client_disconnect_metadata(): data = { @@ -318,9 +333,7 @@ async def test_add_litellm_data_to_request_skips_strip_with_team_opt_in(): "input_cost_per_token": 0.0001, } - user_auth = _user_api_key_auth( - team_metadata={"allow_client_pricing_override": True} - ) + user_auth = _user_api_key_auth(team_metadata={"allow_client_pricing_override": True}) updated = await add_litellm_data_to_request( data=data, request=_make_request_mock(), diff --git a/tests/unit/proxy/test_proxy_token_counter.py b/tests/unit/proxy/test_proxy_token_counter.py index e7a32816464..e4c4ddce64f 100644 --- a/tests/unit/proxy/test_proxy_token_counter.py +++ b/tests/unit/proxy/test_proxy_token_counter.py @@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest from dotenv import load_dotenv +from pydantic import JsonValue load_dotenv() @@ -20,6 +21,7 @@ from litellm import Router from litellm.llms.bedrock.common_utils import BedrockError from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import ProxyException, TokenCountRequest from litellm.proxy.anthropic_endpoints.endpoints import ( count_tokens as anthropic_count_tokens, @@ -753,41 +755,45 @@ def test_vertex_ai_partner_models_token_counting_endpoint(vertex_location): ) +class _FailingRuntimeHandler(BedrockCountTokensHandler): + def __init__(self, error: Exception) -> None: + super().__init__() + self._error = error + + async def handle_count_tokens_request( + self, + request_data: dict[str, object], + litellm_params: dict[str, object], + resolved_model: str, + client: AsyncHTTPHandler | None = None, + ) -> dict[str, JsonValue]: + raise self._error + + @pytest.mark.asyncio async def test_bedrock_token_counter_error_propagation_bedrock_error(): """ Test that BedrockTokenCounter properly returns error response when BedrockError is raised. Verifies that the status code and error message are preserved. """ - counter = BedrockTokenCounter() + counter = BedrockTokenCounter( + runtime_handler=_FailingRuntimeHandler(BedrockError(status_code=429, message="Rate limit exceeded")) + ) - # Mock the handler to raise BedrockError with specific status code - with patch.object( - counter, "count_tokens", wraps=counter.count_tokens - ) as mock_count: - # We need to patch at the handler level - with patch( - "litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler" - ) as MockHandler: - mock_handler_instance = MockHandler.return_value - mock_handler_instance.handle_count_tokens_request = AsyncMock( - side_effect=BedrockError(status_code=429, message="Rate limit exceeded") - ) + result = await counter.count_tokens( + model_to_use="anthropic.claude-3-sonnet", + messages=[{"role": "user", "content": "hello"}], + contents=None, + deployment={"litellm_params": {}}, + request_model="bedrock/anthropic.claude-3-sonnet", + ) - result = await counter.count_tokens( - model_to_use="anthropic.claude-3-sonnet", - messages=[{"role": "user", "content": "hello"}], - contents=None, - deployment={"litellm_params": {}}, - request_model="bedrock/anthropic.claude-3-sonnet", - ) - - assert result is not None - assert result.error is True - assert result.status_code == 429 - assert "Rate limit exceeded" in result.error_message - assert result.tokenizer_type == "bedrock_api" - assert result.total_tokens == 0 + assert result is not None + assert result.error is True + assert result.status_code == 429 + assert "Rate limit exceeded" in result.error_message + assert result.tokenizer_type == "bedrock_api" + assert result.total_tokens == 0 @pytest.mark.asyncio @@ -795,28 +801,20 @@ async def test_bedrock_token_counter_error_propagation_generic_exception(): """ Test that BedrockTokenCounter returns error response with 500 status for generic exceptions. """ - counter = BedrockTokenCounter() + counter = BedrockTokenCounter(runtime_handler=_FailingRuntimeHandler(Exception("Unexpected error"))) - with patch( - "litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler" - ) as MockHandler: - mock_handler_instance = MockHandler.return_value - mock_handler_instance.handle_count_tokens_request = AsyncMock( - side_effect=Exception("Unexpected error") - ) + result = await counter.count_tokens( + model_to_use="anthropic.claude-3-sonnet", + messages=[{"role": "user", "content": "hello"}], + contents=None, + deployment={"litellm_params": {}}, + request_model="bedrock/anthropic.claude-3-sonnet", + ) - result = await counter.count_tokens( - model_to_use="anthropic.claude-3-sonnet", - messages=[{"role": "user", "content": "hello"}], - contents=None, - deployment={"litellm_params": {}}, - request_model="bedrock/anthropic.claude-3-sonnet", - ) - - assert result is not None - assert result.error is True - assert result.status_code == 500 - assert "Unexpected error" in result.error_message + assert result is not None + assert result.error is True + assert result.status_code == 500 + assert "Unexpected error" in result.error_message @pytest.mark.asyncio diff --git a/tests/unit/proxy/test_proxy_utils.py b/tests/unit/proxy/test_proxy_utils.py index ba61e393281..92bb163d6ec 100644 --- a/tests/unit/proxy/test_proxy_utils.py +++ b/tests/unit/proxy/test_proxy_utils.py @@ -2,7 +2,7 @@ import asyncio import json import os from datetime import datetime -from typing import Any, Dict, List, Optional, Union +from typing import Any, Dict, Final, List, Optional, Union from unittest.mock import Mock import pytest @@ -3329,3 +3329,22 @@ def test_handle_exception_on_proxy_preserves_auth_error_status_code(): result = handle_exception_on_proxy(auth_error) assert int(result.code) == 401, f"Expected 401, got {result.code}" + + +class _RedisDown: + async def async_delete_cache(self, key: str) -> None: + raise ConnectionError("redis is down") + + +async def test_evict_config_param_clears_the_local_layer_and_survives_a_redis_outage() -> None: + from litellm.caching.caching import DualCache + from litellm.proxy.utils import _config_cache_key, evict_config_param + + cache: Final = DualCache(redis_cache=_RedisDown()) + await cache.in_memory_cache.async_set_cache( + _config_cache_key("router_settings"), {"param_name": "router_settings", "param_value": {"fallbacks": []}} + ) + + await evict_config_param("router_settings", cache=cache) + + assert await cache.in_memory_cache.async_get_cache(_config_cache_key("router_settings")) is None diff --git a/tests/unit/proxy/test_route_llm_request.py b/tests/unit/proxy/test_route_llm_request.py index a891e1079d8..a4b34825d95 100644 --- a/tests/unit/proxy/test_route_llm_request.py +++ b/tests/unit/proxy/test_route_llm_request.py @@ -516,12 +516,10 @@ def test_e2e_proxy_config_opts_in_to_the_mock_params_its_suite_sends(): for param in GATED_MOCK_PARAM_NAMES if f"{param}=" in source or f'"{param}"' in source ) - assert senders, "expected the E2E suite to still exercise the gated mock testing params" - config = yaml.safe_load((repo_root / "proxy_server_config.yaml").read_text(encoding="utf-8")) general_settings = config.get("general_settings") or {} - assert general_settings.get(MOCK_TESTING_CONFIG_KEY) is True, ( + assert not senders or general_settings.get(MOCK_TESTING_CONFIG_KEY) is True, ( f"proxy_server_config.yaml must set general_settings.{MOCK_TESTING_CONFIG_KEY}: true — " f"the E2E suite sends gated mock testing params ({', '.join(sorted(senders))}) " "and the proxy rejects them with a 400 otherwise" @@ -1728,3 +1726,37 @@ async def test_route_request_without_model_on_model_routed_endpoint_is_a_400(): assert exc_info.value.code == "400" assert exc_info.value.param == "model" + + +@pytest.mark.asyncio +async def test_route_request_router_settings_override_skips_null_fields(): + """ + A key or team saved from the dashboard stores every unset router setting as null. Those nulls + must not reach the router as explicit per-request values, or they switch the router-level + fallbacks and retries off for that key. + """ + data: Final = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "Hello"}], + "stream": True, + "router_settings_override": { + "fallbacks": None, + "context_window_fallbacks": None, + "num_retries": None, + "model_group_retry_policy": None, + "timeout": 600, + }, + } + + llm_router: Final = MagicMock() + llm_router.acompletion.return_value = "success" + + response: Final = await route_request(data, llm_router, None, "acompletion") + + assert response == "success" + call_kwargs: Final = llm_router.acompletion.call_args[1] + assert call_kwargs["timeout"] == 600 + assert "fallbacks" not in call_kwargs + assert "context_window_fallbacks" not in call_kwargs + assert "num_retries" not in call_kwargs + assert "model_group_retry_policy" not in call_kwargs diff --git a/tests/unit/proxy/utils/helpers/test_month_end_projection.py b/tests/unit/proxy/utils/helpers/test_month_end_projection.py index 8c1a4ae746d..1f2a88b209d 100644 --- a/tests/unit/proxy/utils/helpers/test_month_end_projection.py +++ b/tests/unit/proxy/utils/helpers/test_month_end_projection.py @@ -1,4 +1,5 @@ -from datetime import date, timedelta +from datetime import date, datetime, timedelta, timezone +from typing import Final import pytest @@ -211,3 +212,181 @@ def test_get_projected_spend_over_limit_raises_when_today_missing(monkeypatch): monkeypatch.setattr("litellm.proxy.utils.date", _Broken) with pytest.raises(RuntimeError): get_projected_spend_over_limit(current_spend=1.0, soft_budget_limit=1.0) + + +UTC_RESET_JAN_16: Final = datetime(2024, 1, 16, tzinfo=timezone.utc) + + +def test_daily_budget_projects_pace_to_the_reset_not_month_end(): + result: Final = get_projected_spend_over_limit( + current_spend=8.0, + soft_budget_limit=10.0, + budget_duration="1d", + budget_reset_at=UTC_RESET_JAN_16, + now=datetime(2024, 1, 15, 12, tzinfo=timezone.utc), + ) + assert result == (16.0, date(2024, 1, 15)) + + +def test_daily_budget_late_in_the_day_does_not_double_todays_spend(): + assert ( + get_projected_spend_over_limit( + current_spend=6.0, + soft_budget_limit=10.0, + budget_duration="1d", + budget_reset_at=UTC_RESET_JAN_16, + now=datetime(2024, 1, 15, 23, tzinfo=timezone.utc), + ) + is None + ) + + +def test_naive_reset_time_is_read_as_utc(): + result: Final = get_projected_spend_over_limit( + current_spend=8.0, + soft_budget_limit=10.0, + budget_duration="1d", + budget_reset_at=datetime(2024, 1, 16), + now=datetime(2024, 1, 15, 12), + ) + assert result == (16.0, date(2024, 1, 15)) + + +def test_weekly_budget_measures_pace_from_the_previous_reset(): + result: Final = get_projected_spend_over_limit( + current_spend=8.0, + soft_budget_limit=10.0, + budget_duration="7d", + budget_reset_at=datetime(2024, 1, 18, tzinfo=timezone.utc), + now=datetime(2024, 1, 15, tzinfo=timezone.utc), + ) + assert result == (14.0, date(2024, 1, 16)) + + +def test_monthly_budget_window_starts_on_the_previous_reset_day(): + result: Final = get_projected_spend_over_limit( + current_spend=8.0, + soft_budget_limit=10.0, + budget_duration="1mo", + budget_reset_at=datetime(2024, 3, 1, tzinfo=timezone.utc), + now=datetime(2024, 2, 5, tzinfo=timezone.utc), + ) + assert result == (pytest.approx(58.0), date(2024, 2, 6)) + + +def test_sub_day_budget_is_measured_in_hours_not_days(): + result: Final = get_projected_spend_over_limit( + current_spend=3.0, + soft_budget_limit=5.0, + budget_duration="4h", + budget_reset_at=datetime(2024, 1, 15, 12, tzinfo=timezone.utc), + now=datetime(2024, 1, 15, 10, tzinfo=timezone.utc), + ) + assert result == (pytest.approx(6.0), date(2024, 1, 15)) + + +def test_first_minutes_after_a_reset_do_not_project_a_burst_across_the_window(): + assert ( + get_projected_spend_over_limit( + current_spend=0.10, + soft_budget_limit=10.0, + budget_duration="1d", + budget_reset_at=UTC_RESET_JAN_16, + now=datetime(2024, 1, 15, 0, 1, tzinfo=timezone.utc), + ) + is None + ) + + +def test_exceed_date_is_reported_in_the_reset_timezone(): + result: Final = get_projected_spend_over_limit( + current_spend=9.5, + soft_budget_limit=10.0, + budget_duration="1d", + budget_reset_at=datetime(2024, 1, 17, tzinfo=timezone(timedelta(hours=2))), + now=datetime(2024, 1, 15, 22, 10, tzinfo=timezone.utc), + ) + assert result is not None + assert result[1] == date(2024, 1, 16) + + +@pytest.mark.parametrize("soft_budget_limit", [9.99, 13.9, 13.99]) +def test_exceed_date_never_lands_after_the_reset(soft_budget_limit): + result: Final = get_projected_spend_over_limit( + current_spend=8.0, + soft_budget_limit=soft_budget_limit, + budget_duration="7d", + budget_reset_at=datetime(2024, 1, 18, tzinfo=timezone.utc), + now=datetime(2024, 1, 15, tzinfo=timezone.utc), + ) + assert result is not None + assert result[1] <= date(2024, 1, 18) + + +@pytest.mark.parametrize("current_spend, expected", [(4.0, False), (8.0, True)]) +def test_is_projected_spend_over_limit_follows_the_reset_window(current_spend, expected): + assert ( + is_projected_spend_over_limit( + current_spend=current_spend, + soft_budget_limit=10.0, + budget_duration="1d", + budget_reset_at=UTC_RESET_JAN_16, + now=datetime(2024, 1, 15, 12, tzinfo=timezone.utc), + ) + is expected + ) + + +def test_without_duration_keeps_month_end_behavior(monkeypatch): + _freeze_today(monkeypatch, date(2024, 1, 15)) + result: Final = get_projected_spend_over_limit(current_spend=8.0, soft_budget_limit=10.0) + assert result is not None + assert result[0] == pytest.approx(8.0 + (8.0 / 14) * 16) + + +def test_thirty_day_budget_measures_pace_from_the_first_of_the_month(): + result: Final = get_projected_spend_over_limit( + current_spend=1.0, + soft_budget_limit=10.0, + budget_duration="30d", + budget_reset_at=datetime(2024, 11, 1, tzinfo=timezone.utc), + now=datetime(2024, 10, 4, tzinfo=timezone.utc), + ) + assert result == (pytest.approx(1.0 + 672 / 72), date(2024, 10, 31)) + + +def test_hour_spelling_the_scheduler_resets_at_midnight_does_not_alert_before_midnight(): + assert ( + get_projected_spend_over_limit( + current_spend=1.0, + soft_budget_limit=10.0, + budget_duration="1hr", + budget_reset_at=UTC_RESET_JAN_16, + now=datetime(2024, 1, 15, 22, tzinfo=timezone.utc), + ) + is None + ) + + +def test_unrecognized_duration_projects_within_the_daily_window_it_resets_on(): + result: Final = get_projected_spend_over_limit( + current_spend=8.0, + soft_budget_limit=10.0, + budget_duration="fortnightly", + budget_reset_at=UTC_RESET_JAN_16, + now=datetime(2024, 1, 15, 12, tzinfo=timezone.utc), + ) + assert result == (16.0, date(2024, 1, 15)) + + +def test_reset_more_than_one_window_ahead_does_not_project(): + assert ( + get_projected_spend_over_limit( + current_spend=9.0, + soft_budget_limit=10.0, + budget_duration="1d", + budget_reset_at=UTC_RESET_JAN_16, + now=datetime(2024, 1, 14, 12, tzinfo=timezone.utc), + ) + is None + ) diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py b/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py index 0d408de9ec6..cdaf5c1ee3b 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_config_param_cache.py @@ -12,8 +12,9 @@ Symbols pinned here: from __future__ import annotations +import logging from types import SimpleNamespace -from typing import Any, List +from typing import Any, Final, List from unittest.mock import AsyncMock, MagicMock import pytest @@ -214,14 +215,23 @@ async def test_invalidate_config_param_evicts_from_cache( @pytest.mark.asyncio -async def test_invalidate_config_param_propagates_cache_error( - _swap_config_cache: Any, +async def test_invalidate_config_param_survives_a_cache_error( + _swap_config_cache: Any, caplog: pytest.LogCaptureFixture ) -> None: _swap_config_cache.async_delete_cache = AsyncMock( side_effect=ConnectionError("redis down") ) - with pytest.raises(ConnectionError): + caplog.set_level(logging.WARNING, logger="LiteLLM Proxy") + utils_mod.verbose_proxy_logger.addHandler(caplog.handler) + try: await invalidate_config_param("p5") + finally: + utils_mod.verbose_proxy_logger.removeHandler(caplog.handler) + actual: Final = { + "delete_calls": _swap_config_cache.async_delete_cache.await_count, + "warned": "config cache eviction of p5 failed: redis down" in caplog.text, + } + assert actual == {"delete_calls": 1, "warned": True} @pytest.mark.asyncio diff --git a/tests/unit/rag/test_main.py b/tests/unit/rag/test_main.py index 81318a1e113..87599d32c22 100644 --- a/tests/unit/rag/test_main.py +++ b/tests/unit/rag/test_main.py @@ -11,9 +11,12 @@ aquery carries the completion response with real usage and cost. """ import asyncio +import datetime import json +from collections.abc import Callable, Mapping +from concurrent.futures import Future from types import MappingProxyType -from typing import Final +from typing import Final, cast from unittest.mock import patch import httpx @@ -22,8 +25,10 @@ import respx from pydantic import ValidationError import litellm +import litellm.utils as litellm_utils from litellm._internal_context import is_internal_call from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils import litellm_logging from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.types.utils import CallTypes, ModelResponse @@ -45,11 +50,27 @@ async def _drain_logging_worker() -> None: class RecordingLogger(CustomLogger): def __init__(self): super().__init__() - self.success_events = [] + self.success_events: Final[list[dict[str, object]]] = [] + self.sync_success_events: Final[list[dict[str, object]]] = [] - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + async def async_log_success_event( + self, + kwargs: Mapping[str, object], + response_obj: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: self.success_events.append({"kwargs": kwargs, "response_obj": response_obj}) + def log_success_event( + self, + kwargs: Mapping[str, object], + response_obj: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + self.sync_success_events.append({"kwargs": kwargs, "response_obj": response_obj}) + @pytest.mark.asyncio @pytest.mark.parametrize("use_router", [False, True]) @@ -107,6 +128,81 @@ async def test_aquery_single_billing_event_carries_completion_usage_and_cost(use assert standard_logging_object["response_cost"] > 0 +@pytest.mark.asyncio +@pytest.mark.parametrize("use_router", [False, True]) +async def test_aquery_vector_store_search_sub_call_logs_no_sync_success_event_on_the_parent( + use_router: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + await _drain_logging_worker() + recording_logger: Final = RecordingLogger() + + class InlineExecutor: + def submit( + self, + fn: Callable[..., object], + *args: object, + **kwargs: object, + ) -> Future[object]: + future: Final = Future[object]() + future.set_result(fn(*args, **kwargs)) + return future + + monkeypatch.setattr(litellm_utils, "executor", InlineExecutor()) + monkeypatch.setattr(litellm_logging, "executor", InlineExecutor()) + monkeypatch.setattr(litellm, "callbacks", [recording_logger]) + + router_kwargs: Final = ( + { + "router": litellm.Router( + model_list=[ + { + "model_name": "gpt-4o-mini", + "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"}, + } + ] + ) + } + if use_router + else {} + ) + + response: Final = await litellm.aquery( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "What is the secret project codename?"}], + retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"}, + mock_response="The secret project codename is AZURE-FALCON-42.", + **router_kwargs, + ) + assert isinstance(response, ModelResponse), "aquery should return its completion response" + assert is_internal_call.get() is False, "aquery should restore the internal-call context" + await _drain_logging_worker() + + assert len(recording_logger.success_events) == 1, "aquery should run one async success callback" + event: Final = recording_logger.success_events[0] + response_obj: Final = event["response_obj"] + assert isinstance(response_obj, ModelResponse), "the async callback should receive the completion response" + assert response_obj.usage.total_tokens > 0, "the async callback response should include completion usage" + + event_kwargs: Final = cast(Mapping[str, object], event["kwargs"]) + standard_logging_object: Final = cast(Mapping[str, object], event_kwargs["standard_logging_object"]) + assert cast(str, standard_logging_object["call_type"]) == "aquery", "the async event should be for aquery" + assert cast(int, standard_logging_object["total_tokens"]) > 0, ( + "the async event should include total completion tokens" + ) + assert cast(int, standard_logging_object["prompt_tokens"]) > 0, "the async event should include prompt tokens" + assert cast(int, standard_logging_object["completion_tokens"]) > 0, ( + "the async event should include completion tokens" + ) + assert cast(float, standard_logging_object["response_cost"]) > 0, "the async event should include completion cost" + sync_response_types: Final = [ + type(event["response_obj"]).__name__ for event in recording_logger.sync_success_events + ] + assert recording_logger.sync_success_events == [], ( + "an internal sub-call must not log on the parent's logging object; " + f"sync success event response types: {sync_response_types}" + ) + + @pytest.mark.asyncio async def test_aquery_response_hidden_params_carry_completion_cost(): """ diff --git a/tests/unit/repositories/test_repositories.py b/tests/unit/repositories/test_repositories.py index 3aed934f40c..f2e033c9141 100644 --- a/tests/unit/repositories/test_repositories.py +++ b/tests/unit/repositories/test_repositories.py @@ -2221,6 +2221,37 @@ class TestPrismaTableRepository: with pytest.raises(RuntimeError, match="No DB Connected"): _ = repo.table + @pytest.mark.asyncio + async def test_managed_file_repository_updates_existing_file_object_only(self): + from litellm.repositories.managed_file_repository import ManagedFileRepository + from litellm.types.llms.openai import OpenAIFileObject + + class UpdateManyMockTable(MockTable): + async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> int: + return int(await self.update(where, data) is not None) + + file_table = UpdateManyMockTable(pk_field="unified_file_id") + await file_table.create({"unified_file_id": "existing-file", "file_object": "{}"}) + prisma_client = SimpleNamespace(db=SimpleNamespace(litellm_managedfiletable=file_table)) + repository = ManagedFileRepository(prisma_client) + file_object = OpenAIFileObject( + id="existing-file", + object="file", + bytes=836, + created_at=456, + filename="output.jsonl", + purpose="batch_output", + status="processed", + ) + + assert await repository.update_file_object("existing-file", file_object) is True + stored_row = await file_table.find_unique(where={"unified_file_id": "existing-file"}) + assert stored_row is not None + assert stored_row.file_object == file_object.model_dump_json() + + assert await repository.update_file_object("missing-file", file_object) is False + assert await file_table.find_unique(where={"unified_file_id": "missing-file"}) is None + CONFIG_SYNCED_TABLE_NAMES = frozenset( { "litellm_agentstable", diff --git a/tests/unit/responses/litellm_completion_transformation/test_anthropic_responses_bridge.py b/tests/unit/responses/litellm_completion_transformation/test_anthropic_responses_bridge.py index 42af50e39bf..677ee0cd2d9 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_anthropic_responses_bridge.py +++ b/tests/unit/responses/litellm_completion_transformation/test_anthropic_responses_bridge.py @@ -1,8 +1,13 @@ +from collections.abc import Mapping, Sequence +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest +from pydantic import BaseModel import litellm +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.responses.litellm_completion_transformation.handler import ( LiteLLMCompletionTransformationHandler, ) @@ -36,7 +41,9 @@ def test_response_api_handler_merges_metadata_and_service_tier_without_error(): async def test_async_response_api_handler_merges_trace_id_without_error(): handler = LiteLLMCompletionTransformationHandler() - async def fake_session_handler(previous_response_id, litellm_completion_request): + async def fake_session_handler( + previous_response_id: str, litellm_completion_request: dict[str, object], instructions: str | None = None + ) -> dict[str, object]: litellm_completion_request["litellm_trace_id"] = "session-trace" return litellm_completion_request @@ -94,3 +101,201 @@ async def test_aresponses_forwards_timeout_to_acompletion(): "this means Router(timeout=N) silently fails for providers on the " "completion transformation path." ) + + +class _FakeSpendLogsDB: + def __init__(self, spend_logs: Sequence[Mapping[str, object]]) -> None: + self._spend_logs = spend_logs + + async def query_raw(self, query: str, *args: object) -> Sequence[Mapping[str, object]]: + return self._spend_logs + + +class _FakePrismaClient: + def __init__(self, spend_logs: Sequence[Mapping[str, object]]) -> None: + self.db = _FakeSpendLogsDB(spend_logs) + + +class _AnthropicBlock(BaseModel, frozen=True): + type: str + id: str | None = None + tool_use_id: str | None = None + text: str | None = None + + +class _AnthropicMessage(BaseModel, frozen=True): + role: str + content: tuple[_AnthropicBlock, ...] + + +class _AnthropicRequest(BaseModel, frozen=True): + messages: tuple[_AnthropicMessage, ...] + system: tuple[_AnthropicBlock, ...] + + +class _RecordingAnthropicMessages: + def __init__(self) -> None: + self.request: _AnthropicRequest | None = None + + def __call__(self, request: httpx.Request) -> httpx.Response: + self.request = _AnthropicRequest.model_validate_json(request.content) + return httpx.Response( + 200, + json={ + "id": "msg_second_turn", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-5-5", + "content": [{"type": "text", "text": "Il fait 47C."}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 1, "output_tokens": 1}, + }, + request=request, + ) + + +_MODEL: Final = "anthropic/claude-sonnet-5-5" +_FIRST_TURN_INSTRUCTIONS: Final = "Be terse." +_FIRST_TURN: Final = { + "request_id": "chatcmpl-first-turn", + "call_type": "aresponses", + "session_id": "session-1", + "proxy_server_request": { + "model": _MODEL, + "input": "What is the weather in Tokyo?", + "instructions": _FIRST_TURN_INSTRUCTIONS, + }, + "response": { + "id": "chatcmpl-first-turn", + "object": "chat.completion", + "created": 0, + "model": _MODEL, + "choices": [ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "toolu_weather", + "type": "function", + "function": {"name": "get_weather", "arguments": '{"city": "Tokyo"}'}, + } + ], + }, + } + ], + }, +} +_EXPECTED_MESSAGES: Final = [ + ("user", [("text", "What is the weather in Tokyo?")]), + ("assistant", [("tool_use", "toolu_weather")]), + ("user", [("tool_result", "toolu_weather")]), +] + + +async def _continue_first_turn_with_tool_output(instructions: str | None) -> _AnthropicRequest: + anthropic: Final = _RecordingAnthropicMessages() + with patch("litellm.proxy.proxy_server.prisma_client", _FakePrismaClient([_FIRST_TURN])): + await litellm.aresponses( + model=_MODEL, + previous_response_id="chatcmpl-first-turn", + input=[{"type": "function_call_output", "call_id": "toolu_weather", "output": "47C"}], + instructions=instructions, + tools=[ + { + "type": "function", + "name": "get_weather", + "parameters": {"type": "object", "properties": {"city": {"type": "string"}}}, + } + ], + api_key="sk-ant-fake", + client=AsyncHTTPHandler(transport=httpx.MockTransport(anthropic)), + ) + assert anthropic.request is not None + return anthropic.request + + +def _message_shapes(request: _AnthropicRequest) -> list[tuple[str, list[tuple[str, str | None]]]]: + return [ + (message.role, [(block.type, block.id or block.tool_use_id or block.text) for block in message.content]) + for message in request.messages + ] + + +@pytest.mark.asyncio +async def test_previous_response_id_tool_output_with_new_instructions_builds_valid_anthropic_request() -> None: + """ + A continuation that resends `instructions` must not land a system message between the replayed + tool_use and its tool_result, and the previous turn's instructions do not carry over (OpenAI semantics) + """ + request: Final = await _continue_first_turn_with_tool_output(instructions="Answer in French.") + + assert _message_shapes(request) == _EXPECTED_MESSAGES + assert [block.text for block in request.system] == ["Answer in French."] + + +@pytest.mark.asyncio +async def test_previous_response_id_tool_output_without_instructions_keeps_the_previous_turns() -> None: + """ + A continuation that sends no `instructions` keeps the previous turn's instructions as the system prompt + """ + request: Final = await _continue_first_turn_with_tool_output(instructions=None) + + assert _message_shapes(request) == _EXPECTED_MESSAGES + assert [block.text for block in request.system] == [_FIRST_TURN_INSTRUCTIONS] + + +def _second_turn(instructions: str | None) -> Mapping[str, object]: + return { + "request_id": "chatcmpl-second-turn", + "call_type": "aresponses", + "session_id": "session-1", + "proxy_server_request": { + "model": _MODEL, + "input": [{"type": "function_call_output", "call_id": "toolu_weather", "output": "47C"}], + **({"instructions": instructions} if instructions else {}), + }, + "response": { + "id": "chatcmpl-second-turn", + "object": "chat.completion", + "created": 0, + "model": _MODEL, + "choices": [ + {"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "47C in Tokyo."}} + ], + }, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("second_turn_instructions", "expected_system"), + [("Answer in French.", "Answer in French."), (None, _FIRST_TURN_INSTRUCTIONS)], +) +async def test_previous_response_id_without_instructions_carries_the_latest_ones_ahead_of_a_tool_roundtrip( + second_turn_instructions: str | None, expected_system: str +) -> None: + anthropic: Final = _RecordingAnthropicMessages() + with patch( + "litellm.proxy.proxy_server.prisma_client", + _FakePrismaClient([_FIRST_TURN, _second_turn(second_turn_instructions)]), + ): + await litellm.aresponses( + model=_MODEL, + previous_response_id="chatcmpl-second-turn", + input="What about Osaka?", + api_key="sk-ant-fake", + client=AsyncHTTPHandler(transport=httpx.MockTransport(anthropic)), + ) + + assert anthropic.request is not None + assert _message_shapes(anthropic.request) == [ + *_EXPECTED_MESSAGES, + ("assistant", [("text", "47C in Tokyo.")]), + ("user", [("text", "What about Osaka?")]), + ] + assert [block.text for block in anthropic.request.system] == [expected_system] diff --git a/tests/unit/responses/litellm_completion_transformation/test_google_ai_studio_responses_wire.py b/tests/unit/responses/litellm_completion_transformation/test_google_ai_studio_responses_wire.py new file mode 100644 index 00000000000..2b71b70c6bf --- /dev/null +++ b/tests/unit/responses/litellm_completion_transformation/test_google_ai_studio_responses_wire.py @@ -0,0 +1,158 @@ +import json +from typing import Final + +import httpx +import pytest +import respx +from pydantic import TypeAdapter + +import litellm +from litellm.types.llms.openai import ResponseCompletedEvent, ResponsesAPIResponse +from tests.unit.proxy.conftest import httpx_transport + +pytestmark: Final = pytest.mark.usefixtures(httpx_transport.__name__) +_GEMINI_URL: Final = ( + "https://generativelanguage.googleapis.com/v1beta/models/" + "gemini-2.5-flash:(?:generateContent|streamGenerateContent).*" +) +_JSON_OBJECT: Final = TypeAdapter(dict[str, object]) + + +def _function_call_response(signature: str) -> dict[str, object]: + return { + "candidates": [ + { + "content": { + "role": "model", + "parts": [ + { + "functionCall": { + "name": "get_weather", + "args": {"location": "San Francisco"}, + }, + "thoughtSignature": signature, + } + ], + }, + "finishReason": "STOP", + "index": 0, + } + ], + "usageMetadata": { + "promptTokenCount": 4, + "candidatesTokenCount": 2, + "totalTokenCount": 6, + }, + } + + +@pytest.mark.asyncio +async def test_web_search_preview_is_sent_as_google_search() -> None: + response_body: Final = { + "candidates": [ + { + "content": {"role": "model", "parts": [{"text": "Search completed"}]}, + "finishReason": "STOP", + "index": 0, + } + ], + "usageMetadata": { + "promptTokenCount": 2, + "candidatesTokenCount": 2, + "totalTokenCount": 4, + }, + } + + with respx.mock() as mock_router: + route: Final = mock_router.post(url__regex=_GEMINI_URL).mock( + return_value=httpx.Response(status_code=200, json=response_body) + ) + response: Final = await litellm.aresponses( + model="gemini/gemini-2.5-flash", + api_key="test-key", + input="Find current weather", + tools=[{"type": "web_search_preview", "search_context_size": "low"}], + ) + requests: Final = tuple(route.calls) + + assert isinstance(response, ResponsesAPIResponse) + assert response.output_text == "Search completed" + assert len(requests) == 1 + request_body: Final = _JSON_OBJECT.validate_json(requests[0].request.content) + assert request_body["tools"] == [{"googleSearch": {}}] + + +@pytest.mark.asyncio +async def test_gemini_function_call_preserves_thought_signature() -> None: + signature: Final = "gemini-thought-signature" + response_body: Final = _function_call_response(signature) + + with respx.mock() as mock_router: + route: Final = mock_router.post(url__regex=_GEMINI_URL).mock( + return_value=httpx.Response(status_code=200, json=response_body) + ) + response: Final = await litellm.aresponses( + model="gemini/gemini-2.5-flash", + api_key="test-key", + input="What is the weather in San Francisco?", + tools=[ + { + "type": "function", + "name": "get_weather", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string"}}, + }, + } + ], + ) + requests: Final = tuple(route.calls) + + assert isinstance(response, ResponsesAPIResponse) + function_calls: Final = tuple(item for item in response.output if item.type == "function_call") + assert len(function_calls) == 1 + assert function_calls[0].name == "get_weather" + assert function_calls[0].provider_specific_fields == {"thought_signature": signature} + assert len(requests) == 1 + + +@pytest.mark.asyncio +async def test_gemini_streaming_function_call_preserves_thought_signature() -> None: + signature: Final = "gemini-stream-thought-signature" + response_body: Final = _function_call_response(signature) + event_stream: Final = f"data: {json.dumps(response_body)}\n\n" + + with respx.mock() as mock_router: + route: Final = mock_router.post(url__regex=_GEMINI_URL).mock( + return_value=httpx.Response( + status_code=200, + content=event_stream, + headers={"content-type": "text/event-stream"}, + ) + ) + stream: Final = await litellm.aresponses( + model="gemini/gemini-2.5-flash", + api_key="test-key", + input="What is the weather in San Francisco?", + stream=True, + tools=[ + { + "type": "function", + "name": "get_weather", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string"}}, + }, + } + ], + ) + events: Final = tuple([event async for event in stream]) + requests: Final = tuple(route.calls) + + completed_events: Final = tuple(event for event in events if isinstance(event, ResponseCompletedEvent)) + assert len(completed_events) == 1 + function_calls: Final = tuple(item for item in completed_events[0].response.output if item.type == "function_call") + assert len(function_calls) == 1 + assert function_calls[0].name == "get_weather" + assert function_calls[0].provider_specific_fields == {"thought_signature": signature} + assert len(requests) == 1 diff --git a/tests/unit/responses/test_dispatch.py b/tests/unit/responses/test_dispatch.py index befe2d0000e..e6a0193a063 100644 --- a/tests/unit/responses/test_dispatch.py +++ b/tests/unit/responses/test_dispatch.py @@ -165,13 +165,13 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None: request, call_args, call_kwargs = captured[0] assert result is response - assert request.bound["model"] == "anthropic/claude-sonnet-4-5" - assert request.bound["input"] is INPUT - assert request.bound["stream"] is True - assert request.bound["api_key"] == "sk-test" - assert request.bound["base_url"] == "https://example.invalid" - assert request.bound["custom_llm_provider"] == "anthropic" - assert request.bound["extra_headers"] is extra_headers + assert request.resolved["model"] == "anthropic/claude-sonnet-4-5" + assert request.resolved["input"] is INPUT + assert request.resolved["stream"] is True + assert request.resolved["api_key"] == "sk-test" + assert request.resolved["base_url"] == "https://example.invalid" + assert request.resolved["custom_llm_provider"] == "anthropic" + assert request.resolved["extra_headers"] is extra_headers assert request.kwargs == kwargs assert request.kwargs["litellm_metadata"] is metadata assert call_args == args @@ -262,7 +262,7 @@ def test_public_responses_routes_through_dispatch(monkeypatch: pytest.MonkeyPatc finally: NATIVE_RESPONSES.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["gpt-4o"] + assert [request.resolved["model"] for request in captured] == ["gpt-4o"] @pytest.mark.asyncio @@ -284,7 +284,7 @@ async def test_public_aresponses_routes_through_dispatch(monkeypatch: pytest.Mon finally: NATIVE_ARESPONSES.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["gpt-4o"] + assert [request.resolved["model"] for request in captured] == ["gpt-4o"] def test_responses_with_retries_uses_the_dispatch_entrypoint(monkeypatch: pytest.MonkeyPatch) -> None: @@ -307,6 +307,6 @@ def test_positional_parameters_remain_available_to_native_projection() -> None: include: Final = ["reasoning.encrypted_content"] request: Final = _DISPATCH.request((INPUT, "openai/test-model", include, "Be brief", 16), {}) assert request is not None - assert request.bound["include"] is include - assert request.bound["instructions"] == "Be brief" - assert request.bound["max_output_tokens"] == 16 + assert request.resolved["include"] is include + assert request.resolved["instructions"] == "Be brief" + assert request.resolved["max_output_tokens"] == 16 diff --git a/tests/unit/responses/test_responses_api_lifecycle.py b/tests/unit/responses/test_responses_api_lifecycle.py new file mode 100644 index 00000000000..8805bc97d10 --- /dev/null +++ b/tests/unit/responses/test_responses_api_lifecycle.py @@ -0,0 +1,109 @@ +from typing import Final, Literal, TypeAlias + +import httpx +import openai +import pytest +import respx + +import litellm +from tests.unit.proxy.conftest import httpx_transport + +pytestmark: Final = pytest.mark.usefixtures(httpx_transport.__name__) +Provider: TypeAlias = Literal["anthropic", "gemini"] +Operation: TypeAlias = Literal["delete", "get", "cancel"] +CancelProvider: TypeAlias = Literal["openai", "azure"] + + +def _invoke_unsupported_response_operation(provider: Provider, operation: Operation) -> object: + match operation: + case "delete": + return litellm.delete_responses( + response_id="resp_unsupported", + custom_llm_provider=provider, + api_key="sk-test", + ) + case "get": + return litellm.get_responses( + response_id="resp_unsupported", + custom_llm_provider=provider, + api_key="sk-test", + ) + case "cancel": + return litellm.cancel_responses( + response_id="resp_unsupported", + custom_llm_provider=provider, + api_key="sk-test", + ) + + +async def _invoke_cancel_response(sync_mode: bool, call_kwargs: dict[str, object]) -> object: + if sync_mode: + return litellm.cancel_responses(**call_kwargs) + return await litellm.acancel_responses(**call_kwargs) + + +@pytest.mark.parametrize( + ("provider", "operation"), + ( + ("anthropic", "delete"), + ("anthropic", "get"), + ("anthropic", "cancel"), + ("gemini", "delete"), + ("gemini", "get"), + ("gemini", "cancel"), + ), +) +def test_unsupported_response_lifecycle_operation_fails_before_http(provider: Provider, operation: Operation) -> None: + with respx.mock() as mock_router: + with pytest.raises(litellm.APIConnectionError) as exc_info: + _invoke_unsupported_response_operation(provider, operation) + calls: Final = tuple(mock_router.calls) + + assert exc_info.value.status_code == 500 + assert f"not supported for {provider}" in str(exc_info.value) + assert calls == () + + +@pytest.mark.parametrize( + ("provider", "sync_mode"), + (("openai", True), ("openai", False), ("azure", True), ("azure", False)), +) +@pytest.mark.asyncio +async def test_cancel_responses_404_surfaces_openai_api_error_with_exact_url( + provider: CancelProvider, sync_mode: bool +) -> None: + response_id: Final = "resp_missing" + error_body: Final = { + "error": { + "message": "Response was not found", + "type": "invalid_request_error", + "code": "not_found", + } + } + api_base: Final = ( + "https://api.openai.com/v1" if provider == "openai" else "https://example-resource.openai.azure.com" + ) + expected_url: Final = ( + f"{api_base}/responses/{response_id}/cancel" + if provider == "openai" + else f"{api_base}/openai/responses/{response_id}/cancel?api-version=2025-03-01-preview" + ) + + with respx.mock() as mock_router: + route: Final = mock_router.post(expected_url).mock( + return_value=httpx.Response(status_code=404, json=error_body) + ) + call_kwargs: Final = { + "custom_llm_provider": provider, + "api_key": "sk-test", + "api_base": api_base, + "api_version": "2025-03-01-preview", + "response_id": response_id, + } + with pytest.raises(openai.APIError) as exc_info: + await _invoke_cancel_response(sync_mode, call_kwargs) + requests: Final = tuple(route.calls) + + assert exc_info.value.status_code == 404 + assert len(requests) == 1 + assert str(requests[0].request.url) == expected_url diff --git a/tests/unit/responses/test_responses_api_request_body.py b/tests/unit/responses/test_responses_api_request_body.py index 3ef3cd8854a..8c8b347be82 100644 --- a/tests/unit/responses/test_responses_api_request_body.py +++ b/tests/unit/responses/test_responses_api_request_body.py @@ -827,6 +827,62 @@ async def test_injection_points_still_reach_a_native_responses_provider(): assert "cache_control_injection_points" not in body +@pytest.mark.asyncio +async def test_injection_points_reach_a_provider_prefixed_native_responses_model(): + """Bedrock Mantle resolves its Responses config from the price map by bare model name, + so predicting the bridge with the ``bedrock_mantle/``-prefixed name read as bridged and + deferred the system point to a chat-completions pass this native path never runs.""" + injected_client = AsyncHTTPHandler() + mock_post = AsyncMock( + return_value=MockResponse(_minimal_responses_api_payload("resp_mantle", "openai.gpt-5.6-sol"), 200) + ) + injected_client.post = mock_post + + await litellm.aresponses( + model="bedrock_mantle/openai.gpt-5.6-sol", + api_key="fake-bearer-token", + aws_region_name="us-east-1", + input=copy.deepcopy(_INJECTION_POINT_INPUT), + cache_control_injection_points=copy.deepcopy(_SYSTEM_INJECTION_POINT), + client=injected_client, + ) + + assert mock_post.call_args.kwargs["url"].endswith("/openai/v1/responses") + body = _sent_body(mock_post) + assert body["input"][0]["content"][0]["prompt_cache_breakpoint"] == {"mode": "explicit"} + assert body["prompt_cache_options"] == {"mode": "implicit"} + assert "cache_control_injection_points" not in body + + +@pytest.mark.asyncio +async def test_injection_points_reach_a_foundry_deployment_of_an_openai_model(monkeypatch): + """The router hands this layer the provider it resolved with the deployment's api_base, but the + hook reads the provider from the request kwargs, which never carry it, and resolving + ``azure_ai/gpt-6-astra`` by name alone reads the ambient ``AZURE_AI_API_BASE``, so next to an + Azure OpenAI one the Foundry deployment got Anthropic marks the Responses transform then stripped.""" + monkeypatch.setenv("AZURE_AI_API_BASE", "https://other-deployment.openai.azure.com") + injected_client = AsyncHTTPHandler() + mock_post = AsyncMock( + return_value=MockResponse(_minimal_responses_api_payload("resp_foundry", "gpt-6-astra"), 200) + ) + injected_client.post = mock_post + + await litellm.aresponses( + model="azure_ai/gpt-6-astra", + custom_llm_provider="azure_ai", + api_key="fake-api-key", + api_base="https://foundry.services.ai.azure.com", + input=copy.deepcopy(_INJECTION_POINT_INPUT), + cache_control_injection_points=copy.deepcopy(_SYSTEM_INJECTION_POINT), + client=injected_client, + ) + + body = _sent_body(mock_post) + assert body["input"][0]["content"][0]["prompt_cache_breakpoint"] == {"mode": "explicit"} + assert body["prompt_cache_options"] == {"mode": "implicit"} + assert "cache_control_injection_points" not in body + + async def _bridged_body(mock_post, *, points, input, instructions="You are a documentation assistant."): injected_client = AsyncHTTPHandler() injected_client.post = mock_post diff --git a/tests/unit/router_strategy/test_lowest_tpm_rpm.py b/tests/unit/router_strategy/test_lowest_tpm_rpm.py index 625f648bec4..1c659f47279 100644 --- a/tests/unit/router_strategy/test_lowest_tpm_rpm.py +++ b/tests/unit/router_strategy/test_lowest_tpm_rpm.py @@ -1,13 +1,14 @@ +import asyncio from datetime import datetime, timedelta from typing import Final from unittest.mock import AsyncMock import pytest -from litellm import Router +from litellm import Router, token_counter from litellm.caching.dual_cache import DualCache from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage -from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict +from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict, RouterErrors MODEL_GROUP: Final = "lowest-tpm-router" HIGH_USAGE_DEPLOYMENT_ID: Final = "highest-usage" @@ -117,3 +118,69 @@ async def test_v2_subclass_overriding_async_get_available_deployments_with_the_o f"from {HIGH_USAGE_DEPLOYMENT_ID}", f"from {LOW_USAGE_DEPLOYMENT_ID}", } + + +def _rate_limited_router(num_allowed_send: int) -> tuple[Router, tuple[list[dict[str, str]], ...]]: + conversations: Final = tuple( + [{"role": "user", "content": f"{index}. Hey, how's it going?"}] for index in range(num_allowed_send) + ) + tpm: Final = sum(token_counter(model="gpt-4o", messages=messages) + 5 for messages in conversations) + deployment: Final = _deployment(LOW_USAGE_DEPLOYMENT_ID) + router: Final = Router( + model_list=[{**deployment, "rpm": num_allowed_send, "tpm": tpm}], + routing_strategy="usage-based-routing", + enable_pre_call_checks=True, + num_retries=0, + ) + return router, conversations + + +def test_usage_based_routing_v1_serves_sync_calls_within_rpm_and_tpm() -> None: + router, conversations = _rate_limited_router(num_allowed_send=3) + responses: Final = [router.completion(model=MODEL_GROUP, messages=messages) for messages in conversations[:2]] + assert [response.choices[0].message.content for response in responses] == [f"from {LOW_USAGE_DEPLOYMENT_ID}"] * 2 + + +@pytest.mark.asyncio +async def test_usage_based_routing_v1_serves_async_calls_within_rpm_and_tpm() -> None: + router, conversations = _rate_limited_router(num_allowed_send=3) + responses: Final = await asyncio.gather( + *(router.acompletion(model=MODEL_GROUP, messages=messages) for messages in conversations[:2]) + ) + assert [response.choices[0].message.content for response in responses] == [f"from {LOW_USAGE_DEPLOYMENT_ID}"] * 2 + + +RPM_LIMIT: Final = 3 + + +def _router_with_recorded_rpm(recorded: int, enable_pre_call_checks: bool) -> Router: + router: Final = Router( + model_list=[{**_deployment(LOW_USAGE_DEPLOYMENT_ID), "rpm": RPM_LIMIT}], + routing_strategy="usage-based-routing", + enable_pre_call_checks=enable_pre_call_checks, + num_retries=0, + ) + now: Final = datetime.now() + for offset in range(-1, 2): + router.cache.set_cache( + key=f"{MODEL_GROUP}:rpm:{(now + timedelta(minutes=offset)).strftime('%H-%M')}", + value={LOW_USAGE_DEPLOYMENT_ID: recorded}, + ttl=float("inf"), + ) + return router + + +@pytest.mark.parametrize("enable_pre_call_checks", [True, False]) +def test_usage_based_routing_v1_serves_a_deployment_below_its_recorded_rpm_limit(enable_pre_call_checks: bool) -> None: + router: Final = _router_with_recorded_rpm(recorded=1, enable_pre_call_checks=enable_pre_call_checks) + response: Final = router.completion(model=MODEL_GROUP, messages=[{"role": "user", "content": "hello"}]) + assert response.choices[0].message.content == f"from {LOW_USAGE_DEPLOYMENT_ID}" + + +@pytest.mark.parametrize("enable_pre_call_checks", [True, False]) +def test_usage_based_routing_v1_rejects_a_deployment_that_reached_its_recorded_rpm_limit( + enable_pre_call_checks: bool, +) -> None: + router: Final = _router_with_recorded_rpm(recorded=RPM_LIMIT, enable_pre_call_checks=enable_pre_call_checks) + with pytest.raises(ValueError, match=RouterErrors.no_deployments_available.value): + router.completion(model=MODEL_GROUP, messages=[{"role": "user", "content": "hello"}]) diff --git a/tests/unit/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py b/tests/unit/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py index 5b204c4155b..87fbf4433a1 100644 --- a/tests/unit/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py +++ b/tests/unit/router_utils/pre_call_checks/test_encrypted_content_affinity_check.py @@ -2059,6 +2059,71 @@ async def test_real_router_selection_keeps_origin_reasoning_and_strips_foreign_o router.discard() +@pytest.mark.asyncio +async def test_affinity_pin_yields_to_the_fallback_hops_target_order_and_strips_the_origins_reasoning(): + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-6-astra", + "litellm_params": { + "model": "openai/gpt-6-astra", + "api_base": "https://api.openai.com/v1", + "api_key": "key-openai", + "order": 1, + }, + "model_info": {"id": "dep-openai"}, + }, + { + "model_name": "gpt-6-astra", + "litellm_params": { + "model": "bedrock_mantle/openai.gpt-6-astra", + "api_base": "https://bedrock-mantle.us-east-1.api.aws", + "api_key": "key-mantle", + "order": 2, + }, + "model_info": {"id": "dep-mantle"}, + }, + ], + optional_pre_call_checks=["encrypted_content_affinity"], + num_retries=0, + ) + openai_wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("blob-openai", "dep-openai") + + def history() -> list: + return [ + {"type": "message", "role": "user", "content": "first question"}, + { + "type": "reasoning", + "encrypted_content": openai_wrapped, + "summary": [{"type": "summary_text", "text": "openai summary"}], + }, + {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "first answer"}]}, + {"type": "message", "role": "user", "content": "second question"}, + ] + + try: + first_attempt = {"input": history(), "store": False} + pinned = await router.async_get_available_deployment( + model="gpt-6-astra", request_kwargs=first_attempt, input=first_attempt["input"] + ) + assert pinned["model_info"]["id"] == "dep-openai" + assert first_attempt["input"] == history() + + hop = {"input": history(), "store": False, "_target_order": 2, "fallback_depth": 1} + hop_deployment = await router.async_get_available_deployment( + model="gpt-6-astra", request_kwargs=hop, input=hop["input"] + ) + assert hop_deployment["model_info"]["id"] == "dep-mantle" + assert hop["input"] == [ + {"type": "message", "role": "user", "content": "first question"}, + {"type": "reasoning", "summary": [{"type": "summary_text", "text": "openai summary"}]}, + {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "first answer"}]}, + {"type": "message", "role": "user", "content": "second question"}, + ] + finally: + router.discard() + + @pytest.mark.asyncio async def test_affinity_keeps_mixed_origins_on_the_same_encryption_boundary(): from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( @@ -2271,6 +2336,56 @@ async def test_affinity_strips_unknown_origins_but_leaves_unmarked_encrypted_con ] +def _router_without_the_origin(): + return litellm.Router( + model_list=[ + { + "model_name": "gpt-6-astra", + "litellm_params": { + "model": "openai/gpt-6-astra", + "api_base": "https://api.openai.com/v1", + "api_key": "openai-key", + }, + "model_info": {"id": "target-order-2"}, + } + ], + num_retries=0, + ) + + +@pytest.mark.parametrize("router", [None, _router_without_the_origin()], ids=["no router", "origin removed"]) +@pytest.mark.parametrize( + "unmarked_origin", ["origin-removed", None], ids=["failed deployment named", "failed deployment unknown"] +) +def test_hop_strip_drops_unmarked_reasoning_whose_origin_cannot_be_resolved(router, unmarked_origin): + from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import ( + EncryptedContentAffinityCheck, + ) + + request_input = [ + {"type": "message", "role": "user", "content": "What is 17*23?"}, + { + "type": "reasoning", + "id": "rs_unmarked", + "encrypted_content": "gAAAAA-minted-by-a-removed-deployment", + "summary": [{"type": "summary_text", "text": "multiply 17 by 23"}], + }, + ] + target = { + "model_info": {"id": "target-order-2"}, + "litellm_params": {"api_base": "https://api.openai.com/v1", "api_key": "openai-key"}, + } + + EncryptedContentAffinityCheck.strip_reasoning_the_targets_cannot_decrypt( + router, request_input, None, (target,), unmarked_origin=unmarked_origin + ) + + assert request_input == [ + {"type": "message", "role": "user", "content": "What is 17*23?"}, + {"type": "reasoning", "summary": [{"type": "summary_text", "text": "multiply 17 by 23"}]}, + ] + + def _cross_group_request_kwargs(): wrapped = ResponsesAPIRequestUtils.wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a") return { diff --git a/tests/unit/router_utils/test_fallback_event_handlers.py b/tests/unit/router_utils/test_fallback_event_handlers.py index adfce2c34be..38c13decaa7 100644 --- a/tests/unit/router_utils/test_fallback_event_handlers.py +++ b/tests/unit/router_utils/test_fallback_event_handlers.py @@ -30,6 +30,7 @@ from litellm.router_utils.fallback_event_handlers import ( get_pre_routing_selection, log_failure_fallback_event, log_success_fallback_event, + mid_stream_fallback_snapshot_kwargs, mid_stream_retry_kwargs, record_pre_routing_selection, record_retry_attempt, @@ -1471,6 +1472,30 @@ def test_get_fallback_model_group_never_resolves_a_provider_without_a_prefixed_k resolver.assert_not_called() +def test_mid_stream_fallback_snapshot_kwargs_restores_the_popped_lists_and_shares_the_buckets(): + controls: Final = MidStreamFallbackControls( + MappingProxyType({"fallbacks": [{"primary": ["backup"]}], "context_window_fallbacks": None}) + ) + metadata: Final = {"model_group": "primary"} + kwargs: Final = {"messages": [{"role": "user", "content": "hi"}], "stream": True, "metadata": metadata} + + snapshot: Final = mid_stream_fallback_snapshot_kwargs(model="primary", controls=controls, kwargs=kwargs) + + assert snapshot == { + **kwargs, + "fallbacks": [{"primary": ["backup"]}], + "context_window_fallbacks": None, + MID_STREAM_FALLBACK_CONTROLS_KEY: controls, + "model": "primary", + } + assert snapshot["metadata"] is metadata + assert "fallbacks" not in kwargs + + bare: Final = mid_stream_fallback_snapshot_kwargs(model="primary", controls=None, kwargs=kwargs) + assert "fallbacks" not in bare + assert bare[MID_STREAM_FALLBACK_CONTROLS_KEY] == MidStreamFallbackControls(MappingProxyType({})) + + def test_mid_stream_retry_kwargs_strips_what_the_retry_wrapper_pops_and_keeps_the_controls_carrier(): def generic_function(**kwargs) -> None: return None diff --git a/tests/unit/rust_bridge/chat_completions/test_route_host.py b/tests/unit/rust_bridge/chat_completions/test_route_host.py index 92eed17de0f..402c032f340 100644 --- a/tests/unit/rust_bridge/chat_completions/test_route_host.py +++ b/tests/unit/rust_bridge/chat_completions/test_route_host.py @@ -4,7 +4,7 @@ from typing import Final import pytest import litellm -from litellm.rust_bridge.chat_completions.route_host import arguments, connection_defaults, response +from litellm.rust_bridge.chat_completions.route_host import connection_defaults, response from litellm.types.utils import ModelResponse @@ -35,12 +35,6 @@ def test_response_builds_the_public_model_response() -> None: assert built.usage.total_tokens == 5 -def test_arguments_are_the_public_kwargs_view() -> None: - kwargs: Final = MappingProxyType({"metadata": {"user_id": "u"}}) - - assert arguments(kwargs) is kwargs - - @pytest.mark.parametrize( ("provider", "global_key", "provider_key", "expected_key", "expected_base"), ( diff --git a/tests/unit/rust_bridge/messages/test_route_host.py b/tests/unit/rust_bridge/messages/test_route_host.py index dde76ee5e82..e721b9624d2 100644 --- a/tests/unit/rust_bridge/messages/test_route_host.py +++ b/tests/unit/rust_bridge/messages/test_route_host.py @@ -1,8 +1,7 @@ from types import MappingProxyType from typing import Final -from litellm.rust_bridge.messages.route_host import arguments, response -from litellm.rust_bridge.public_call import NativeCall +from litellm.rust_bridge.messages.route_host import response import pytest import litellm from litellm.rust_bridge.messages import route_host @@ -29,26 +28,6 @@ def test_response_is_a_detached_public_messages_dict() -> None: assert "_hidden_params" not in native -def test_arguments_preserve_the_bound_view() -> None: - kwargs: Final = MappingProxyType({"litellm_metadata": {"user_id": "u"}}) - request: Final = NativeCall( - args=(), - kwargs=kwargs, - bound={ - "model": "claude-sonnet-4-5", - "messages": [{"role": "user", "content": "hi"}], - "max_tokens": 16, - "stream": None, - "api_key": None, - "api_base": None, - "custom_llm_provider": "anthropic", - **kwargs, - }, - ) - - assert arguments(request.bound) is request.bound - - def test_settings_project_caller_configuration_without_resolving_a_model(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(litellm, "drop_params", False) monkeypatch.setattr(litellm, "reasoning_auto_summary", True) @@ -102,7 +81,7 @@ def test_native_request_rejections_map_to_the_public_400() -> None: request: Final = NativeCall( args=(), kwargs=MappingProxyType({}), - bound={ + base={ "model": "anthropic/claude-sonnet-5", "messages": (), "max_tokens": 8, @@ -116,14 +95,14 @@ def test_native_request_rejections_map_to_the_public_400() -> None: rejected: Final = ValueError("claude-sonnet-5 does not support top_k=5") rejected.messages_request_error = True # pyright: ignore[reportAttributeAccessIssue] # marker the native host sets - mapped: Final = route_host.map_failure(rejected, request.bound, "anthropic") + mapped: Final = route_host.map_failure(rejected, request.resolved, "anthropic") assert isinstance(mapped, litellm.BadRequestError) assert mapped.status_code == 400 assert "does not support top_k=5" in mapped.message assert mapped.model == "claude-sonnet-5" assert not isinstance( - route_host.map_failure(ValueError("plain"), request.bound, "anthropic"), litellm.BadRequestError + route_host.map_failure(ValueError("plain"), request.resolved, "anthropic"), litellm.BadRequestError ) diff --git a/tests/unit/rust_bridge/messages/test_secrets.py b/tests/unit/rust_bridge/messages/test_secrets.py index 53e1376b978..1dab9604740 100644 --- a/tests/unit/rust_bridge/messages/test_secrets.py +++ b/tests/unit/rust_bridge/messages/test_secrets.py @@ -53,7 +53,7 @@ def _native_request() -> NativeCall: return NativeCall( args=(), kwargs=supplied, - bound={ + base={ "model": MESSAGES_MODEL, "messages": MESSAGES, "max_tokens": 8, diff --git a/tests/unit/rust_bridge/native_route_wheel_test.py b/tests/unit/rust_bridge/native_route_wheel_test.py index 9e7523aa29e..bc3269f5dd9 100644 --- a/tests/unit/rust_bridge/native_route_wheel_test.py +++ b/tests/unit/rust_bridge/native_route_wheel_test.py @@ -20,13 +20,6 @@ from typing import Final REQUEST_STARTED: Final = threading.Event() REQUEST_CANCELLED: Final = threading.Event() -ANTHROPIC_RESPONSE: Final = ( - b'{"id":"msg_native","type":"message","role":"assistant",' - b'"model":"claude-sonnet-4-5","content":[{"type":"text","text":"native-message"}],' - b'"stop_reason":"end_turn","stop_sequence":null,' - b'"usage":{"input_tokens":2,"output_tokens":3}}' -) - class NativeRouteServer(ThreadingHTTPServer): request_queue_size = 64 @@ -49,7 +42,7 @@ class NativeRouteHandler(BaseHTTPRequestHandler): return status: Final = 429 if outcome == "429" else 200 - response_body: Final = native_response(status, route) + response_body: Final = native_response(status) self.send_response(status) self.send_header("content-type", "application/json") @@ -78,32 +71,23 @@ def assert_native_request( headers: HTTPMessage, body: object, ) -> None: - if route not in {"transcription", "chat_completions"}: + if route != "transcription": raise AssertionError(f"unexpected route marker: {route!r}") if outcome not in {"success", "429", "hang"}: raise AssertionError(f"unexpected outcome marker: {outcome!r}") if not isinstance(body, dict): raise TypeError(f"{route} sent {type(body).__name__}, expected a JSON object") - if route == "transcription": - assert path == "/model/mistral.voxtral-mini-3b-2507/converse" - assert headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ") - assert headers.get("x-amz-date") - assert body["messages"][0]["content"][0]["audio"]["source"]["bytes"] == "AQI=" - assert "The audio language is en" in body["messages"][0]["content"][1]["text"] - return - assert path == "/v1/messages" - assert headers.get("x-api-key") == "sk-native" - assert body["model"] == "claude-sonnet-4-5" - assert body["max_tokens"] == 17 - assert body["messages"][0]["content"] == [{"type": "text", "text": "hello-from-chat"}] + assert path == "/model/mistral.voxtral-mini-3b-2507/converse" + assert headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ") + assert headers.get("x-amz-date") + assert body["messages"][0]["content"][0]["audio"]["source"]["bytes"] == "AQI=" + assert "The audio language is en" in body["messages"][0]["content"][1]["text"] -def native_response(status: int, route: str | None) -> bytes: +def native_response(status: int) -> bytes: if status == 429: return b'{"error":"native-rate-limit"}' - if route == "transcription": - return b'{"output":{"message":{"content":[{"text":"native-transcription"}]}}}' - return ANTHROPIC_RESPONSE + return b'{"output":{"message":{"content":[{"text":"native-transcription"}]}}}' def load_native(native_path: Path) -> object: @@ -117,7 +101,7 @@ def load_native(native_path: Path) -> object: def route_call(route: str, api_base: str, outcome: str) -> SimpleNamespace: fields: Final = route_kwargs(route, api_base, outcome) - return SimpleNamespace(args=(), kwargs=fields, bound=fields) + return SimpleNamespace(args=(), kwargs=fields, base={}) def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]: @@ -138,38 +122,23 @@ def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]: "language": "en", }, } - if route == "chat_completions": - return common | { - "model": "anthropic/claude-sonnet-4-5", - "messages": [{"role": "user", "content": "hello-from-chat"}], - "optional_params": {"max_tokens": 17}, - "api_key": "sk-native", - } raise AssertionError(f"unknown route: {route}") def assert_success(route: str, response: object) -> None: if not isinstance(response, dict): raise TypeError(f"{route} returned {type(response).__name__}, expected dict") - actual: Final = success_value(route, response) - expected: Final = "native-transcription" if route == "transcription" else "native-message" - if actual != expected: - raise AssertionError(f"{route} returned {actual!r}, expected {expected!r}") - - -def success_value(route: str, response: dict[object, object]) -> object: - if route == "transcription": - return response["text"] - return response["choices"][0]["message"]["content"] + if response["text"] != "native-transcription": + raise AssertionError(f"{route} returned {response['text']!r}, expected 'native-transcription'") def assert_rate_limit(route: str, error: BaseException) -> None: - if error.args != (429, native_response(429, route).decode()): + if error.args != (429, native_response(429).decode()): raise AssertionError(f"{route} returned the wrong 429 error: {error!r}") def exercise_sync(native: object, api_base: str) -> None: - for route in ("transcription", "chat_completions"): + for route in ("transcription",): function: Final = getattr(native, route) assert_success(route, function(route_call(route, api_base, "success"))) try: @@ -181,7 +150,7 @@ def exercise_sync(native: object, api_base: str) -> None: async def exercise_async(native: object, api_base: str) -> None: - for route in ("transcription", "chat_completions"): + for route in ("transcription",): function: Final = getattr(native, f"a{route}") assert_success(route, await function(route_call(route, api_base, "success"))) try: @@ -192,13 +161,11 @@ async def exercise_async(native: object, api_base: str) -> None: raise AssertionError(f"a{route} accepted a 429 response") responses: Final = await asyncio.wait_for( - asyncio.gather( - *(native.achat_completions(route_call("chat_completions", api_base, "success")) for _ in range(32)) - ), + asyncio.gather(*(native.atranscription(route_call("transcription", api_base, "success")) for _ in range(32))), timeout=15, ) for response in responses: - assert_success("chat_completions", response) + assert_success("transcription", response) def exercise_routes(native_path: Path, api_base: str) -> object: @@ -210,8 +177,8 @@ def exercise_routes(native_path: Path, api_base: str) -> object: def exercise_signal(native: object, api_base: str) -> int: try: - native.chat_completions( - route_call("chat_completions", api_base, "hang"), + native.transcription( + route_call("transcription", api_base, "hang"), ) except KeyboardInterrupt: sys.stdout.write("KeyboardInterrupt\n") diff --git a/tests/unit/rust_bridge/ocr/test_route_host.py b/tests/unit/rust_bridge/ocr/test_route_host.py index 0b8af515a64..879edb951fc 100644 --- a/tests/unit/rust_bridge/ocr/test_route_host.py +++ b/tests/unit/rust_bridge/ocr/test_route_host.py @@ -10,7 +10,7 @@ from litellm.rust_bridge.public_call import NativeCall REQUEST: Final = NativeCall( args=(), kwargs={"req_format": "markdown"}, - bound={ + base={ "model": "mistral/mistral-ocr-latest", "document": {"type": "document_url", "document_url": "https://example.com/file.pdf"}, "api_key": "test-key", @@ -18,7 +18,6 @@ REQUEST: Final = NativeCall( "timeout": None, "custom_llm_provider": None, "extra_headers": None, - **{"req_format": "markdown"}, }, ) @@ -53,7 +52,7 @@ def test_rust_ocr_response_retains_provider_native_response(): def test_map_failure_builds_public_error_from_upstream_status_and_headers() -> None: error: Final = RustUpstreamError(429, '{"message": "slow down"}', (("retry-after", "7"),)) - public_error: Final = map_failure(error, REQUEST.bound, "mistral") + public_error: Final = map_failure(error, REQUEST.resolved, "mistral") assert isinstance(public_error, litellm.RateLimitError) assert public_error.status_code == 429 @@ -66,7 +65,7 @@ def test_map_failure_builds_public_error_from_upstream_status_and_headers() -> N def test_map_failure_maps_upstream_401_to_authentication_error() -> None: error: Final = RustUpstreamError(401, '{"message": "Unauthorized"}', ()) - public_error: Final = map_failure(error, REQUEST.bound, "mistral") + public_error: Final = map_failure(error, REQUEST.resolved, "mistral") assert isinstance(public_error, litellm.AuthenticationError) assert public_error.status_code == 401 @@ -77,7 +76,7 @@ def test_map_failure_maps_upstream_401_to_authentication_error() -> None: def test_map_failure_leaves_non_upstream_errors_unwrapped() -> None: error: Final = RuntimeError("bridge exploded") - public_error: Final = map_failure(error, REQUEST.bound, "mistral") + public_error: Final = map_failure(error, REQUEST.resolved, "mistral") assert not isinstance(public_error, UpstreamFailure) assert isinstance(public_error, litellm.APIConnectionError) @@ -86,4 +85,4 @@ def test_map_failure_leaves_non_upstream_errors_unwrapped() -> None: def test_map_failure_reports_invalid_request_format_as_unsupported_params() -> None: with pytest.raises(litellm.UnsupportedParamsError, match="Invalid `req_format`: 'markdown'"): - raise map_failure(RustFormatError(), REQUEST.bound, "mistral") + raise map_failure(RustFormatError(), REQUEST.resolved, "mistral") diff --git a/tests/unit/rust_bridge/ocr/test_secrets.py b/tests/unit/rust_bridge/ocr/test_secrets.py index 5d051fbffe0..262fba1429b 100644 --- a/tests/unit/rust_bridge/ocr/test_secrets.py +++ b/tests/unit/rust_bridge/ocr/test_secrets.py @@ -64,7 +64,7 @@ def _native_request(api_base: str) -> NativeCall: return NativeCall( args=(), kwargs=supplied, - bound={ + base={ "model": OCR_MODEL, "document": OCR_DOCUMENT, "api_key": None, diff --git a/tests/unit/rust_bridge/responses/test_route_host.py b/tests/unit/rust_bridge/responses/test_route_host.py index 1d67f2c368c..e110f16350b 100644 --- a/tests/unit/rust_bridge/responses/test_route_host.py +++ b/tests/unit/rust_bridge/responses/test_route_host.py @@ -5,8 +5,7 @@ import pytest from pydantic import ValidationError import litellm -from litellm.rust_bridge.responses.route_host import arguments, connection_defaults, map_failure, response -from litellm.rust_bridge.public_call import NativeCall +from litellm.rust_bridge.responses.route_host import connection_defaults, map_failure, response from litellm.types.llms.openai import ResponsesAPIResponse @@ -42,26 +41,6 @@ def test_response_rejects_a_payload_missing_required_fields() -> None: response(MappingProxyType({"object": "response"})) -def test_arguments_preserve_the_bound_view() -> None: - kwargs: Final = MappingProxyType({"litellm_metadata": {"user_id": "u"}}) - request: Final = NativeCall( - args=(), - kwargs=kwargs, - bound={ - "model": "gpt-4o", - "input": "hi", - "stream": None, - "api_key": None, - "api_base": None, - "custom_llm_provider": "openai", - "extra_headers": None, - **kwargs, - }, - ) - - assert arguments(request.bound) is request.bound - - @pytest.mark.parametrize( ("global_key", "provider_key", "expected"), ( diff --git a/tests/unit/rust_bridge/test_public_call.py b/tests/unit/rust_bridge/test_public_call.py index 24d5db057d6..63029e288e2 100644 --- a/tests/unit/rust_bridge/test_public_call.py +++ b/tests/unit/rust_bridge/test_public_call.py @@ -3,7 +3,7 @@ from typing import Final import pytest -from litellm.rust_bridge.public_call import bind, native_call, signature +from litellm.rust_bridge.public_call import native_call, signature def _messages( @@ -17,45 +17,48 @@ def _messages( return None +_SIGNATURE: Final = signature(_messages) + + @pytest.mark.parametrize("supplied", ({}, {"api_key": None}, {"api_key": "explicit"})) -def test_native_call_preserves_omission_separately_from_bound_defaults(supplied: Mapping[str, object]) -> None: +def test_base_holds_positionals_and_defaults_and_never_a_keyword(supplied: Mapping[str, object]) -> None: messages: Final[Sequence[object]] = [{"role": "user", "content": "hello"}] args: Final = (128, messages, "model", 0.25) - fields: Final = bind(signature(_messages), args, supplied) - assert fields is not None - call: Final = native_call(args, supplied, fields) + call: Final = native_call(_SIGNATURE, args, supplied) assert call.args is args assert call.kwargs is supplied - assert call.bound == { + assert call.base == { "max_tokens": 128, "messages": messages, "model": "model", "temperature": 0.25, - "api_key": supplied.get("api_key"), + "api_key": None, } - assert call.bound["messages"] is messages + assert call.base["messages"] is messages + assert call.resolved["api_key"] == supplied.get("api_key") assert ("api_key" in call.kwargs) == ("api_key" in supplied) -def test_native_call_keeps_extra_option_objects_without_nested_kwargs() -> None: +def test_resolved_lays_the_keywords_over_the_base_and_equals_the_full_binding() -> None: messages: Final[Sequence[object]] = [] metadata: Final = {"trace": "caller"} - supplied: Final = {"metadata": metadata} - args: Final = (128, messages, "model") - fields: Final = bind(signature(_messages), args, supplied) - assert fields is not None + supplied: Final = {"model": "keyword-model", "metadata": metadata} + args: Final = (128, messages) - call: Final = native_call(args, supplied, fields) + call: Final = native_call(_SIGNATURE, args, supplied) - assert call.bound == { + assert "model" not in call.base + assert "kwargs" not in call.base + assert "metadata" not in call.base + assert call.resolved == { "max_tokens": 128, "messages": messages, - "model": "model", + "model": "keyword-model", "temperature": None, "api_key": None, "metadata": metadata, } - assert call.bound["metadata"] is metadata - assert supplied == {"metadata": metadata} + assert call.resolved["metadata"] is metadata + assert supplied == {"model": "keyword-model", "metadata": metadata} diff --git a/tests/unit/rust_bridge/test_tokenizer.py b/tests/unit/rust_bridge/test_tokenizer.py index c5093cdb0ce..42a31a911fb 100644 --- a/tests/unit/rust_bridge/test_tokenizer.py +++ b/tests/unit/rust_bridge/test_tokenizer.py @@ -1,4 +1,5 @@ from typing import Final +from unittest.mock import patch import pytest import tiktoken @@ -47,3 +48,24 @@ def test_native_custom_tokenizer_matches_python() -> None: assert native.encode("Hello World").ids == reference.encode("Hello World").ids assert native.decode(reference.encode("Hello World").ids) == reference.decode(reference.encode("Hello World").ids) + + +def test_python_tokenizer_missing_dependency_is_actionable() -> None: + with patch.dict("sys.modules", {"tokenizers": None}): + with pytest.raises(ImportError, match="pip install tokenizers") as error: + tokenizer._python_tokenizer() + assert isinstance(error.value.__cause__, ModuleNotFoundError) + assert error.value.__cause__.name == "tokenizers" + + +def test_python_tokenizer_factory_preserves_installed_interface() -> None: + result: Final = tokenizer._python_tokenizer().from_str(TOKENIZER_JSON) + assert result.encode("Hello World").ids == Tokenizer.from_str(TOKENIZER_JSON).encode("Hello World").ids + + +def test_python_tokenizer_preserves_unrelated_import_failure() -> None: + failure: Final = ModuleNotFoundError("broken tokenizer installation", name="unrelated_dependency") + with patch("builtins.__import__", side_effect=failure): + with pytest.raises(ModuleNotFoundError) as error: + tokenizer._python_tokenizer() + assert error.value is failure diff --git a/tests/unit/secret_managers/test_aws_secret_manager_v2.py b/tests/unit/secret_managers/test_aws_secret_manager_v2.py index 1972782bc01..ab358e0bf8b 100644 --- a/tests/unit/secret_managers/test_aws_secret_manager_v2.py +++ b/tests/unit/secret_managers/test_aws_secret_manager_v2.py @@ -664,3 +664,16 @@ async def test_end_to_end_iam_role_secret_write(): print("Delete Response:", delete_response) except Exception as e: print(f"Cleanup failed: {e}") + + +def test_missing_botocore_keeps_dependency_identity(): + from unittest.mock import patch + + import pytest + + from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2 + + with patch.dict("sys.modules", {"botocore": None}): + with pytest.raises(ModuleNotFoundError, match="pip install boto3") as caught: + AWSSecretsManagerV2()._prepare_request(action="GetSecretValue", secret_name="test-secret") + assert caught.value.name == "botocore" diff --git a/tests/unit/secret_managers/test_google_secret_manager.py b/tests/unit/secret_managers/test_google_secret_manager.py new file mode 100644 index 00000000000..ed93aa2861d --- /dev/null +++ b/tests/unit/secret_managers/test_google_secret_manager.py @@ -0,0 +1,83 @@ +import base64 +from dataclasses import dataclass +from typing import Final + +import litellm +import pytest +import respx + +from litellm.secret_managers.google_secret_manager import GoogleSecretManager + + +@dataclass(frozen=True, slots=True) +class _CachedVertexCredentials: + token: str + quota_project_id: str | None + expired: bool = False + + def refresh(self, request: object) -> None: + return None + + +def _google_secret_manager( + monkeypatch: pytest.MonkeyPatch, + project_id: str, +) -> GoogleSecretManager: + monkeypatch.setenv("GOOGLE_SECRET_MANAGER_PROJECT_ID", project_id) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + secret_manager: Final = GoogleSecretManager() + credentials: Final = _CachedVertexCredentials( + token="test-gsm-token", + quota_project_id=project_id, + ) + vertex_chat_completion: Final = litellm.vertex_chat_completion + monkeypatch.setitem( + vertex_chat_completion._credentials_project_mapping, + (None, project_id), + (credentials, project_id), + ) + return secret_manager + + +def test_google_secret_manager_decodes_secret_and_requests_latest_version( + monkeypatch: pytest.MonkeyPatch, +) -> None: + project_id: Final = "test-secret-project" + secret_manager: Final = _google_secret_manager(monkeypatch, project_id) + secret_url: Final = ( + f"https://secretmanager.googleapis.com/v1/projects/{project_id}/secrets/OPENAI_API_KEY/versions/latest:access" + ) + encoded_secret: Final = base64.b64encode(b"anything").decode("ascii") + upstream: Final[respx.MockRouter] + + with respx.mock(assert_all_called=True) as upstream: + secret_route: Final = upstream.get(secret_url).respond( + 200, + json={"payload": {"data": encoded_secret}}, + ) + + result: Final = secret_manager.get_secret_from_google_secret_manager("OPENAI_API_KEY") + + assert result == "anything" + assert secret_route.called + assert len(upstream.calls) == 1 + assert str(upstream.calls.last.request.url) == secret_url + assert upstream.calls.last.request.headers["Authorization"] == "Bearer test-gsm-token" + + +def test_google_secret_manager_returns_cached_values_without_http( + monkeypatch: pytest.MonkeyPatch, +) -> None: + project_id: Final = "test-secret-project" + secret_manager: Final = _google_secret_manager(monkeypatch, project_id) + secret_manager.cache.set_cache("cached-none", None) + secret_manager.cache.set_cache("cached-string", "lite-llm") + upstream: Final[respx.MockRouter] + + with respx.mock() as upstream: + missing_value: Final = secret_manager.get_secret_from_google_secret_manager("cached-none") + cached_value: Final = secret_manager.get_secret_from_google_secret_manager("cached-string") + + assert missing_value is None + assert cached_value == "lite-llm" + assert upstream.calls.call_count == 0 diff --git a/tests/unit/secret_managers/test_secret_managers_main.py b/tests/unit/secret_managers/test_secret_managers_main.py index 32040251795..d32555c3669 100644 --- a/tests/unit/secret_managers/test_secret_managers_main.py +++ b/tests/unit/secret_managers/test_secret_managers_main.py @@ -3,6 +3,7 @@ import json import logging import os import time +from typing import Final from unittest.mock import Mock, patch import pytest @@ -230,6 +231,13 @@ def test_oidc_circleci_success(monkeypatch): assert result == "circleci_token" +def test_oidc_circleci_v2_returns_the_environment_token(monkeypatch: pytest.MonkeyPatch) -> None: + token: Final = "circleci-v2-token" + monkeypatch.setenv("CIRCLE_OIDC_TOKEN_V2", token) + + assert get_secret("oidc/circleci_v2/test-audience") == token + + def test_oidc_circleci_failure(monkeypatch): monkeypatch.delenv("CIRCLE_OIDC_TOKEN", raising=False) secret_name = "oidc/circleci/test-audience" diff --git a/tests/unit/test_anthropic_beta_headers_filtering.py b/tests/unit/test_anthropic_beta_headers_filtering.py index 1a6899f16ba..2a85122d91c 100644 --- a/tests/unit/test_anthropic_beta_headers_filtering.py +++ b/tests/unit/test_anthropic_beta_headers_filtering.py @@ -10,7 +10,7 @@ This test validates: import json import os -from typing import Dict, List +from typing import Dict, Final, List from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -458,6 +458,20 @@ class TestAnthropicBetaHeadersFiltering: assert filtered == ["dangerous-tool-use-2026-09-03"] + @pytest.mark.parametrize("provider", ["anthropic", "vertex_ai"]) + def test_inline_tools_forwarded(self, provider: str) -> None: + """inline-tools-2026-09-15 lets a `tool_addition` system message carry a new tool mid-conversation. + Anthropic documents it for the Claude API + (https://platform.claude.com/docs/en/build-with-claude/mid-conversation-system-messages, fetched 2026-10-07), + and a live Vertex rawPredict call on claude-opus-5-5 honored it on 2026-10-07 while answering 400 on the + `tool_addition` block without it.""" + filtered: Final = filter_and_transform_beta_headers( + beta_headers=["inline-tools-2026-09-15"], + provider=provider, + ) + + assert filtered == ["inline-tools-2026-09-15"] + def test_null_value_headers_filtered(self): """Test that headers with null values are always filtered out.""" for provider in [ diff --git a/tests/unit/test_audio_transcription_rust_bridge.py b/tests/unit/test_audio_transcription_rust_bridge.py index 31ad5b72cc9..e3a6781dee7 100644 --- a/tests/unit/test_audio_transcription_rust_bridge.py +++ b/tests/unit/test_audio_transcription_rust_bridge.py @@ -99,7 +99,7 @@ def test_dispatch_marshals_audio_into_rust_call() -> None: "timeout_seconds": 5.0, } assert response.text == "rust" - assert bridge.calls == (NativeCall(args=(), kwargs=expected, bound=expected),) + assert bridge.calls == (NativeCall(args=(), kwargs=expected, base={}),) @pytest.mark.parametrize("disable", ("process", "environment")) @@ -164,7 +164,7 @@ async def test_async_dispatch_marshals_audio_into_rust_call() -> None: "timeout_seconds": 5.0, } assert response.text == "async rust" - assert bridge.calls == (NativeCall(args=(), kwargs=expected, bound=expected),) + assert bridge.calls == (NativeCall(args=(), kwargs=expected, base={}),) def test_bedrock_transcription_dispatches_to_rust_from_sdk_entrypoint() -> None: @@ -175,8 +175,8 @@ def test_bedrock_transcription_dispatches_to_rust_from_sdk_entrypoint() -> None: assert isinstance(response, litellm.TranscriptionResponse) assert response.text == "rust" - assert bridge.calls[0].bound["model"] == MODEL.removeprefix("bedrock/") - assert frozenset(bridge.calls[0].bound) == TRANSCRIPTION_FIELDS + assert bridge.calls[0].resolved["model"] == MODEL.removeprefix("bedrock/") + assert frozenset(bridge.calls[0].resolved) == TRANSCRIPTION_FIELDS @pytest.mark.asyncio @@ -187,8 +187,8 @@ async def test_bedrock_atranscription_dispatches_to_rust_from_sdk_entrypoint() - response: Final = await litellm.atranscription(model=MODEL, file=AUDIO_FILE) assert response.text == "async rust" - assert tuple(call.bound["model"] for call in bridge.calls) == (MODEL.removeprefix("bedrock/"),) - assert frozenset(bridge.calls[0].bound) == TRANSCRIPTION_FIELDS + assert tuple(call.resolved["model"] for call in bridge.calls) == (MODEL.removeprefix("bedrock/"),) + assert frozenset(bridge.calls[0].resolved) == TRANSCRIPTION_FIELDS @pytest.mark.asyncio diff --git a/tests/unit/test_budget_ratchet_check.py b/tests/unit/test_budget_ratchet_check.py deleted file mode 100644 index b359b8b42e5..00000000000 --- a/tests/unit/test_budget_ratchet_check.py +++ /dev/null @@ -1,189 +0,0 @@ -"""Tests for scripts/budget_ratchet_check.py. - -The guard's contract is "limits may only fall": a raised limit, a dropped rule, or -a deleted file is a regression, while a lowered/equal limit, a brand-new rule, or a -brand-new budget file is fine. Each branch is pinned here. -""" - -import importlib.util -import subprocess -import sys -from pathlib import Path -from typing import Final - -_MODULE_PATH = ( - Path(__file__).resolve().parents[2] / "scripts" / "budget_ratchet_check.py" -) -_spec = importlib.util.spec_from_file_location("budget_ratchet_check", _MODULE_PATH) -ratchet = importlib.util.module_from_spec(_spec) -_spec.loader.exec_module(ratchet) - - -def _spec_of(limit): - return {"limit": limit} - - -def test_limits_read_the_limit_and_skip_malformed(): - limits = ratchet._limits({"LIT006": _spec_of(1023), "junk": 5}) - assert limits == {"LIT006": 1023} # malformed (non-dict) spec ignored - - -def test_limits_fall_back_to_legacy_baseline_plus_slack(): - # The base side of a diff can predate the `limit` migration; its ceiling is - # baseline + slack, read on the same footing as a new-schema `limit`. - assert ratchet._limits({"LIT006": {"baseline": 1013, "slack": 10}}) == {"LIT006": 1023} - - -def test_migration_from_legacy_schema_to_equal_limit_is_clean(): - # baseline+slack (1023) -> limit 1023 is the same ceiling, so no regression. - base = {"LIT006": {"baseline": 1013, "slack": 10}} - assert ratchet.regressions_for("b.json", base, {"LIT006": _spec_of(1023)}) == [] - # ...and a genuine raise across the migration is still caught. - regs = ratchet.regressions_for("b.json", base, {"LIT006": _spec_of(1024)}) - assert [r.rule for r in regs] == ["LIT006"] and "1023 -> 1024" in regs[0].detail - - -def test_raised_limit_is_a_regression(): - base = {"LIT006": _spec_of(1023)} - head = {"LIT006": _spec_of(1024)} - regs = ratchet.regressions_for("b.json", base, head) - assert [r.rule for r in regs] == ["LIT006"] - assert "1023 -> 1024" in regs[0].detail - - -def test_lowered_or_equal_limit_is_clean(): - base = {"LIT006": _spec_of(1023)} - # limit drops - assert ratchet.regressions_for("b.json", base, {"LIT006": _spec_of(1000)}) == [] - # nothing changes - assert ratchet.regressions_for("b.json", base, {"LIT006": _spec_of(1023)}) == [] - - -def test_dropped_rule_is_a_regression(): - regs = ratchet.regressions_for("b.json", {"LIT007": _spec_of(0)}, {}) - assert [r.rule for r in regs] == ["LIT007"] - assert "dropped" in regs[0].detail - - -def test_new_rule_in_head_is_clean(): - assert ratchet.regressions_for("b.json", {}, {"new-rule": _spec_of(5)}) == [] - - -def test_dropped_rule_that_graduated_to_a_hard_failing_config_is_clean(): - base = {"UP006": _spec_of(0)} - assert ratchet.regressions_for("b.json", base, {}, graduated=("UP006",)) == [] - - -def test_graduation_matches_by_prefix_like_ruff_selectors_do(): - base = {"ANN202": _spec_of(865)} - assert ratchet.regressions_for("b.json", base, {}, graduated=("ANN",)) == [] - - -def test_an_unrelated_graduation_does_not_excuse_a_dropped_rule(): - base = {"C901": _spec_of(3)} - regs = ratchet.regressions_for("b.json", base, {}, graduated=("UP006", "SIM118")) - assert [r.rule for r in regs] == ["C901"] - assert "dropped" in regs[0].detail - - -def test_graduation_never_excuses_a_raised_limit(): - base = {"UP006": _spec_of(0)} - regs = ratchet.regressions_for("b.json", base, {"UP006": _spec_of(7)}, graduated=("UP006",)) - assert [r.rule for r in regs] == ["UP006"] - assert "0 -> 7" in regs[0].detail - - -def test_dropped_rule_the_checker_retired_is_clean(): - base: Final = {"TQ008": _spec_of(10993)} - assert ratchet.regressions_for("b.json", base, {}, retired=frozenset({"TQ008"})) == [] - - -def test_dropped_rule_the_checker_still_emits_is_a_regression(): - base: Final = {"TQ001": _spec_of(5), "TQ008": _spec_of(10993)} - regs: Final = ratchet.regressions_for("b.json", base, {}, retired=frozenset({"TQ008"})) - assert [r.rule for r in regs] == ["TQ001"] - assert "dropped" in regs[0].detail - - -def test_retirement_never_excuses_a_raised_limit(): - base: Final = {"TQ008": _spec_of(0)} - regs: Final = ratchet.regressions_for("b.json", base, {"TQ008": _spec_of(7)}, retired=frozenset({"TQ008"})) - assert [r.rule for r in regs] == ["TQ008"] - assert "0 -> 7" in regs[0].detail - - -def test_retired_rules_come_from_the_paired_checker(): - base: Final = {"TQ001": _spec_of(5), "TQ008": _spec_of(10993)} - assert ratchet.retired_rules("test-quality-budget.json", base) == frozenset({"TQ008"}) - - -def test_budgets_without_a_paired_checker_never_retire(): - base: Final = {"TQ008": _spec_of(1)} - for rel in ("ruff-strict-budget.json", "type-discipline-budget.json", "basedpyright-code-budget.json"): - assert ratchet.retired_rules(rel, base) == frozenset() - - -def test_graduated_selectors_come_from_the_paired_ruff_config(): - selectors = ratchet.graduated_selectors("ruff-strict-budget.json") - assert "UP006" in selectors - assert "ANN" not in selectors - - -def test_budgets_without_a_paired_config_can_never_graduate(): - assert ratchet.graduated_selectors("type-discipline-budget.json") == () - assert ratchet.graduated_selectors("basedpyright-code-budget.json") == () - - -def test_a_selector_the_config_also_ignores_does_not_count_as_graduated(): - lint = {"ignore": ["UP006"], "extend-select": ["UP006", "SIM118"]} - assert ratchet.selectors_hard_failed_by(lint) == ("SIM118",) - - -def test_selectors_hard_failed_by_reads_a_config_with_no_ignore_list(): - assert ratchet.selectors_hard_failed_by({"extend-select": ["UP006"]}) == ("UP006",) - - -def test_deleted_budget_file_is_a_regression(): - regs = ratchet.regressions_for("b.json", {"LIT006": _spec_of(1)}, None) - assert [r.rule for r in regs] == ["*"] - assert "deleted" in regs[0].detail - - -def test_new_budget_file_has_nothing_to_ratchet(): - assert ratchet.regressions_for("b.json", None, {"LIT006": _spec_of(1)}) == [] - - -def test_default_budgets_watch_every_budget_file_in_the_repo(): - # This job is the repo's only ceiling-raise alarm, so every *-budget.json on disk must be - # watched; a budget left out of DEFAULT_BUDGETS (e.g. basedpyright-code-budget.json) can be - # loosened with no signal. Equality also catches a phantom entry that no longer exists. - repo_root = _MODULE_PATH.parents[1] - on_disk = frozenset(p.name for p in repo_root.glob("*budget*.json")) - assert on_disk == frozenset(ratchet.DEFAULT_BUDGETS) - - -# --------------------------------------------------------------------------- # -# Base-ref resolution: a bad ref must fail loudly, never pass vacuously -# --------------------------------------------------------------------------- # - - -def test_ref_is_commit_distinguishes_real_from_bogus(): - assert ratchet._ref_is_commit("HEAD") is True - assert ratchet._ref_is_commit("definitely-not-a-real-ref-zzz") is False - - -def test_load_base_reads_a_present_file_and_none_for_an_absent_one(): - # A real budget file exists at HEAD; a made-up path is absent at the same (valid) ref. - assert ratchet._load_base("type-discipline-budget.json", "HEAD") is not None - assert ratchet._load_base("scripts/no-such-budget-xyz.json", "HEAD") is None - - -def test_unresolvable_base_ref_exits_nonzero_instead_of_skipping(): - proc = subprocess.run( - [sys.executable, str(_MODULE_PATH), "--base", "definitely-not-a-real-ref-zzz"], - cwd=_MODULE_PATH.parents[1], - capture_output=True, - text=True, - ) - assert proc.returncode == 1 - assert "does not resolve to a commit" in proc.stderr diff --git a/tests/unit/test_check_type_discipline.py b/tests/unit/test_check_type_discipline.py index ee598ff353d..378d46a3f3c 100644 --- a/tests/unit/test_check_type_discipline.py +++ b/tests/unit/test_check_type_discipline.py @@ -7,12 +7,11 @@ a test fail. The comment-scanner cases are the regression for the readline path: """ import importlib.util -import json import os -import re import subprocess import sys from pathlib import Path +from typing import Final import pytest @@ -599,6 +598,277 @@ def test_writable_ok_without_reason_is_lit005_and_does_not_suppress(tmp_path): assert "LIT012" in codes +# --------------------------------------------------------------------------- # +# Unfrozen pydantic models (LIT015) +# --------------------------------------------------------------------------- # + + +def test_unfrozen_basemodel_is_flagged(tmp_path): + src = "from pydantic import BaseModel\nclass P(BaseModel):\n a: int\n" + assert "LIT015" in _codes(tmp_path, src) + + +def test_configdict_frozen_true_is_clean(tmp_path): + src = ( + "from pydantic import BaseModel, ConfigDict\n" + "class P(BaseModel):\n" + " model_config = ConfigDict(extra='allow', frozen=True)\n" + ) + assert "LIT015" not in _codes(tmp_path, src) + + +def test_shared_frozen_configdict_constant_is_clean(tmp_path): + src: Final = ( + "from typing import Final\n" + "from pydantic import BaseModel, ConfigDict\n" + "_RESPONSE_CONFIG: Final = ConfigDict(frozen=True)\n" + "class Foo(BaseModel):\n" + " model_config = _RESPONSE_CONFIG\n" + " x: int\n" + ) + assert "LIT015" not in _codes(tmp_path, src) + + +def test_mutable_config_before_class_is_not_hidden_by_frozen_reassignment(tmp_path): + src: Final = ( + "from pydantic import BaseModel, ConfigDict\n" + "_RESPONSE_CONFIG = ConfigDict(frozen=False)\n" + "class Foo(BaseModel):\n" + " model_config = _RESPONSE_CONFIG\n" + " x: int\n" + "_RESPONSE_CONFIG = ConfigDict(frozen=True)\n" + ) + assert "LIT015" in _codes(tmp_path, src) + + +def test_frozen_config_before_class_survives_mutable_reassignment(tmp_path): + src: Final = ( + "from pydantic import BaseModel, ConfigDict\n" + "_RESPONSE_CONFIG = ConfigDict(frozen=True)\n" + "class Foo(BaseModel):\n" + " model_config = _RESPONSE_CONFIG\n" + " x: int\n" + "_RESPONSE_CONFIG = ConfigDict(frozen=False)\n" + ) + assert "LIT015" not in _codes(tmp_path, src) + + +def test_config_assigned_only_after_class_is_unresolved(tmp_path): + src: Final = ( + "from pydantic import BaseModel, ConfigDict\n" + "class Foo(BaseModel):\n" + " model_config = _RESPONSE_CONFIG\n" + " x: int\n" + "_RESPONSE_CONFIG = ConfigDict(frozen=True)\n" + ) + assert "LIT015" in _codes(tmp_path, src) + + +def test_unknown_model_config_name_is_flagged(tmp_path): + src: Final = ( + "from pydantic import BaseModel\nclass Foo(BaseModel):\n model_config = UNKNOWN_CONFIG\n x: int\n" + ) + assert "LIT015" in _codes(tmp_path, src) + + +def test_dict_literal_model_config_frozen_true_is_clean(tmp_path): + src = "from pydantic import BaseModel\nclass P(BaseModel):\n model_config = {'frozen': True, 'extra': 'allow'}\n" + assert "LIT015" not in _codes(tmp_path, src) + + +def test_subclass_of_in_file_frozen_model_is_clean(tmp_path): + src = ( + "from pydantic import BaseModel, ConfigDict\n" + "class Base(BaseModel):\n" + " model_config = ConfigDict(frozen=True)\n" + "class Child(Base):\n" + " a: int\n" + ) + assert "LIT015" not in _codes(tmp_path, src) + + +def test_subclass_of_in_file_unfrozen_model_flags_both(tmp_path): + src = "from pydantic import BaseModel\nclass Base(BaseModel):\n pass\nclass Child(Base):\n a: int\n" + assert _codes(tmp_path, src).count("LIT015") == 2 + + +def test_litellm_pydantic_object_base_without_frozen_is_flagged(tmp_path): + src = "class P(LiteLLMPydanticObjectBase):\n a: int\n" + assert "LIT015" in _codes(tmp_path, src) + + +def test_inner_config_class_frozen_true_is_clean(tmp_path): + src = "from pydantic import BaseModel\nclass P(BaseModel):\n class Config:\n frozen = True\n" + assert "LIT015" not in _codes(tmp_path, src) + + +def test_frozen_false_is_flagged(tmp_path): + src = ( + "from pydantic import BaseModel, ConfigDict\nclass P(BaseModel):\n model_config = ConfigDict(frozen=False)\n" + ) + assert "LIT015" in _codes(tmp_path, src) + + +def test_later_model_config_frozen_false_overrides_earlier_frozen_true(tmp_path): + src = ( + "from pydantic import BaseModel, ConfigDict\n" + "class P(BaseModel):\n" + " model_config = ConfigDict(frozen=True)\n" + " model_config = ConfigDict(frozen=False)\n" + ) + assert "LIT015" in _codes(tmp_path, src) + + +def test_subclass_frozen_false_overrides_frozen_parent(tmp_path): + src = ( + "from pydantic import BaseModel, ConfigDict\n" + "class Base(BaseModel):\n" + " model_config = ConfigDict(frozen=True)\n" + "class Writable(Base):\n" + " model_config = ConfigDict(frozen=False)\n" + "class StillFrozen(Base):\n" + " model_config = ConfigDict(extra='allow')\n" + ) + assert _codes(tmp_path, src).count("LIT015") == 1 + + +def test_root_model_without_frozen_is_flagged(tmp_path): + src = "from pydantic import RootModel\nclass P(RootModel):\n root: int\n" + assert "LIT015" in _codes(tmp_path, src) + + +def test_qualified_pydantic_basemodel_is_flagged(tmp_path): + src = "import pydantic\nclass P(pydantic.BaseModel):\n a: int\n" + assert "LIT015" in _codes(tmp_path, src) + + +def test_keyword_frozen_model_is_clean(tmp_path): + src = "from pydantic import BaseModel\nclass M(BaseModel, frozen=True):\n a: int\n" + assert "LIT015" not in _codes(tmp_path, src) + + +def test_subclass_of_keyword_frozen_model_is_clean(tmp_path): + src = ( + "from pydantic import BaseModel\n" + "class Parent(BaseModel, frozen=True):\n" + " pass\n" + "class Child(Parent):\n" + " a: int\n" + ) + assert "LIT015" not in _codes(tmp_path, src) + + +def test_frozen_false_keyword_overrides_body_frozen_config(tmp_path): + src = ( + "from pydantic import BaseModel, ConfigDict\n" + "class M(BaseModel, frozen=False):\n" + " model_config = ConfigDict(frozen=True)\n" + ) + assert "LIT015" in _codes(tmp_path, src) + + +def test_frozen_ok_with_reason_suppresses_and_is_not_lit013(tmp_path): + src = ( + "from pydantic import BaseModel\n" + "class P(BaseModel): # frozen-ok: mutated during build before handoff\n" + " a: int\n" + ) + codes = _codes(tmp_path, src) + assert "LIT015" not in codes + assert "LIT013" not in codes + + +def test_frozen_ok_on_frozen_model_is_lit013(tmp_path): + src = ( + "from pydantic import BaseModel\n" + "class P(BaseModel, frozen=True): # frozen-ok: mutable during construction\n" + " a: int\n" + ) + assert _codes(tmp_path, src) == ["LIT013"] + + +def test_frozen_ok_without_reason_is_lit005_and_does_not_suppress(tmp_path): + src = "from pydantic import BaseModel\nclass P(BaseModel): # frozen-ok\n a: int\n" + codes = _codes(tmp_path, src) + assert "LIT005" in codes + assert "LIT015" in codes + + +def test_typeddict_and_plain_classes_are_not_models(tmp_path): + src = "from typing import TypedDict\nclass T(TypedDict):\n a: int\nclass C:\n a: int\n" + assert "LIT015" not in _codes(tmp_path, src) + + +def test_extra_allow_does_not_exempt(tmp_path): + src = ( + "from pydantic import BaseModel, ConfigDict\n" + "class P(BaseModel):\n" + " model_config = ConfigDict(extra='allow')\n" + ) + assert "LIT015" in _codes(tmp_path, src) + + +def test_litellm_base_model_is_flagged(tmp_path): + src = "class Foo(LiteLLMBaseModel):\n x: int\n" + assert "LIT015" in _codes(tmp_path, src) + + +def test_litellm_openai_response_base_is_flagged(tmp_path): + src = "class R(BaseLiteLLMOpenAIResponseObject):\n x: int\n" + assert "LIT015" in _codes(tmp_path, src) + + +def test_base_settings_is_flagged(tmp_path): + src = "class S(BaseSettings):\n x: int\n" + assert "LIT015" in _codes(tmp_path, src) + + +def test_settings_config_dict_frozen_true_is_clean(tmp_path): + src = ( + "from pydantic_settings import BaseSettings, SettingsConfigDict\n" + "class S(BaseSettings):\n" + " model_config = SettingsConfigDict(frozen=True)\n" + ) + assert "LIT015" not in _codes(tmp_path, src) + + +def test_openai_object_base_is_flagged_and_can_be_frozen(tmp_path): + assert "LIT015" in _codes(tmp_path, "class Foo(OpenAIObject):\n x: int\n") + assert "LIT015" not in _codes( + tmp_path, + "from pydantic import ConfigDict\n" + "class Foo(OpenAIObject):\n" + " model_config = ConfigDict(frozen=True)\n" + " x: int\n", + ) + + +def test_same_named_models_use_each_classes_own_override(tmp_path): + src = ( + "from pydantic import BaseModel\n" + "class Dup(BaseModel, frozen=True):\n" + " x: int\n" + "class Dup(BaseModel):\n" + " x: int\n" + ) + assert _codes(tmp_path, src) == ["LIT015"] + + +def test_pydantic_model_discovery_uses_each_class_nodes_bases(tmp_path): + src: Final = ( + "from pydantic import BaseModel\n" + "class Config(BaseModel):\n" + " x: int\n" + "class Consumer(BaseModel, frozen=True):\n" + " class Config:\n" + " arbitrary_types_allowed = True\n" + ) + path: Final = tmp_path / "snippet.py" + path.write_text(src, encoding="utf-8") + violations: Final = checker.check_file(path) + assert [(violation.line, violation.code) for violation in violations] == [(2, "LIT015")] + + # --------------------------------------------------------------------------- # # Stacked comprehension clauses (LIT014) # --------------------------------------------------------------------------- # @@ -751,20 +1021,6 @@ def test_violation_message_names_the_clause_counts(tmp_path: Path): assert "2 `for` clauses and 1 `if` clause" in messages[0] -# --------------------------------------------------------------------------- # -# Budget integrity: every emittable LIT rule (bar the LIT000 read/parse error) is gated -# --------------------------------------------------------------------------- # - - -def test_budget_covers_exactly_the_checker_rules(): - budget = json.loads((_REPO_ROOT / "type-discipline-budget.json").read_text()) - emitted = set(re.findall(r"LIT\d{3}", _MODULE_PATH.read_text(encoding="utf-8"))) - {"LIT000"} - assert set(budget) == emitted - for spec in budget.values(): - assert isinstance(spec["limit"], int) - assert spec["limit"] >= 0 - - _FANS_OUT = checker._worker_count(checker.PARALLEL_MIN_PATHS) > 1 _SERIAL_ONLY = "one usable core, so scan_paths stays serial and there is no fan-out to compare" diff --git a/tests/unit/test_default_branch.py b/tests/unit/test_default_branch.py index 0401bec324f..bb04ea2a0fa 100644 --- a/tests/unit/test_default_branch.py +++ b/tests/unit/test_default_branch.py @@ -28,7 +28,7 @@ def remote_and_clone(tmp_path: Path) -> tuple[Path, Path]: (seed / "scripts").mkdir() for name in ( "default_branch.py", - "budget_ratchet_check.py", + "lint_base_counts.py", "ruff_strict_gate.py", "type_discipline_gate.py", "test_quality_gate.py", @@ -39,11 +39,9 @@ def remote_and_clone(tmp_path: Path) -> tuple[Path, Path]: shutil.copyfile(ROOT / "Makefile", seed / "Makefile") (seed / "litellm").mkdir() (seed / "litellm" / "example.py").write_text("value = 0\n") - (seed / "ruff-strict-budget.json").write_text('{"C901": {"limit": 1}}\n') _commit(seed, "staging base") _git(seed, "checkout", "-qb", "main") (seed / "litellm" / "example.py").write_text("value = 1\n") - (seed / "ruff-strict-budget.json").write_text('{"C901": {"limit": 0}}\n') _commit(seed, "main base") remote: Final = tmp_path / "remote.git" _git(tmp_path, "clone", "-q", "--bare", str(seed), str(remote)) @@ -121,28 +119,6 @@ def test_explicit_base_works_without_remote_access( assert "No changed litellm Python files" in checked.stdout -def test_budget_ratchet_compares_against_new_default(remote_and_clone: tuple[Path, Path]) -> None: - remote, repo = remote_and_clone - _git(remote, "symbolic-ref", "HEAD", "refs/heads/main") - resolved: Final = _resolve(repo) - assert resolved.returncode == 0, resolved.stderr - _git(repo, "checkout", "-qb", "litellm_feature", "origin/main") - (repo / "ruff-strict-budget.json").write_text('{"C901": {"limit": 1}}\n') - command: Final = [sys.executable, "scripts/budget_ratchet_check.py"] - checked: Final = subprocess.run(command, cwd=repo, capture_output=True, text=True, check=False) - assert checked.returncode == 1 - assert "limit raised 0 -> 1" in checked.stdout - assert "base origin/main" in checked.stdout - overridden: Final = subprocess.run( - [*command, "--base", "origin/release_branch"], - cwd=repo, - capture_output=True, - text=True, - check=False, - ) - assert overridden.returncode == 0, overridden.stdout + overridden.stderr - - def _freshness(repo: Path, *args: str) -> subprocess.CompletedProcess[str]: return subprocess.run( [ @@ -192,7 +168,6 @@ def test_migration_freshness_refuses_unavailable_remote(remote_and_clone: tuple[ @pytest.mark.parametrize( "gate", [ - "budget_ratchet_check", "ruff_strict_gate", "type_discipline_gate", "test_quality_gate", @@ -213,14 +188,11 @@ def test_each_gate_refuses_an_unverifiable_default(remote_and_clone: tuple[Path, assert "Cannot verify the base branch against origin" in result.stderr -@pytest.mark.parametrize( - "target", ["lint-format-check-changed", "lint-test-quality", "lint-test-quality-budget-update"] -) +@pytest.mark.parametrize("target", ["lint-format-check-changed", "lint-test-quality"]) def test_direct_make_target_fetches_default_once(remote_and_clone: tuple[Path, Path], target: str) -> None: _, repo = remote_and_clone trace: Final = repo.parent / "git-trace.jsonl" shutil.copyfile(ROOT / "scripts" / "check_test_quality.py", repo / "scripts" / "check_test_quality.py") - shutil.copyfile(ROOT / "test-quality-budget.json", repo / "test-quality-budget.json") (repo / "tests").mkdir() result: Final = subprocess.run( ["make", "-o", "install-dev", target, "LINT_DEP_INSTALL=", "UV_RUN=env"], diff --git a/tests/unit/test_lazy_imports.py b/tests/unit/test_lazy_imports.py index 10986f8a140..0416020a9b5 100644 --- a/tests/unit/test_lazy_imports.py +++ b/tests/unit/test_lazy_imports.py @@ -47,7 +47,7 @@ def test_import_litellm_does_not_load_fastapi_or_bpe_table(): [ sys.executable, "-c", - "import sys, litellm; print(','.join(m for m in ('fastapi','starlette','litellm.litellm_core_utils.default_encoding') if m in sys.modules))", + "import sys, litellm; print(','.join(m for m in ('fastapi','starlette','litellm.proxy.proxy_cli','litellm.litellm_core_utils.default_encoding') if m in sys.modules))", ], check=True, capture_output=True, diff --git a/tests/unit/test_lint_base_counts.py b/tests/unit/test_lint_base_counts.py new file mode 100644 index 00000000000..84f8a0936c0 --- /dev/null +++ b/tests/unit/test_lint_base_counts.py @@ -0,0 +1,454 @@ +"""Tests for scripts/lint_base_counts.py, the merge-base counting shared by the +four lint gates: the ceiling rule with its optional per-rule caps, the on-disk cache and +its eviction, the CI artifact fetch, the artifact emit, and the merge-base +resolution.""" + +import fnmatch +import io +import json +import os +import zipfile +from collections.abc import Callable, Mapping, Sequence +from pathlib import Path +from typing import Final, NamedTuple, NoReturn + +import pytest + +import lint_base_counts as counts + +_CHECKER: Final = counts.Checker("basedpyright", ("f1", "f2")) +_OTHER_CHECKER: Final = counts.Checker("ruff-strict", ("f1", "f2")) + + +def test_evaluate_passes_a_rule_that_did_not_grow() -> None: + assert counts.evaluate({"LIT006": 12}, {"LIT006": 12}) == () + + +def test_evaluate_blames_one_new_violation_of_an_uncapped_rule() -> None: + assert counts.evaluate({"LIT006": 13}, {"LIT006": 12}) == (counts.Breach("LIT006", 13, 12, 1),) + + +def test_evaluate_lets_a_capped_rule_grow_up_to_its_cap_and_no_further() -> None: + caps: Final = {"reportAny": 110} + assert counts.evaluate({"reportAny": 110}, {"reportAny": 100}, caps) == () + assert counts.evaluate({"reportAny": 111}, {"reportAny": 100}, caps) == (counts.Breach("reportAny", 111, 110, 11),) + + +def test_evaluate_never_blames_a_bystander_for_a_base_already_over_the_cap() -> None: + caps: Final = {"reportAny": 110} + assert counts.evaluate({"reportAny": 120}, {"reportAny": 120}, caps) == () + assert counts.evaluate({"reportAny": 119}, {"reportAny": 120}, caps) == () + + +def test_evaluate_holds_a_base_over_its_cap_to_no_growth() -> None: + caps: Final = {"reportAny": 110} + assert counts.evaluate({"reportAny": 121}, {"reportAny": 120}, caps) == (counts.Breach("reportAny", 121, 120, 1),) + + +def test_evaluate_cap_applies_only_to_the_rule_it_names() -> None: + caps: Final = {"reportAny": 110} + assert counts.evaluate({"LIT006": 13}, {"LIT006": 12}, caps) == (counts.Breach("LIT006", 13, 12, 1),) + + +def test_evaluate_counts_a_rule_absent_from_the_base_as_zero() -> None: + assert counts.evaluate({"NEW99": 1}, {}) == (counts.Breach("NEW99", 1, 0, 1),) + + +def test_evaluate_never_blames_a_change_that_reduced_a_rule() -> None: + assert counts.evaluate({"LIT006": 11}, {"LIT006": 12}) == () + + +def test_evaluate_reports_every_grown_rule_sorted_by_name() -> None: + head: Final = {"TQ008": 3, "TQ001": 2, "TQ003": 5} + base: Final = {"TQ008": 2, "TQ001": 1, "TQ003": 5} + assert [b.rule for b in counts.evaluate(head, base)] == ["TQ001", "TQ008"] + + +def test_evaluate_ignores_a_base_rule_the_head_fixed_entirely() -> None: + assert counts.evaluate({}, {"LIT006": 12}) == () + + +def test_cache_key_changes_with_base_point_and_each_fingerprint() -> None: + key: Final = counts.cache_key("abc", ("cfg", "lock")) + assert counts.cache_key("abc", ("cfg", "lock")) == key + assert counts.cache_key("def", ("cfg", "lock")) != key + assert counts.cache_key("abc", ("cfg2", "lock")) != key + assert counts.cache_key("abc", ("cfg", "lock2")) != key + + +def test_checker_names_its_artifact_and_cache_file_by_the_same_key() -> None: + key: Final = counts.cache_key("abc123", ("f1", "f2")) + assert _CHECKER.artifact_name("abc123") == f"basedpyright-counts-{key}" + assert _CHECKER.cache_file_name("abc123") == f"basedpyright-base-{key}.json" + assert fnmatch.fnmatch(_CHECKER.cache_file_name("abc123"), _CHECKER.cache_glob()) + + +def test_checkers_with_the_same_fingerprints_never_share_a_name() -> None: + assert _CHECKER.artifact_name("abc123") != _OTHER_CHECKER.artifact_name("abc123") + assert not fnmatch.fnmatch(_OTHER_CHECKER.cache_file_name("abc123"), _CHECKER.cache_glob()) + + +def test_cached_counts_round_trip(tmp_path: Path) -> None: + path: Final = counts.store_counts(tmp_path, _CHECKER, "abc123", {"reportAny": 3, "reportCall": 1}) + assert path == tmp_path / _CHECKER.cache_file_name("abc123") + assert counts.load_cached_counts(path) == {"reportAny": 3, "reportCall": 1} + + +@pytest.mark.parametrize( + "content", + [ + None, + "{not json", + json.dumps(["counts"]), + json.dumps({"base_point": "abc"}), + json.dumps({"counts": {"reportAny": "three"}}), + json.dumps({"counts": {"reportAny": True}}), + ], +) +def test_missing_corrupt_or_misshapen_cache_reads_as_none(tmp_path: Path, content: str | None) -> None: + path: Final = tmp_path / "cache.json" + if content is not None: + path.write_text(content) + assert counts.load_cached_counts(path) is None + + +def test_scratch_is_invisible_to_the_prune_glob() -> None: + scratch: Final = counts.scratch_path(Path("/c") / _CHECKER.cache_file_name("abc")) + assert not fnmatch.fnmatch(scratch.name, _CHECKER.cache_glob()) + + +def test_store_prune_spares_a_concurrent_runs_in_flight_scratch(tmp_path: Path) -> None: + foreign: Final = counts.scratch_path(tmp_path / _CHECKER.cache_file_name("other")) + foreign.write_text("{}") + mine: Final = counts.store_counts(tmp_path, _CHECKER, "mine", {"reportAny": 1}) + assert foreign.exists() + assert counts.load_cached_counts(mine) == {"reportAny": 1} + + +def test_store_keeps_a_concurrent_worktrees_entry_for_another_branch_point(tmp_path: Path) -> None: + old: Final = counts.store_counts(tmp_path, _CHECKER, "old", {"reportAny": 1}) + new: Final = counts.store_counts(tmp_path, _CHECKER, "new", {"reportAny": 2}) + assert counts.load_cached_counts(old) == {"reportAny": 1} + assert counts.load_cached_counts(new) == {"reportAny": 2} + + +def test_store_evicts_only_the_oldest_entries_beyond_the_cap(tmp_path: Path) -> None: + aged: Final = tuple( + counts.store_counts(tmp_path, _CHECKER, f"base{age}", {"reportAny": age}) + for age in range(counts.CACHE_KEEP_ENTRIES) + ) + for age, path in enumerate(aged): + os.utime(path, (age, age)) + newest: Final = counts.store_counts(tmp_path, _CHECKER, "newest", {"reportAny": 99}) + assert not aged[0].exists() + assert all(path.exists() for path in aged[1:]) + assert counts.load_cached_counts(newest) == {"reportAny": 99} + + +def test_store_never_evicts_the_entry_it_just_wrote_even_on_mtime_ties(tmp_path: Path) -> None: + for index in range(counts.CACHE_KEEP_ENTRIES + 2): + os.utime(counts.store_counts(tmp_path, _CHECKER, f"base{index}", {"reportAny": 1}), (9_999_999_999,) * 2) + mine: Final = counts.store_counts(tmp_path, _CHECKER, "mine", {"reportAny": 2}) + assert counts.load_cached_counts(mine) == {"reportAny": 2} + assert len(list(tmp_path.glob(_CHECKER.cache_glob()))) == counts.CACHE_KEEP_ENTRIES + + +def test_store_eviction_never_touches_another_checkers_entries(tmp_path: Path) -> None: + other: Final = counts.store_counts(tmp_path, _OTHER_CHECKER, "base", {"E501": 1}) + os.utime(other, (1, 1)) + for index in range(counts.CACHE_KEEP_ENTRIES + 1): + counts.store_counts(tmp_path, _CHECKER, f"base{index}", {"reportAny": 1}) + assert counts.load_cached_counts(other) == {"E501": 1} + + +def _no_fetch(checker: counts.Checker, base_point: str) -> None: + return None + + +def _never(reason: str) -> Callable[..., NoReturn]: + def callback(*args: object) -> NoReturn: + raise AssertionError(reason) + + return callback + + +def test_base_counts_cached_returns_the_hit_without_recomputing(tmp_path: Path) -> None: + counts.store_counts(tmp_path, _CHECKER, "abc123", {"reportAny": 7}) + assert counts.base_counts_cached( + _CHECKER, + "abc123", + _never("a cache hit must not re-run the base pass"), + cache_dir=tmp_path, + fetch=_never("a cache hit must not reach for CI"), + ) == {"reportAny": 7} + + +def test_base_counts_cached_computes_once_then_hits(tmp_path: Path) -> None: + calls: Final[list[str]] = [] + + def fake(ref: str) -> counts.Counts: + calls.append(ref) + return {"reportAny": 4} + + first: Final = counts.base_counts_cached(_CHECKER, "abc123", fake, cache_dir=tmp_path, fetch=_no_fetch) + second: Final = counts.base_counts_cached(_CHECKER, "abc123", fake, cache_dir=tmp_path, fetch=_no_fetch) + assert first == second == {"reportAny": 4} + assert calls == ["abc123"] + + +def test_base_counts_cached_keeps_each_checker_apart(tmp_path: Path) -> None: + counts.store_counts(tmp_path, _OTHER_CHECKER, "abc123", {"E501": 7}) + assert counts.base_counts_cached( + _CHECKER, "abc123", lambda ref: {"reportAny": 4}, cache_dir=tmp_path, fetch=_no_fetch + ) == {"reportAny": 4} + + +def test_an_empty_base_pass_is_never_cached(tmp_path: Path) -> None: + calls: Final[list[str]] = [] + + def crashed(ref: str) -> counts.Counts: + calls.append(ref) + return {} + + assert counts.base_counts_cached(_CHECKER, "abc123", crashed, cache_dir=tmp_path, fetch=_no_fetch) == {} + assert counts.base_counts_cached(_CHECKER, "abc123", crashed, cache_dir=tmp_path, fetch=_no_fetch) == {} + assert calls == ["abc123", "abc123"] + assert list(tmp_path.iterdir()) == [] + + +def test_base_counts_cached_uses_fetched_counts_and_persists_them(tmp_path: Path) -> None: + fetched: Final = counts.base_counts_cached( + _CHECKER, + "abc123", + _never("fetched counts must skip the local base pass"), + cache_dir=tmp_path, + fetch=lambda checker, base_point: {"reportAny": 9}, + ) + assert fetched == {"reportAny": 9} + assert counts.load_cached_counts(tmp_path / _CHECKER.cache_file_name("abc123")) == {"reportAny": 9} + assert counts.base_counts_cached( + _CHECKER, + "abc123", + _never("the persisted fetch must satisfy later runs"), + cache_dir=tmp_path, + fetch=_never("the persisted fetch must satisfy later runs"), + ) == {"reportAny": 9} + + +def test_base_counts_cached_hands_the_fetcher_the_checker_and_base_point(tmp_path: Path) -> None: + seen: Final[list[tuple[counts.Checker, str]]] = [] + + def fetch(checker: counts.Checker, base_point: str) -> None: + seen.append((checker, base_point)) + + counts.base_counts_cached(_CHECKER, "abc123", lambda ref: {"reportAny": 4}, cache_dir=tmp_path, fetch=fetch) + assert seen == [(_CHECKER, "abc123")] + + +def test_base_counts_cached_falls_back_to_compute_on_a_fetch_miss(tmp_path: Path) -> None: + calls: Final[list[str]] = [] + + def local(ref: str) -> counts.Counts: + calls.append(ref) + return {"reportAny": 4} + + assert counts.base_counts_cached(_CHECKER, "abc123", local, cache_dir=tmp_path, fetch=_no_fetch) == { + "reportAny": 4 + } + assert calls == ["abc123"] + + +def test_base_counts_cached_treats_empty_fetched_counts_as_a_miss(tmp_path: Path) -> None: + assert counts.base_counts_cached( + _CHECKER, + "abc123", + lambda ref: {"reportAny": 2}, + cache_dir=tmp_path, + fetch=lambda checker, base_point: {}, + ) == {"reportAny": 2} + assert counts.load_cached_counts(tmp_path / _CHECKER.cache_file_name("abc123")) == {"reportAny": 2} + + +def test_origin_slug_parsing_supports_ssh_and_https_github_forms() -> None: + assert counts.parse_origin_slug("git@github.com:BerriAI/litellm.git") == "BerriAI/litellm" + assert counts.parse_origin_slug("git@github.com:BerriAI/litellm") == "BerriAI/litellm" + assert counts.parse_origin_slug("https://github.com/BerriAI/litellm.git") == "BerriAI/litellm" + assert counts.parse_origin_slug("https://github.com/BerriAI/litellm") == "BerriAI/litellm" + assert counts.parse_origin_slug("https://github.com/BerriAI/litellm/") == "BerriAI/litellm" + + +def test_origin_slug_parsing_rejects_non_github_urls() -> None: + assert counts.parse_origin_slug("https://gitlab.com/BerriAI/litellm.git") is None + assert counts.parse_origin_slug("git@bitbucket.org:BerriAI/litellm.git") is None + assert counts.parse_origin_slug("not a url") is None + assert counts.parse_origin_slug("") is None + + +def _artifact_zip(payload: Mapping[str, object]) -> bytes: + buffer: Final = io.BytesIO() + with zipfile.ZipFile(buffer, "w") as archive: + archive.writestr("counts.json", json.dumps(payload)) + return buffer.getvalue() + + +def _gh_stub( + listing: Mapping[str, object], zip_bytes: bytes, seen: list[tuple[str, ...]] | None = None +) -> counts.GhOutput: + def gh_output(args: Sequence[str]) -> bytes: + if seen is not None: + seen.append(tuple(args)) + if args[-1].startswith("repos/"): + return json.dumps(listing).encode() + return zip_bytes + + return gh_output + + +def _live_listing() -> Mapping[str, object]: + return {"artifacts": [{"expired": False, "archive_download_url": "https://api.github.com/x/zip"}]} + + +def test_fetcher_returns_counts_from_a_matching_artifact(capsys: pytest.CaptureFixture[str]) -> None: + payload: Final = {"base_point": "abc123", "counts": {"reportAny": 3}} + fetched: Final = counts.fetch_ci_base_counts( + _CHECKER, "abc123", gh=_gh_stub(_live_listing(), _artifact_zip(payload)) + ) + assert fetched == {"reportAny": 3} + assert "fetched from CI artifact" in capsys.readouterr().err + + +def test_fetcher_asks_for_the_artifact_named_by_the_checker_and_base_point() -> None: + seen: Final[list[tuple[str, ...]]] = [] + payload: Final = {"base_point": "abc123", "counts": {"reportAny": 3}} + counts.fetch_ci_base_counts(_CHECKER, "abc123", gh=_gh_stub(_live_listing(), _artifact_zip(payload), seen)) + listing_request: Final = seen[0][-1] + assert f"name={_CHECKER.artifact_name('abc123')}" in listing_request + assert seen[1][-1] == "https://api.github.com/x/zip" + + +def test_fetcher_rejects_an_artifact_for_a_different_base_point() -> None: + payload: Final = {"base_point": "someothersha", "counts": {"reportAny": 3}} + assert ( + counts.fetch_ci_base_counts(_CHECKER, "abc123", gh=_gh_stub(_live_listing(), _artifact_zip(payload))) is None + ) + + +@pytest.mark.parametrize("bad_counts", [{}, {"reportAny": "three"}, {"reportAny": True}]) +def test_fetcher_rejects_empty_or_misshapen_artifact_counts(bad_counts: Mapping[str, object]) -> None: + payload: Final = {"base_point": "abc123", "counts": bad_counts} + assert ( + counts.fetch_ci_base_counts(_CHECKER, "abc123", gh=_gh_stub(_live_listing(), _artifact_zip(payload))) is None + ) + + +def test_fetcher_rejects_an_expired_artifact() -> None: + listing: Final = {"artifacts": [{"expired": True, "archive_download_url": "https://api.github.com/x/zip"}]} + payload: Final = {"base_point": "abc123", "counts": {"reportAny": 3}} + assert counts.fetch_ci_base_counts(_CHECKER, "abc123", gh=_gh_stub(listing, _artifact_zip(payload))) is None + + +def test_fetcher_misses_when_no_artifact_is_published() -> None: + assert counts.fetch_ci_base_counts(_CHECKER, "abc123", gh=_gh_stub({"artifacts": []}, b"")) is None + + +def test_fetcher_misses_when_gh_is_unusable(capsys: pytest.CaptureFixture[str]) -> None: + assert counts.fetch_ci_base_counts(_CHECKER, "abc123", gh=lambda args: None) is None + assert "computing base counts locally" in capsys.readouterr().err + + +def test_fetcher_misses_on_a_corrupt_artifact_archive() -> None: + assert counts.fetch_ci_base_counts(_CHECKER, "abc123", gh=_gh_stub(_live_listing(), b"not a zip")) is None + + +def test_emit_writes_the_artifact_json_named_by_the_head_key( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + counts.emit_counts(_CHECKER, {"reportAny": 3, "aRule": 1}, tmp_path, "deadbeef") + name: Final = _CHECKER.artifact_name("deadbeef") + assert json.loads((tmp_path / f"{name}.json").read_text()) == { + "base_point": "deadbeef", + "counts": {"aRule": 1, "reportAny": 3}, + } + summary: Final = capsys.readouterr().out + assert "deadbeef" in summary + assert name in summary + assert "4" in summary + + +def test_emit_refuses_to_publish_empty_counts(tmp_path: Path) -> None: + with pytest.raises(SystemExit): + counts.emit_counts(_CHECKER, {}, tmp_path, "deadbeef") + assert list(tmp_path.iterdir()) == [] + + +def test_emitted_file_is_the_one_the_fetcher_looks_up(tmp_path: Path) -> None: + written: Final = counts.emit_counts(_CHECKER, {"reportAny": 3}, tmp_path, "deadbeef") + payload: Final = json.loads(written.read_text()) + assert counts.counts_for_base(payload, "deadbeef") == {"reportAny": 3} + assert counts.counts_for_base(payload, "someothersha") is None + listing_zip: Final = _artifact_zip(payload) + assert counts.fetch_ci_base_counts(_CHECKER, "deadbeef", gh=_gh_stub(_live_listing(), listing_zip)) == { + "reportAny": 3 + } + + +class _History(NamedTuple): + parents: Mapping[str, tuple[str, ...]] + refs: Mapping[str, str] + + def ancestry(self, commit: str) -> frozenset[str]: + return frozenset((commit,)).union(*(self.ancestry(parent) for parent in self.parents[commit])) + + def merge_base(self, left: str, right: str) -> str: + common: Final = self.ancestry(self.refs.get(left, left)) & self.ancestry(self.refs.get(right, right)) + return next(c for c in common if not any(c != other and c in self.ancestry(other) for other in common)) + + +def _git_over(history: _History) -> counts.Git: + def git(args: Sequence[str]) -> str: + match tuple(args): + case ("merge-base", left, right): + return f"{history.merge_base(left, right)}\n" + case ("rev-parse", "HEAD"): + return f"{history.refs['HEAD']}\n" + case ("rev-parse", "--verify", "--quiet", ref): + return f"{history.refs[ref]}\n" if ref in history.refs else "" + case ("rev-parse", "--path-format=absolute", "--git-common-dir"): + return "/repo/.git\n" + case ("rev-parse", "--path-format=absolute", "--git-dir"): + return "/repo/.git/worktrees/feature\n" + case _: + raise AssertionError(f"unexpected git call: {args}") + + return git + + +_FEATURE_OFF_MAIN: Final = _History( + parents={"shared": (), "feature": ("shared",), "drift": ("shared",)}, + refs={"main": "drift", "HEAD": "feature"}, +) + + +def test_base_point_is_the_branch_point_when_no_merge_is_in_progress() -> None: + assert counts.resolve_base_point("main", _git_over(_FEATURE_OFF_MAIN)) == "shared" + + +def test_base_point_mid_merge_advances_to_the_merged_in_base_tip() -> None: + merging_main: Final = _FEATURE_OFF_MAIN._replace(refs={**_FEATURE_OFF_MAIN.refs, "MERGE_HEAD": "drift"}) + assert counts.resolve_base_point("main", _git_over(merging_main)) == "drift" + + +def test_base_point_mid_merge_of_an_older_side_branch_keeps_the_newer_branch_point() -> None: + merging_old_side: Final = _History( + parents={"shared": (), "old": ("shared",), "drift": ("shared",), "feature": ("drift",)}, + refs={"main": "drift", "HEAD": "feature", "MERGE_HEAD": "old"}, + ) + assert counts.resolve_base_point("main", _git_over(merging_old_side)) == "drift" + + +def test_head_sha_is_the_checked_out_commit() -> None: + assert counts.head_sha(_git_over(_FEATURE_OFF_MAIN)) == "feature" + + +def test_default_cache_dir_is_shared_by_every_worktree() -> None: + assert counts.default_cache_dir(_git_over(_FEATURE_OFF_MAIN)) == Path("/repo/.git") / counts.CACHE_DIR_NAME diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py index 6bbab41658d..1929cbc871f 100644 --- a/tests/unit/test_main.py +++ b/tests/unit/test_main.py @@ -25,6 +25,8 @@ import litellm from litellm import acompletion, completion from litellm import main as litellm_main from litellm.constants import CONTROL_OPTIONS_KEY +from litellm.caching.base_cache import BaseCache +from litellm.caching.caching import Cache from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.custom_prompt_management import CustomPromptManagement from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs @@ -33,7 +35,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.prompt_templates.factory import anthropic_messages_pt from litellm.llms.custom_httpx.http_handler import HTTPHandler from litellm.types.litellm_params import ControlOptions -from litellm.types.llms.openai import AllMessageValues +from litellm.types.llms.openai import AllMessageValues, HttpxBinaryResponseContent from litellm.types.prompts.init_prompts import PromptSpec from litellm.types.utils import Delta, ModelResponseStream, StandardCallbackDynamicParams, StreamingChoices, Usage @@ -4628,6 +4630,45 @@ def test_completion_rejects_an_invalid_stream_chunk_size_before_the_mcp_gateway( assert exc_info.value.param == "stream_chunk_size" +@pytest.mark.asyncio +@pytest.mark.parametrize("use_async", [False, True]) +@pytest.mark.parametrize("missing_tenacity", [False, True]) +@pytest.mark.parametrize("route", ["bedrock", "bedrock/invoke"]) +async def test_bedrock_stream_missing_dependency_remains_actionable_with_retries( + monkeypatch, use_async, missing_tenacity, route +): + import builtins + + original_import = builtins.__import__ + + def import_without_aws_or_retry_dependencies(name, *args, **kwargs): + if name.split(".")[0] == "botocore" or (name == "tenacity" and missing_tenacity): + raise ModuleNotFoundError(name=name.split(".")[0]) + return original_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", import_without_aws_or_retry_dependencies) + monkeypatch.setattr(litellm, "num_retries", None) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + with respx.mock as upstream: + response = upstream.post(url__regex=r"https://bedrock-test\.invalid/.*").respond(200, content=b"") + arguments = dict( + model=f"{route}/anthropic.claude-3-sonnet-20240229-v1:0", + messages=[{"role": "user", "content": "ping"}], + api_key="test-bearer", + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint="https://bedrock-test.invalid", + stream=True, + num_retries=1, + ) + if use_async: + with pytest.raises(ImportError, match="pip install boto3"): + await litellm.acompletion(**arguments) + else: + with pytest.raises(ImportError, match="pip install boto3"): + litellm.completion(**arguments) + assert response.call_count == 1 + + def test_drop_params_false_still_rejects_an_invalid_stream_chunk_size() -> None: with pytest.raises(litellm.BadRequestError): litellm.completion( @@ -5831,3 +5872,235 @@ def test_completion_openai_metadata(monkeypatch, enable_preview_features): } else: assert "metadata" not in mock_completion.call_args.kwargs + + +AZURE_TTS_BASE: Final = "https://tts.example.azure.com" +SPEECH_INPUT: Final = "the quick brown fox jumped over the lazy dogs" + + +@pytest.mark.parametrize("sync_mode", [True, False]) +async def test_speech_azure_returns_binary_audio( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, sync_mode: bool +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + route: Final = respx_mock.post( + url__regex=rf"{AZURE_TTS_BASE}/openai/deployments/tts/audio/speech\?api-version=.+" + ).mock(return_value=httpx.Response(200, content=b"ID3-fake-mp3")) + + speech_kwargs: Final = { + "model": "azure/tts", + "input": SPEECH_INPUT, + "voice": "alloy", + "api_base": AZURE_TTS_BASE, + "api_key": "fake-key", + "max_retries": 1, + "timeout": 60, + } + response: Final = litellm.speech(**speech_kwargs) if sync_mode else await litellm.aspeech(**speech_kwargs) + + assert route.call_count == 1 + assert route.calls[0].request.headers["api-key"] == "fake-key" + assert json.loads(route.calls[0].request.content) == {"model": "tts", "input": SPEECH_INPUT, "voice": "alloy"} + assert isinstance(response, HttpxBinaryResponseContent) + assert response.content == b"ID3-fake-mp3" + + +@pytest.mark.parametrize("sync_mode", [True, False]) +async def test_speech_openai_returns_binary_audio( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, sync_mode: bool +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + route: Final = respx_mock.post("https://api.openai.com/v1/audio/speech").mock( + return_value=httpx.Response(200, content=b"ID3-fake-mp3") + ) + + speech_kwargs: Final = { + "model": "openai/tts-1", + "input": SPEECH_INPUT, + "voice": "alloy", + "api_key": "fake-key", + "max_retries": 1, + "timeout": 60, + } + response: Final = litellm.speech(**speech_kwargs) if sync_mode else await litellm.aspeech(**speech_kwargs) + + assert route.call_count == 1 + assert route.calls[0].request.headers["authorization"] == "Bearer fake-key" + assert json.loads(route.calls[0].request.content) == {"model": "tts-1", "input": SPEECH_INPUT, "voice": "alloy"} + assert isinstance(response, HttpxBinaryResponseContent) + assert response.content == b"ID3-fake-mp3" + + +class _SignallingCache(Cache): + def __init__(self, loop: asyncio.AbstractEventLoop) -> None: + super().__init__() + self.loop: Final = loop + self.written: Final = asyncio.Event() + + async def async_add_cache( + self, result: object, dynamic_cache_object: BaseCache | None = None, **kwargs: object + ) -> None: + await super().async_add_cache(result, dynamic_cache_object=dynamic_cache_object, **kwargs) + self.loop.call_soon_threadsafe(self.written.set) + + +GETTYSBURG_WAV: Final = ("gettysburg.wav", b"RIFF\x00\x00\x00\x00WAVE-gettysburg", "audio/wav") +EAGLE_WAV: Final = ("eagle.wav", b"RIFF\x00\x00\x00\x00WAVE-eagle", "audio/wav") + + +async def test_transcription_caching_hit_same_file_miss_different_file( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + cache: Final = _SignallingCache(asyncio.get_running_loop()) + monkeypatch.setattr(litellm, "cache", cache) + route: Final = respx_mock.post("https://api.openai.com/v1/audio/transcriptions").mock( + side_effect=[ + httpx.Response(200, json={"text": "gettysburg transcript"}), + httpx.Response(200, json={"text": "eagle transcript"}), + ] + ) + + response_1: Final = await litellm.atranscription(model="openai/whisper-1", file=GETTYSBURG_WAV, api_key="fake-key") + await asyncio.wait_for(cache.written.wait(), 30) + + response_2: Final = await litellm.atranscription(model="openai/whisper-1", file=GETTYSBURG_WAV, api_key="fake-key") + assert response_2._hidden_params["cache_hit"] is True + assert response_2.text == response_1.text == "gettysburg transcript" + + response_3: Final = await litellm.atranscription(model="openai/whisper-1", file=EAGLE_WAV, api_key="fake-key") + assert response_3._hidden_params.get("cache_hit") is not True + assert response_3.text == "eagle transcript" + assert route.call_count == 2 + + +async def test_whisper_log_pre_call_fires_once(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + respx_mock.post("https://api.openai.com/v1/audio/transcriptions").mock( + return_value=httpx.Response(200, json={"text": "hello"}) + ) + + class _PreCallRecorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.models: tuple[str, ...] = () + + def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None: + self.models = (*self.models, model) + + recorder: Final = _PreCallRecorder() + monkeypatch.setattr(litellm, "callbacks", [recorder]) + + await litellm.atranscription(model="openai/whisper-1", file=GETTYSBURG_WAV, api_key="fake-key") + + assert recorder.models == ("whisper-1",) + + +@pytest.mark.parametrize("model", ["gpt-4o-mini-transcribe", "gpt-4o-transcribe", "whisper-1"]) +async def test_transcription_model_names_pass_through( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, model: str +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + route: Final = respx_mock.post("https://api.openai.com/v1/audio/transcriptions").mock( + return_value=httpx.Response(200, json={"text": "hello"}) + ) + + response: Final = await litellm.atranscription( + model=f"openai/{model}", + file=GETTYSBURG_WAV, + api_key="fake-key", + response_format="json", + ) + + assert response._hidden_params["model"] == model + assert response._hidden_params["custom_llm_provider"] == "openai" + assert response.text == "hello" + assert route.call_count == 1 + assert f'name="model"\r\n\r\n{model}\r\n'.encode() in route.calls[0].request.content + + +@pytest.mark.asyncio +@pytest.mark.parametrize("use_async", [False, True]) +@pytest.mark.parametrize("model", ["sagemaker/test-endpoint", "sagemaker_chat/test-endpoint"]) +async def test_sagemaker_missing_dependency_remains_actionable_with_retries(monkeypatch, use_async, model): + import sys + + monkeypatch.setattr(litellm, "num_retries", None) + for dependency in ("botocore", "boto3", "tenacity"): + monkeypatch.setitem(sys.modules, dependency, None) + if use_async: + with pytest.raises(ModuleNotFoundError, match="pip install boto3") as caught: + await litellm.acompletion(model=model, messages=[{"role": "user", "content": "ping"}], num_retries=1) + else: + with pytest.raises(ModuleNotFoundError, match="pip install boto3") as caught: + litellm.completion(model=model, messages=[{"role": "user", "content": "ping"}], num_retries=1) + assert caught.value.name == "botocore" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("use_async", [False, True]) +async def test_polly_missing_dependency_remains_actionable_with_retries(monkeypatch, use_async): + import sys + + monkeypatch.setattr(litellm, "num_retries", None) + for dependency in ("botocore", "boto3", "tenacity"): + monkeypatch.setitem(sys.modules, dependency, None) + if use_async: + with pytest.raises(ModuleNotFoundError, match="pip install boto3") as caught: + await litellm.aspeech(model="aws_polly/standard", input="ping", voice="Joanna", num_retries=1) + else: + with pytest.raises(ModuleNotFoundError, match="pip install boto3") as caught: + litellm.speech(model="aws_polly/standard", input="ping", voice="Joanna", num_retries=1) + assert caught.value.name == "botocore" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("missing_tenacity", [False, True]) +@pytest.mark.parametrize("use_async", [False, True]) +async def test_mantle_responses_missing_dependency_is_not_retried(monkeypatch, missing_tenacity, use_async): + import builtins + + original_import = builtins.__import__ + attempts = [] + + def import_without_aws(name, *args, **kwargs): + if name == "botocore": + attempts.append(name) + raise ModuleNotFoundError(name="botocore") + if name == "tenacity" and missing_tenacity: + attempts.append(name) + raise ModuleNotFoundError(name="tenacity") + return original_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", import_without_aws) + monkeypatch.setattr(litellm, "num_retries", None) + for name in ("AWS_BEARER_TOKEN_BEDROCK", "BEDROCK_MANTLE_API_KEY"): + monkeypatch.delenv(name, raising=False) + if use_async: + with pytest.raises(ModuleNotFoundError, match="pip install boto3"): + await litellm.aresponses(model="bedrock_mantle/openai.gpt-oss-120b", input="ping", num_retries=1) + else: + with pytest.raises(ModuleNotFoundError, match="pip install boto3"): + litellm.responses(model="bedrock_mantle/openai.gpt-oss-120b", input="ping", num_retries=1) + assert attempts == ["botocore"] + + +@pytest.mark.asyncio +async def test_async_responses_still_retries_provider_server_errors(monkeypatch): + monkeypatch.setattr(litellm, "num_retries", None) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + with respx.mock as upstream: + response = upstream.post("https://openai-test.invalid/v1/responses").mock(side_effect=[ + httpx.Response(500, json={"error": {"message": "temporary provider failure", "type": "server_error"}}), + httpx.Response(200, json={ + "id": "resp-retry", "object": "response", "created_at": 1, "status": "completed", + "model": "test-model", "output": [], + "usage": {"input_tokens": 3, "output_tokens": 2, "total_tokens": 5}, + }), + ]) + result = await litellm.aresponses( + model="openai/test-model", input="ping", api_key="test-key", + api_base="https://openai-test.invalid/v1", num_retries=1, max_retries=0, + ) + assert result.status == "completed" + assert response.call_count == 2 diff --git a/tests/unit/test_mistral_large_4_model_metadata.py b/tests/unit/test_mistral_large_4_model_metadata.py index 924265d2b41..f45c53920dd 100644 --- a/tests/unit/test_mistral_large_4_model_metadata.py +++ b/tests/unit/test_mistral_large_4_model_metadata.py @@ -31,7 +31,7 @@ def test_metadata_and_backup(model): info = litellm.get_model_info(model) assert info["litellm_provider"] == "mistral" assert info["mode"] == "chat" - assert info["max_input_tokens"] == 524288 + assert info["max_input_tokens"] == 1048576 assert info["input_cost_per_token"] == 6.8e-07 assert info["output_cost_per_token"] == 2.09e-06 assert info["cache_read_input_token_cost"] == 6.8e-08 diff --git a/tests/unit/test_pre_commit_lint.py b/tests/unit/test_pre_commit_lint.py index 471c8b41b5c..7c9b98c1c94 100644 --- a/tests/unit/test_pre_commit_lint.py +++ b/tests/unit/test_pre_commit_lint.py @@ -14,9 +14,9 @@ from tests._process_helpers import process_is_gone ROOT = Path(__file__).resolve().parents[2] SCRIPT = ROOT / "scripts" / "pre_commit_lint.sh" WHOLE_TREE_RUFF = "run --no-sync ruff check --config ruff-tests.toml tests" -TEST_TREE_RAN = "ran: test-tree lint (ruff-tests.toml + test-quality budget)" +TEST_TREE_RAN = "ran: test-tree lint (ruff-tests.toml + test-quality gate)" TEST_TREE_SKIPPED = ( - "skipped: test-tree lint (ruff-tests.toml + test-quality budget) " + "skipped: test-tree lint (ruff-tests.toml + test-quality gate) " "(no tests/ Python files or test-tree lint inputs in scope)" ) @@ -470,7 +470,6 @@ def test_tests_only_change_runs_the_whole_test_tree_ruff_and_the_quality_gate(tm "changed", [ "ruff-tests.toml", - "test-quality-budget.json", "scripts/check_test_quality.py", "scripts/test_quality_gate.py", "tests/e2e/test_x.py", @@ -523,7 +522,7 @@ def test_a_failing_quality_gate_fails_a_tests_only_run(tmp_path: Path) -> None: _stage_file(repo, "tests/test_a.py", "def test_a() -> None: ...\n") proc = _run(repo, bin_dir, {"STUB_FAIL": "test-quality"}) assert proc.returncode == 1 - assert "Test-quality budget failed" in proc.stdout + proc.stderr + assert "Test-quality gate failed" in proc.stdout + proc.stderr assert "check: FAIL" in proc.stdout @@ -566,7 +565,7 @@ def test_partial_staging_warns_when_test_files_are_left_unstaged(tmp_path: Path) (repo / "tests" / "test_a.py").write_text("def test_a() -> None:\n assert True\n") proc = _run(repo, bin_dir, {"STUB_ARGS_DIR": str(args_dir)}) assert proc.returncode == 0, proc.stdout + proc.stderr - assert "SKIPPED test-tree lint (ruff-tests.toml + test-quality budget)" in proc.stdout + assert "SKIPPED test-tree lint (ruff-tests.toml + test-quality gate)" in proc.stdout assert "tests/test_a.py" in proc.stdout assert _recorded(args_dir, "ruff_tests.args") == [] assert _recorded(args_dir, "make.args") == [] diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index ec9e4861fdd..4356488a2f8 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -63,7 +63,10 @@ from litellm.router_utils.cooldown_handlers import ( async_get_cooldown_deployments, get_cooldown_deployments, ) -from litellm.router_utils.fallback_event_handlers import DISABLE_FALLBACKS_METADATA_KEY +from litellm.router_utils.fallback_event_handlers import ( + DISABLE_FALLBACKS_METADATA_KEY, + MID_STREAM_FALLBACK_CONTROLS_KEY, +) from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute from litellm.scheduler import FlowItem from litellm.types.llms.openai import ChatCompletionRequest @@ -3868,6 +3871,136 @@ async def test_acompletion_mid_stream_fallback_walks_every_entry_of_the_configur ] +class _DiesBeforeFirstChunk(CustomStreamWrapper): + def __init__(self, model: str): + super().__init__(completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock()) + + def _mid_stream_error(self) -> MidStreamFallbackError: + return MidStreamFallbackError( + message=f"provider 500 from {self.model}", + model=self.model, + llm_provider="openai", + generated_content="", + is_pre_first_chunk=True, + original_exception=litellm.InternalServerError( + message=f"provider 500 from {self.model}", model=self.model, llm_provider="openai" + ), + ) + + def __iter__(self): + return self + + def __next__(self): + raise self._mid_stream_error() + + def __aiter__(self): + return self + + async def __anext__(self): + raise self._mid_stream_error() + + +class _Answers(_DiesBeforeFirstChunk): + def __init__(self, model: str): + super().__init__(model) + self._chunks = iter( + [litellm.ModelResponseStream(choices=[{"index": 0, "delta": {"content": f"ok-from-{model}"}}])] + ) + + def __next__(self): + return next(self._chunks) + + async def __anext__(self): + try: + return next(self._chunks) + except StopIteration: + raise StopAsyncIteration from None + + +def _primary_and_backup_router(**settings: object) -> litellm.Router: + return litellm.Router( + model_list=[ + {"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "fake-key"}}, + {"model_name": "backup", "litellm_params": {"model": "openai/backup-model", "api_key": "fake-key"}}, + ], + num_retries=0, + **settings, + ) + + +def _stream_for(**kwargs: object) -> CustomStreamWrapper: + model: Final = str(kwargs["model"]) + return _Answers(model) if "backup" in model else _DiesBeforeFirstChunk(model) + + +def _groups_called(provider_calls: MagicMock) -> list[str]: + return [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] + + +def _router_internals_reached_the_provider(provider_calls: MagicMock) -> bool: + leaked: Final = frozenset( + ("fallbacks", "context_window_fallbacks", "content_policy_fallbacks", MID_STREAM_FALLBACK_CONTROLS_KEY) + ) + return any(leaked & call.kwargs.keys() for call in provider_calls.call_args_list) + + +def test_completion_mid_stream_fallback_honors_the_per_request_list(): + router: Final = _primary_and_backup_router() + + with patch("litellm.completion", side_effect=_stream_for) as provider_calls: + response: Final = router.completion( + model="primary", + messages=[{"role": "user", "content": "hi"}], + stream=True, + fallbacks=[{"primary": ["backup"]}], + ) + content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in response if chunk is not None) + + assert content == "ok-from-openai/backup-model" + assert _groups_called(provider_calls) == ["primary", "backup"] + assert not _router_internals_reached_the_provider(provider_calls) + + +@pytest.mark.asyncio +async def test_acompletion_mid_stream_fallback_honors_the_per_request_list(): + router: Final = _primary_and_backup_router() + + async def fake_acompletion(**kwargs): + return _stream_for(**kwargs) + + with patch("litellm.acompletion", side_effect=fake_acompletion) as provider_calls: + response: Final = await router.acompletion( + model="primary", + messages=[{"role": "user", "content": "hi"}], + stream=True, + fallbacks=[{"primary": ["backup"]}], + ) + content: Final = "".join( + [chunk.choices[0].delta.content or "" async for chunk in response if chunk is not None] + ) + + assert content == "ok-from-openai/backup-model" + assert _groups_called(provider_calls) == ["primary", "backup"] + assert not _router_internals_reached_the_provider(provider_calls) + + +@pytest.mark.asyncio +async def test_acompletion_mid_stream_fallback_honors_a_per_request_fallbacks_none(): + router: Final = _primary_and_backup_router(fallbacks=[{"primary": ["backup"]}]) + + async def fake_acompletion(**kwargs): + return _stream_for(**kwargs) + + with patch("litellm.acompletion", side_effect=fake_acompletion) as provider_calls: + response: Final = await router.acompletion( + model="primary", messages=[{"role": "user", "content": "hi"}], stream=True, fallbacks=None + ) + with pytest.raises(litellm.InternalServerError, match="provider 500 from openai/primary-model"): + [chunk async for chunk in response] + + assert _groups_called(provider_calls) == ["primary"] + + def test_refusal_on_the_last_fallback_hop_is_returned_instead_of_raised(): """LIT-7400 follow-up: a refusal on the final hop of an exhausted list passes through.""" from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets @@ -4590,6 +4723,235 @@ async def test_aresponses_streaming_iterator_fallback(): assert call_kwargs["disable_fallbacks"] is False +@pytest.mark.asyncio +async def test_aresponses_mid_stream_order_fallback_hop_drops_the_encrypted_reasoning_the_next_provider_cannot_decrypt( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + """A Codex-style multi-turn history replays the order-1 provider's encrypted reasoning. When that + provider's stream breaks before its first output chunk, the order-2 hop must not replay reasoning + the next provider cannot decrypt; the readable summary stays.""" + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + def history() -> list: + return [ + {"type": "message", "role": "user", "content": "What is 17*23?"}, + { + "type": "reasoning", + "id": "rs_order1", + "encrypted_content": "gAAAAA-minted-by-order-1", + "summary": [{"type": "summary_text", "text": "multiply 17 by 23"}], + }, + {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "391"}]}, + {"type": "message", "role": "user", "content": "And 19*21?"}, + ] + + def response_body(response_id: str, model: str, status: str, output: list) -> dict: + return { + "id": response_id, + "object": "response", + "created_at": 0, + "status": status, + "model": model, + "output": output, + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2} if status == "completed" else None, + } + + def sse(events: list) -> httpx.Response: + body: Final = "".join(f"data: {json.dumps(event)}\n\n" for event in events) + return httpx.Response(200, content=body, headers={"content-type": "text/event-stream"}) + + openai_opened: Final = response_body("resp_openai", "gpt-6-astra", "in_progress", []) + openai_route: Final = respx_mock.post("https://api.openai.com/v1/responses").mock( + return_value=sse( + [ + {"type": "response.created", "sequence_number": 0, "response": openai_opened}, + {"type": "response.in_progress", "sequence_number": 1, "response": openai_opened}, + { + "type": "error", + "sequence_number": 2, + "error": { + "type": "server_error", + "code": "server_error", + "message": "The server had an error while processing your request", + "param": None, + }, + }, + ] + ) + ) + mantle_answer: Final = [ + { + "id": "msg_mantle", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "399", "annotations": []}], + } + ] + mantle_route: Final = respx_mock.post("https://bedrock-mantle.us-east-1.api.aws/openai/v1/responses").mock( + return_value=sse( + [ + { + "type": "response.created", + "sequence_number": 0, + "response": response_body("resp_mantle", "openai.gpt-6-astra", "in_progress", []), + }, + { + "type": "response.completed", + "sequence_number": 1, + "response": response_body("resp_mantle", "openai.gpt-6-astra", "completed", mantle_answer), + }, + ] + ) + ) + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-6-astra", + "litellm_params": {"model": "openai/gpt-6-astra", "api_key": "openai-key", "order": 1}, + "model_info": {"id": "openai-order-1"}, + }, + { + "model_name": "gpt-6-astra", + "litellm_params": { + "model": "bedrock_mantle/openai.gpt-6-astra", + "api_key": "mantle-bearer-token", + "aws_region_name": "us-east-1", + "order": 2, + }, + "model_info": {"id": "mantle-order-2"}, + }, + ], + num_retries=0, + ) + stream = await router.aresponses(model="gpt-6-astra", input=history(), store=False, stream=True) + collected = [event async for event in stream] + + assert [event.type for event in collected] == ["response.created", "response.completed"] + assert json.loads(openai_route.calls.last.request.read())["input"] == history() + assert json.loads(mantle_route.calls.last.request.read())["input"] == [ + {"type": "message", "role": "user", "content": "What is 17*23?"}, + {"type": "reasoning", "summary": [{"type": "summary_text", "text": "multiply 17 by 23"}]}, + {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "391"}]}, + {"type": "message", "role": "user", "content": "And 19*21?"}, + ] + + +@pytest.mark.asyncio +async def test_aresponses_mid_stream_order_fallback_hop_keeps_the_encrypted_reasoning_the_same_boundary_can_decrypt( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + """The mid-stream hop re-enters the chain on a snapshot taken before routing, so the snapshot has to + carry the deployment that streamed and failed: a same-boundary order-2 deployment can decrypt that + deployment's unmarked reasoning and must receive it unchanged.""" + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + def history() -> list: + return [ + {"type": "message", "role": "user", "content": "What is 17*23?"}, + { + "type": "reasoning", + "id": "rs_order1", + "encrypted_content": "gAAAAA-minted-by-order-1", + "summary": [{"type": "summary_text", "text": "multiply 17 by 23"}], + }, + {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "391"}]}, + {"type": "message", "role": "user", "content": "And 19*21?"}, + ] + + def response_body(response_id: str, model: str, status: str, output: list) -> dict: + return { + "id": response_id, + "object": "response", + "created_at": 0, + "status": status, + "model": model, + "output": output, + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2} if status == "completed" else None, + } + + def sse(events: list) -> httpx.Response: + body: Final = "".join(f"data: {json.dumps(event)}\n\n" for event in events) + return httpx.Response(200, content=body, headers={"content-type": "text/event-stream"}) + + order_1_opened: Final = response_body("resp_order1", "gpt-6-astra", "in_progress", []) + order_2_answer: Final = [ + { + "id": "msg_order2", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "399", "annotations": []}], + } + ] + openai_route: Final = respx_mock.post("https://api.openai.com/v1/responses").mock( + side_effect=[ + sse( + [ + {"type": "response.created", "sequence_number": 0, "response": order_1_opened}, + {"type": "response.in_progress", "sequence_number": 1, "response": order_1_opened}, + { + "type": "error", + "sequence_number": 2, + "error": { + "type": "server_error", + "code": "server_error", + "message": "The server had an error while processing your request", + "param": None, + }, + }, + ] + ), + sse( + [ + { + "type": "response.created", + "sequence_number": 0, + "response": response_body("resp_order2", "gpt-6-astra-mini", "in_progress", []), + }, + { + "type": "response.completed", + "sequence_number": 1, + "response": response_body("resp_order2", "gpt-6-astra-mini", "completed", order_2_answer), + }, + ] + ), + ] + ) + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-6-astra", + "litellm_params": { + "model": "openai/gpt-6-astra", + "api_base": "https://api.openai.com/v1", + "api_key": "openai-key", + "order": 1, + }, + "model_info": {"id": "openai-order-1"}, + }, + { + "model_name": "gpt-6-astra", + "litellm_params": { + "model": "openai/gpt-6-astra-mini", + "api_base": "https://api.openai.com/v1", + "api_key": "openai-key", + "order": 2, + }, + "model_info": {"id": "openai-order-2"}, + }, + ], + num_retries=0, + ) + stream = await router.aresponses(model="gpt-6-astra", input=history(), store=False, stream=True) + collected = [event async for event in stream] + + assert [event.type for event in collected] == ["response.created", "response.completed"] + assert [json.loads(call.request.read())["input"] for call in openai_route.calls] == [history(), history()] + + @pytest.mark.asyncio async def test_aresponses_streaming_content_policy_error_event_routes_to_content_policy_fallback(): """Regression: a mid-stream content_policy_violation error event never reached @@ -24593,3 +24955,88 @@ class CompletionCustomHandler( except Exception: print(f"Assertion Error: {traceback.format_exc()}") self.errors.append(traceback.format_exc()) + + +_FALLBACK_WIRE_PRIMARY: Final = "http://primary.wire.test/v1" +_FALLBACK_WIRE_BACKUP: Final = "http://backup.wire.test/v1" +_FALLBACK_WIRE_BACKUP_REPLY: Final = { + "id": "chatcmpl-backup", + "object": "chat.completion", + "created": 1, + "model": "gpt-5.6", + "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "pong"}}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, +} + + +def _fallback_wire_router() -> Router: + return Router( + model_list=[ + { + "model_name": "primary", + "litellm_params": { + "model": "openai/gpt-5.6", + "api_key": "sk-primary", + "api_base": _FALLBACK_WIRE_PRIMARY, + "max_retries": 0, + }, + }, + { + "model_name": "backup", + "litellm_params": { + "model": "openai/gpt-5.6", + "api_key": "sk-backup", + "api_base": _FALLBACK_WIRE_BACKUP, + "max_retries": 0, + }, + }, + ], + fallbacks=[{"primary": ["backup"]}], + num_retries=0, + ) + + +def _mock_fallback_wire(respx_mock: respx.MockRouter) -> tuple[respx.Route, respx.Route]: + primary = respx_mock.post(f"{_FALLBACK_WIRE_PRIMARY}/chat/completions").mock( + return_value=httpx.Response(500, json={"error": {"message": "primary overloaded", "type": "server_error"}}) + ) + backup = respx_mock.post(f"{_FALLBACK_WIRE_BACKUP}/chat/completions").mock( + return_value=httpx.Response(200, json=_FALLBACK_WIRE_BACKUP_REPLY) + ) + return primary, backup + + +def _assert_fallback_errors_reached_the_caller_and_not_the_wire( + response: object, primary: respx.Route, backup: respx.Route +) -> None: + for route in (primary, backup): + assert route.called + for call in route.calls: + assert "include_fallback_errors" not in json.loads(call.request.content) + assert isinstance(response, litellm.ModelResponse) + assert response.choices[0].message.content == "pong" + headers = response._hidden_params["additional_headers"] + assert headers["x-litellm-attempted-fallbacks"] == 1 + errors = json.loads(headers["x-litellm-fallback-errors"]) + assert len(errors) == 1 + assert "primary overloaded" in errors[0]["message"] + + +def test_sync_completion_keeps_include_fallback_errors_off_the_wire_and_returns_the_errors(): + with respx.mock(assert_all_called=True) as respx_mock: + primary, backup = _mock_fallback_wire(respx_mock) + response = _fallback_wire_router().completion( + model="primary", messages=[{"role": "user", "content": "hi"}], include_fallback_errors=True + ) + _assert_fallback_errors_reached_the_caller_and_not_the_wire(response, primary, backup) + + +@pytest.mark.asyncio +async def test_acompletion_keeps_include_fallback_errors_off_the_wire_and_returns_the_errors(monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + with respx.mock(assert_all_called=True) as respx_mock: + primary, backup = _mock_fallback_wire(respx_mock) + response = await _fallback_wire_router().acompletion( + model="primary", messages=[{"role": "user", "content": "hi"}], include_fallback_errors=True + ) + _assert_fallback_errors_reached_the_caller_and_not_the_wire(response, primary, backup) diff --git a/tests/unit/test_router/test_router_callback_hook_sequence.py b/tests/unit/test_router/test_router_callback_hook_sequence.py new file mode 100644 index 00000000000..a7c41964950 --- /dev/null +++ b/tests/unit/test_router/test_router_callback_hook_sequence.py @@ -0,0 +1,509 @@ +import asyncio +import inspect +import json +from collections import Counter +from collections.abc import Callable, Mapping, Sequence +from datetime import datetime +from typing import Final, Literal, NamedTuple, TypedDict + +from typing_extensions import ReadOnly + +import httpx +import pytest +import respx +from pydantic import TypeAdapter + +import litellm +from litellm.caching.caching import Cache +from litellm.integrations.custom_logger import CustomLogger +from litellm.router import Router + +_PRIMARY: Final = "https://hooks-primary.openai.azure.com" +_FALLBACK: Final = "https://hooks-fallback.openai.azure.com" +_API_VERSION: Final = "2024-10-21" +_MESSAGES: Final = [{"role": "user", "content": "Hi - i'm openai"}] +_USAGE: Final = {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7} +_COMPLETION: Final = { + "id": "chatcmpl-hooks", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4.1-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hello"}, "finish_reason": "stop"}], + "usage": _USAGE, +} +_EMBEDDING_VECTOR: Final = [0.1, 0.2, 0.3] +_EMBEDDING: Final = { + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": _EMBEDDING_VECTOR}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 3, "total_tokens": 3}, +} +_STREAM: Final = ( + "".join( + f"data: {json.dumps(chunk)}\n\n" + for chunk in ( + { + "id": "chatcmpl-hooks", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": "gpt-4.1-mini", + "choices": [{"index": 0, "delta": {"role": "assistant", "content": "hel"}, "finish_reason": None}], + }, + { + "id": "chatcmpl-hooks", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": "gpt-4.1-mini", + "choices": [{"index": 0, "delta": {"content": "lo"}, "finish_reason": "stop"}], + }, + ) + ) + + "data: [DONE]\n\n" +) +_AUTH_ERROR: Final = httpx.Response( + 401, + json={ + "error": {"message": "Incorrect API key provided", "type": "invalid_request_error", "code": "invalid_api_key"} + }, +) + +_OUR_MODEL_GROUPS: Final = frozenset({"hooks-group", "primary-group", "fallback-group"}) + +_State = Literal[ + "sync_pre_api_call", + "post_api_call", + "async_stream", + "sync_success", + "async_success", + "sync_failure", + "async_failure", +] + + +class _HookEvent(NamedTuple): + state: _State + model: object + kwargs: Mapping[str, object] + response: object + + +def _router_context_problems(kwargs: Mapping[str, object]) -> tuple[str, ...]: + litellm_params: Final = kwargs.get("litellm_params") + if not isinstance(litellm_params, dict): + return ("litellm_params",) + metadata: Final = litellm_params.get("metadata") + model_info: Final = litellm_params.get("model_info") + checks: Final = { + "metadata": isinstance(metadata, dict), + "model_group": isinstance(metadata, dict) and isinstance(metadata.get("model_group"), str), + "deployment": isinstance(metadata, dict) and isinstance(metadata.get("deployment"), str), + "model_info": isinstance(model_info, dict), + "model_info id": isinstance(model_info, dict) and isinstance(model_info.get("id"), str), + "proxy_server_request": isinstance(litellm_params.get("proxy_server_request"), (str, type(None))), + "preset_cache_key": isinstance(litellm_params.get("preset_cache_key"), (str, type(None))), + "stream_response": isinstance(litellm_params.get("stream_response"), dict), + } + return tuple(name for name, ok in checks.items() if not ok) + + +def _request_problems(kwargs: Mapping[str, object]) -> tuple[str, ...]: + checks: Final = { + "model": isinstance(kwargs.get("model"), str), + "messages": isinstance(kwargs.get("messages"), list), + "optional_params": isinstance(kwargs.get("optional_params"), dict), + "start_time": isinstance(kwargs.get("start_time"), (datetime, type(None))), + "stream": isinstance(kwargs.get("stream"), bool), + "user": isinstance(kwargs.get("user"), (str, type(None))), + } + return (*(name for name, ok in checks.items() if not ok), *_router_context_problems(kwargs)) + + +def _call_detail_problems(kwargs: Mapping[str, object]) -> tuple[str, ...]: + original_response: Final = kwargs.get("original_response") + checks: Final = { + "input": isinstance(kwargs.get("input"), (list, dict, str)), + "api_key": isinstance(kwargs.get("api_key"), (str, type(None))), + "original_response": isinstance(original_response, (str, litellm.CustomStreamWrapper, type(None))) + or inspect.iscoroutine(original_response) + or inspect.isasyncgen(original_response), + "additional_args": isinstance(kwargs.get("additional_args"), (dict, type(None))), + "log_event_type": isinstance(kwargs.get("log_event_type"), str), + } + return tuple(name for name, ok in checks.items() if not ok) + + +def _is_from_this_test(kwargs: Mapping[str, object]) -> bool: + litellm_params: Final = kwargs.get("litellm_params") + metadata: Final = litellm_params.get("metadata") if isinstance(litellm_params, dict) else None + return isinstance(metadata, dict) and metadata.get("model_group") in _OUR_MODEL_GROUPS + + +class _HookRecorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.events: Final[list[_HookEvent]] = [] + self.errors: Final[list[str]] = [] + self.loop: asyncio.AbstractEventLoop | None = None + self.waiters: Final[list[tuple[Callable[[Sequence[_State]], bool], asyncio.Event]]] = [] + + @property + def states(self) -> list[_State]: + return [event.state for event in self.events] + + def _record( + self, state: _State, model: object, kwargs: Mapping[str, object], response: object, problems: Sequence[str] + ) -> None: + if not _is_from_this_test(kwargs): + return + self.errors.extend(f"{state}: {problem}" for problem in problems) + self.events.append(_HookEvent(state, model, kwargs, response)) + if self.loop is not None: + self.loop.call_soon_threadsafe(self._notify) + + def _notify(self) -> None: + for predicate, event in self.waiters: + if predicate(tuple(self.states)): + event.set() + + async def until(self, predicate: Callable[[Sequence[_State]], bool]) -> tuple[_State, ...]: + self.loop = asyncio.get_running_loop() + if not predicate(tuple(self.states)): + event: Final = asyncio.Event() + self.waiters.append((predicate, event)) + await asyncio.wait_for(event.wait(), timeout=10) + return tuple(self.states) + + def log_pre_api_call(self, model: object, messages: object, kwargs: Mapping[str, object]) -> None: + problems: Final = ( + *(("model",) if not isinstance(model, str) else ()), + *(("messages",) if not isinstance(messages, list) else ()), + *_request_problems(kwargs), + ) + self._record("sync_pre_api_call", model, kwargs, messages, problems) + + def log_post_api_call( + self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + problems: Final = ( + *(("start_time",) if not isinstance(start_time, datetime) else ()), + *(("end_time",) if end_time is not None else ()), + *(("response_obj",) if response_obj is not None else ()), + *_request_problems(kwargs), + *_call_detail_problems(kwargs), + ) + self._record("post_api_call", kwargs.get("model"), kwargs, response_obj, problems) + + def log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + self._record("sync_success", kwargs.get("model"), kwargs, response_obj, ()) + + def log_failure_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + self._record("sync_failure", kwargs.get("model"), kwargs, response_obj, ()) + + async def async_log_stream_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + self._record("async_stream", kwargs.get("model"), kwargs, response_obj, ()) + + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + problems: Final = ( + *(("times",) if not (isinstance(start_time, datetime) and isinstance(end_time, datetime)) else ()), + *( + ("response_obj",) + if not isinstance(response_obj, (litellm.ModelResponse, litellm.EmbeddingResponse)) + else () + ), + *(("cache_hit",) if not isinstance(kwargs.get("cache_hit"), (bool, type(None))) else ()), + *_request_problems(kwargs), + *_call_detail_problems(kwargs), + ) + self._record("async_success", kwargs.get("model"), kwargs, response_obj, problems) + + async def async_log_failure_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + problems: Final = ( + *(("times",) if not (isinstance(start_time, datetime) and isinstance(end_time, datetime)) else ()), + *(("response_obj",) if response_obj is not None else ()), + *(("exception",) if not isinstance(kwargs.get("exception"), Exception) else ()), + *_request_problems(kwargs), + *_call_detail_problems(kwargs), + ) + self._record("async_failure", kwargs.get("model"), kwargs, response_obj, problems) + + +def _settled(terminal: int, posts: int) -> Callable[[Sequence[_State]], bool]: + def satisfied(states: Sequence[_State]) -> bool: + terminals: Final = sum(state in ("async_success", "async_failure") for state in states) + return terminals >= terminal and states.count("post_api_call") >= posts + + return satisfied + + +@pytest.fixture +def recorder(monkeypatch: pytest.MonkeyPatch) -> _HookRecorder: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + hook_recorder: Final = _HookRecorder() + monkeypatch.setattr(litellm, "callbacks", [hook_recorder]) + return hook_recorder + + +def _router(model: str, api_base: str = _PRIMARY) -> Router: + return Router( + model_list=[ + { + "model_name": "hooks-group", + "litellm_params": { + "model": model, + "api_key": "sk-unit-test", + "api_base": api_base, + "api_version": _API_VERSION, + }, + "model_info": {"base_model": model}, + } + ], + num_retries=0, + ) + + +class _RouterMetadata(TypedDict): + model_group: ReadOnly[str] + deployment: ReadOnly[str] + + +class _RouterModelInfo(TypedDict): + id: ReadOnly[str] + + +class _RouterParams(TypedDict): + metadata: ReadOnly[_RouterMetadata] + model_info: ReadOnly[_RouterModelInfo] + + +def _router_params(event: _HookEvent) -> _RouterParams: + return TypeAdapter(_RouterParams).validate_python(event.kwargs["litellm_params"]) + + +def _model_group(event: _HookEvent) -> str: + return _router_params(event)["metadata"]["model_group"] + + +def _model_id(event: _HookEvent) -> str: + return _router_params(event)["model_info"]["id"] + + +def _of_state(recorder: _HookRecorder, state: _State) -> list[_HookEvent]: + return [event for event in recorder.events if event.state == state] + + +def _completion(event: _HookEvent) -> litellm.ModelResponse: + assert isinstance(event.response, litellm.ModelResponse) + return event.response + + +@pytest.mark.asyncio +async def test_router_chat_success_streaming_and_failure_fire_the_hooks_in_order( + recorder: _HookRecorder, respx_mock: respx.MockRouter +) -> None: + route: Final = respx_mock.post(url__startswith=f"{_PRIMARY}/openai/deployments/gpt-4.1-mini/chat/completions") + router: Final = _router("azure/gpt-4.1-mini") + + route.mock(return_value=httpx.Response(200, json=_COMPLETION)) + await router.acompletion(model="hooks-group", messages=_MESSAGES) + await recorder.until(_settled(1, posts=1)) + assert recorder.states == ["sync_pre_api_call", "post_api_call", "async_success"] + pre, post, success = recorder.events + assert pre.model == "gpt-4.1-mini" + assert pre.response == _MESSAGES + assert {_model_group(event) for event in recorder.events} == {"hooks-group"} + assert len({_model_id(event) for event in recorder.events}) == 1 + assert post.kwargs["messages"] == _MESSAGES + assert success.kwargs["stream"] is False + assert _completion(success).choices[0].message.content == "hello" + assert _completion(success).usage.total_tokens == _USAGE["total_tokens"] + + route.mock(return_value=httpx.Response(200, text=_STREAM, headers={"content-type": "text/event-stream"})) + stream: Final = await router.acompletion(model="hooks-group", messages=_MESSAGES, stream=True) + assert "".join([chunk.choices[0].delta.content or "" async for chunk in stream]) == "hello" + await recorder.until(_settled(2, posts=2)) + streamed: Final = recorder.events[3:] + assert sorted(event.state for event in streamed[:2]) == ["post_api_call", "sync_pre_api_call"] + assert [event.state for event in streamed[2:]] == ["async_success"] + assert all(event.kwargs["stream"] is True for event in streamed) + assert _completion(streamed[2]).choices[0].message.content == "hello" + + route.mock(return_value=_AUTH_ERROR) + with pytest.raises(litellm.AuthenticationError): + await router.acompletion(model="hooks-group", messages=_MESSAGES) + await recorder.until(_settled(3, posts=3)) + failed: Final = recorder.events[6:] + assert [event.state for event in failed] == ["sync_pre_api_call", "post_api_call", "async_failure"] + assert failed[2].response is None + assert isinstance(failed[2].kwargs["exception"], litellm.AuthenticationError) + assert recorder.errors == [] + + +@pytest.mark.asyncio +async def test_router_embedding_success_and_failure_fire_the_hooks_in_order( + recorder: _HookRecorder, respx_mock: respx.MockRouter +) -> None: + route: Final = respx_mock.post(url__startswith=f"{_PRIMARY}/openai/deployments/text-embedding-3-small/embeddings") + router: Final = _router("azure/text-embedding-3-small") + + route.mock(return_value=httpx.Response(200, json=_EMBEDDING)) + await router.aembedding(model="hooks-group", input=["hello"]) + await recorder.until(_settled(1, posts=1)) + assert recorder.states == ["sync_pre_api_call", "post_api_call", "async_success"] + assert {event.model for event in recorder.events} == {"text-embedding-3-small"} + assert {_model_group(event) for event in recorder.events} == {"hooks-group"} + embedding: Final = recorder.events[2].response + assert isinstance(embedding, litellm.EmbeddingResponse) + assert embedding.model_dump()["data"][0]["embedding"] == _EMBEDDING_VECTOR + assert embedding.usage.prompt_tokens == 3 + + route.mock(return_value=_AUTH_ERROR) + with pytest.raises(litellm.AuthenticationError): + await router.aembedding(model="hooks-group", input=["hello"]) + await recorder.until(_settled(2, posts=2)) + assert recorder.states[3:] == ["sync_pre_api_call", "post_api_call", "async_failure"] + assert recorder.events[5].response is None + assert isinstance(recorder.events[5].kwargs["exception"], litellm.AuthenticationError) + assert recorder.errors == [] + + +@pytest.mark.asyncio +async def test_router_fallback_fires_failure_then_success_hooks( + recorder: _HookRecorder, respx_mock: respx.MockRouter +) -> None: + respx_mock.post(url__startswith=f"{_PRIMARY}/openai/deployments/gpt-4.1-mini/chat/completions").mock( + return_value=_AUTH_ERROR + ) + fallback: Final = respx_mock.post( + url__startswith=f"{_FALLBACK}/openai/deployments/gpt-4.1-mini/chat/completions" + ).mock(return_value=httpx.Response(200, json=_COMPLETION)) + router: Final = Router( + model_list=[ + { + "model_name": "primary-group", + "litellm_params": { + "model": "azure/gpt-4.1-mini", + "api_key": "my-bad-key", + "api_base": _PRIMARY, + "api_version": _API_VERSION, + }, + }, + { + "model_name": "fallback-group", + "litellm_params": { + "model": "azure/gpt-4.1-mini", + "api_key": "sk-unit-test", + "api_base": _FALLBACK, + "api_version": _API_VERSION, + }, + }, + ], + fallbacks=[{"primary-group": ["fallback-group"]}], + num_retries=0, + ) + + await router.acompletion(model="primary-group", messages=_MESSAGES) + await recorder.until(_settled(2, posts=2)) + + assert fallback.call_count == 1 + assert recorder.states == [ + "sync_pre_api_call", + "post_api_call", + "async_failure", + "sync_pre_api_call", + "post_api_call", + "async_success", + ] + assert [_model_group(event) for event in recorder.events] == ["primary-group"] * 3 + ["fallback-group"] * 3 + assert _model_id(recorder.events[0]) != _model_id(recorder.events[3]) + assert isinstance(recorder.events[2].kwargs["exception"], litellm.AuthenticationError) + assert _completion(recorder.events[5]).choices[0].message.content == "hello" + assert recorder.errors == [] + + +@pytest.mark.asyncio +async def test_router_completion_cache_hit_fires_a_second_success_hook( + recorder: _HookRecorder, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(litellm, "cache", Cache()) + route: Final = respx_mock.post(url__startswith=f"{_PRIMARY}/openai/deployments/gpt-4.1-mini/chat/completions").mock( + return_value=httpx.Response(200, json=_COMPLETION) + ) + router: Final = _router("azure/gpt-4.1-mini") + + await router.acompletion(model="hooks-group", messages=_MESSAGES, caching=True) + await recorder.until(_settled(1, posts=1)) + await router.acompletion(model="hooks-group", messages=_MESSAGES, caching=True) + await recorder.until(_settled(2, posts=1)) + + assert route.call_count == 1 + assert recorder.states == ["sync_pre_api_call", "post_api_call", "async_success", "async_success"] + first, second = _of_state(recorder, "async_success") + assert first.kwargs.get("cache_hit") is not True + assert second.kwargs.get("cache_hit") is True + assert _completion(second).choices[0].message.content == "hello" + assert recorder.errors == [] + + +@pytest.mark.asyncio +async def test_router_streaming_cache_hit_still_fires_the_success_hook( + recorder: _HookRecorder, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(litellm, "cache", Cache()) + route: Final = respx_mock.post(url__startswith=f"{_PRIMARY}/openai/deployments/gpt-4.1-mini/chat/completions").mock( + return_value=httpx.Response(200, text=_STREAM, headers={"content-type": "text/event-stream"}) + ) + router: Final = _router("azure/gpt-4.1-mini") + + first: Final = await router.acompletion(model="hooks-group", messages=_MESSAGES, stream=True, caching=True) + first_text: Final = "".join([chunk.choices[0].delta.content or "" async for chunk in first]) + await recorder.until(_settled(1, posts=1)) + states_after_first: Final = len(recorder.states) + second: Final = await router.acompletion(model="hooks-group", messages=_MESSAGES, stream=True, caching=True) + second_text: Final = "".join([chunk.choices[0].delta.content or "" async for chunk in second]) + await recorder.until(_settled(2, posts=1)) + + assert route.call_count == 1 + assert first_text == second_text == "hello" + assert sorted(recorder.states[:2]) == ["post_api_call", "sync_pre_api_call"] + assert recorder.states[2:states_after_first] == ["async_success"] + assert recorder.states[states_after_first:] == ["async_success"] + first_success, second_success = _of_state(recorder, "async_success") + assert first_success.kwargs.get("cache_hit") is not True + assert second_success.kwargs.get("cache_hit") is True + assert _completion(second_success).choices[0].message.content == "hello" + assert recorder.errors == [] + + +@pytest.mark.asyncio +async def test_router_embedding_cache_hit_fires_a_second_success_hook( + recorder: _HookRecorder, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(litellm, "cache", Cache()) + route: Final = respx_mock.post( + url__startswith=f"{_PRIMARY}/openai/deployments/text-embedding-3-small/embeddings" + ).mock(return_value=httpx.Response(200, json=_EMBEDDING)) + router: Final = _router("azure/text-embedding-3-small") + + await router.aembedding(model="hooks-group", input=["hello"], caching=True) + await recorder.until(_settled(1, posts=1)) + await router.aembedding(model="hooks-group", input=["hello"], caching=True) + await recorder.until(_settled(2, posts=1)) + + assert route.call_count == 1 + assert recorder.states == ["sync_pre_api_call", "post_api_call", "async_success", "async_success"] + first, second = _of_state(recorder, "async_success") + assert first.kwargs.get("cache_hit") is not True + assert second.kwargs.get("cache_hit") is True + assert isinstance(second.response, litellm.EmbeddingResponse) + assert second.response.model_dump()["data"][0]["embedding"] == _EMBEDDING_VECTOR + assert recorder.errors == [] diff --git a/tests/unit/test_router/test_router_provider_endpoints.py b/tests/unit/test_router/test_router_provider_endpoints.py new file mode 100644 index 00000000000..708322f4fec --- /dev/null +++ b/tests/unit/test_router/test_router_provider_endpoints.py @@ -0,0 +1,732 @@ +import asyncio +import contextlib +import io +import json +import uuid +from collections.abc import Callable, Iterator +from datetime import datetime, timezone +from typing import Final + +import httpx +import litellm +import pytest +import respx +from litellm import CustomLogger, Router +from litellm.caching.dual_cache import DualCache +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.constants import ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.router_utils.router_callbacks.track_deployment_metrics import ( + get_deployment_successes_for_current_minute, +) +from litellm.types.caching import RedisPipelineIncrementOperation +from litellm.types.llms.openai import HttpxBinaryResponseContent +from litellm.types.router import RoutingStrategy +from litellm.types.utils import ImageResponse, ModelResponse, RerankResponse +from openai import AsyncAzureOpenAI, AsyncOpenAI + +EVENT_TIMEOUT_SECONDS: Final = 5 +ROUTING_SELECTIONS: Final = 40 +ROUTING_MESSAGES: Final = ({"role": "user", "content": "route this request"},) +EXPENSIVE_COSTS: Final = {"input_cost_per_token": 1.0, "output_cost_per_token": 1.0} +CHEAP_COSTS: Final = {"input_cost_per_token": 1e-9, "output_cost_per_token": 1e-9} +CHAT_RESPONSE: Final = { + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 5, "completion_tokens": 1, "total_tokens": 6}, +} + + +@pytest.fixture(autouse=True) +def httpx_transport(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.fixture +def router_minute_pinned(monkeypatch: pytest.MonkeyPatch) -> None: + pinned: Final = datetime(2026, 1, 1, 12, 0, 30, tzinfo=timezone.utc) + monkeypatch.setattr("litellm.router.get_utc_datetime", lambda: pinned) + + +class _RouterLoggingCapture(CustomLogger): + def __init__(self, model_id: str) -> None: + super().__init__() + self.model_id: Final = model_id + self.success_events: asyncio.Queue[tuple[object | None, object | None]] = asyncio.Queue() + + async def async_log_success_event( + self, + kwargs: dict[str, object], + response_obj: object, + start_time: object, + end_time: object, + ) -> None: + standard_logging_object: Final = kwargs.get("standard_logging_object") + if not isinstance(standard_logging_object, dict) or standard_logging_object.get("model_id") != self.model_id: + return + self.success_events.put_nowait((kwargs.get("client"), standard_logging_object)) + + async def next_event(self) -> tuple[object | None, object | None]: + return await asyncio.wait_for(self.success_events.get(), timeout=EVENT_TIMEOUT_SECONDS) + + +def _rpm_tpm_router(model_id: str) -> Router: + return Router( + model_list=[ + { + "model_name": "gpt-5-mini", + "litellm_params": {"model": "gpt-5-mini", "api_key": "sk-fake", "tpm": 1000, "rpm": 100}, + "model_info": {"id": model_id}, + } + ] + ) + + +def _ratelimit_headers(response: ModelResponse | CustomStreamWrapper) -> dict[str, int]: + return {k: v for k, v in response._hidden_params["additional_headers"].items() if k.startswith("x-ratelimit-")} + + +def _openai_router_client() -> AsyncOpenAI: + return AsyncOpenAI(api_key="sk-fake") + + +def _azure_router_client() -> AsyncAzureOpenAI: + return AsyncAzureOpenAI(api_key="sk-fake", azure_endpoint="https://azure.test", api_version="2025-02-01-preview") + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("litellm_params", "build_client", "route_url"), + [ + ( + {"model": "whisper-1", "api_key": "sk-fake"}, + _openai_router_client, + "https://api.openai.com/v1/audio/transcriptions", + ), + ( + { + "model": "azure/whisper", + "api_base": "https://azure.test", + "api_key": "sk-fake", + "api_version": "2025-02-01-preview", + }, + _azure_router_client, + "https://azure.test/openai/deployments/whisper/audio/transcriptions?api-version=2025-02-01-preview", + ), + ], + ids=["openai", "azure"], +) +async def test_router_transcription_reuses_router_level_client_for_each_deployment( + litellm_params: dict[str, str], + build_client: Callable[[], AsyncOpenAI], + route_url: str, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + model_id: Final = f"whisper-{uuid.uuid4().hex}" + capture: Final = _RouterLoggingCapture(model_id) + monkeypatch.setattr(litellm, "callbacks", [capture]) + router: Final = Router( + model_list=[{"model_name": "whisper", "litellm_params": dict(litellm_params), "model_info": {"id": model_id}}] + ) + router_level_client: Final = build_client() + router.cache.set_cache(key=f"{model_id}_async_client", value=router_level_client, local_only=True) + route: Final = respx_mock.post(route_url).respond(200, json={"text": "hello"}) + + response: Final = await router.atranscription( + model="whisper", file=("speech.wav", io.BytesIO(b"offline audio"), "audio/wav") + ) + public_client, public_logging = await capture.next_event() + internal_response: Final = await router._atranscription( + model="whisper", file=("speech.wav", io.BytesIO(b"offline audio"), "audio/wav") + ) + internal_client, _ = await capture.next_event() + upstream_requests: Final = tuple(call.request for call in route.calls) + + assert public_client == str(router_level_client) + assert internal_client == str(router_level_client) + assert tuple(str(request.url) for request in upstream_requests) == (route_url, route_url) + assert all( + request.headers.get("content-type", "").startswith("multipart/form-data") + and b'filename="speech.wav"' in request.content + and b"offline audio" in request.content + for request in upstream_requests + ) + assert isinstance(public_logging, dict) + assert public_logging.get("model_group") == "whisper" + assert response.text == "hello" + assert internal_response.text == "hello" + + +@pytest.mark.asyncio +async def test_router_speech_returns_binary_content_and_logs_model_group( + respx_mock: respx.MockRouter, + monkeypatch: pytest.MonkeyPatch, +) -> None: + model_id: Final = f"tts-{uuid.uuid4().hex}" + capture: Final = _RouterLoggingCapture(model_id) + monkeypatch.setattr(litellm, "callbacks", [capture]) + router: Final = Router( + model_list=[ + { + "model_name": "tts", + "litellm_params": {"model": "openai/tts-1", "api_key": "sk-fake"}, + "model_info": {"id": model_id}, + } + ] + ) + route: Final = respx_mock.post("https://api.openai.com/v1/audio/speech").respond(200, content=b"audio") + + response: Final = await router.aspeech(model="tts", input="hello", voice="alloy") + _, standard_logging_object = await capture.next_event() + + assert route.call_count == 1 + assert json.loads(route.calls[0].request.content) == {"model": "tts-1", "input": "hello", "voice": "alloy"} + assert isinstance(response, HttpxBinaryResponseContent) + assert response.content == b"audio" + assert isinstance(standard_logging_object, dict) + assert standard_logging_object["model_group"] == "tts" + + +@pytest.mark.asyncio +async def test_router_rerank_returns_valid_response_from_public_and_underlying_calls( + respx_mock: respx.MockRouter, +) -> None: + router: Final = Router( + model_list=[ + { + "model_name": "cohere-rerank", + "litellm_params": {"model": "cohere/rerank-english-v3.0", "api_key": "sk-fake"}, + } + ] + ) + route: Final = respx_mock.post("https://api.cohere.com/v2/rerank").respond( + 200, + json={ + "id": "rerank-1", + "results": [{"index": 0, "relevance_score": 0.9}], + "meta": {"api_version": {"version": "2"}}, + }, + ) + + public_response: Final = await router.arerank( + model="cohere-rerank", + query="hello", + documents=["hello", "world"], + top_n=1, + ) + underlying_response: Final = await router._arerank( + model="cohere-rerank", + query="hello", + documents=["hello", "world"], + top_n=1, + ) + + assert route.call_count == 2 + request_bodies: Final = tuple(json.loads(call.request.content) for call in route.calls) + assert all( + body["model"] == "rerank-english-v3.0" + and body["query"] == "hello" + and body["documents"] == ["hello", "world"] + and body["top_n"] == 1 + for body in request_bodies + ) + public_validated: Final = RerankResponse.model_validate(public_response) + assert public_validated.id == "rerank-1" + assert public_validated.results[0]["relevance_score"] == 0.9 + underlying_validated: Final = RerankResponse.model_validate(underlying_response) + assert underlying_validated.id == "rerank-1" + assert underlying_validated.results[0]["relevance_score"] == 0.9 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "expected_model", "expected_api_key"), + [ + ("omni-moderation-latest", "omni-moderation-latest", "sk-catch-all"), + ("openai/omni-moderation-latest", "omni-moderation-latest", "sk-openai-wildcard"), + (None, None, "sk-env"), + ], +) +async def test_router_moderation_routes_through_wildcard_deployments( + model: str | None, + expected_model: str | None, + expected_api_key: str, + respx_mock: respx.MockRouter, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("OPENAI_API_KEY", "sk-env") + router: Final = Router( + model_list=[ + { + "model_name": "openai/*", + "litellm_params": {"model": "openai/*", "api_key": "sk-openai-wildcard"}, + }, + { + "model_name": "*", + "litellm_params": {"model": "openai/*", "api_key": "sk-catch-all"}, + }, + ] + ) + route: Final = respx_mock.post("https://api.openai.com/v1/moderations").respond( + 200, + json={ + "id": "modr-1", + "model": "omni-moderation-latest", + "results": [{"flagged": False, "categories": {}, "category_scores": {}}], + }, + ) + + response: Final = await router.amoderation(model=model, input="hello") + + assert route.call_count == 1 + upstream_request: Final = route.calls[0].request + expected_body: Final = {"input": "hello"} if expected_model is None else {"input": "hello", "model": expected_model} + assert json.loads(upstream_request.content) == expected_body + assert upstream_request.headers["authorization"] == f"Bearer {expected_api_key}" + assert response.id == "modr-1" + assert response.model == "omni-moderation-latest" + assert response.results[0].flagged is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sync_mode", [True, False]) +async def test_router_image_generation_returns_valid_image_response( + sync_mode: bool, + respx_mock: respx.MockRouter, +) -> None: + router: Final = Router( + model_list=[ + { + "model_name": "gpt-image-1", + "litellm_params": {"model": "openai/gpt-image-1", "api_key": "sk-fake"}, + } + ] + ) + route: Final = respx_mock.post("https://api.openai.com/v1/images/generations").respond( + 200, + json={"created": 1700000000, "data": [{"url": "https://images.test/result.png"}]}, + ) + + response: Final = ( + router._image_generation(model="gpt-image-1", prompt="a cat") + if sync_mode + else await router._aimage_generation(model="gpt-image-1", prompt="a cat") + ) + + assert route.call_count == 1 + request_body: Final = json.loads(route.calls[0].request.content) + assert request_body["model"] == "gpt-image-1" + assert request_body["prompt"] == "a cat" + validated: Final = ImageResponse.model_validate(response) + assert validated.data[0].url == "https://images.test/result.png" + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("router_minute_pinned") +async def test_router_acompletion_headers_read_post_increment_counter_and_count_once( + monkeypatch: pytest.MonkeyPatch, +) -> None: + capture: Final = _RouterLoggingCapture("lit-3058-async") + monkeypatch.setattr(litellm, "callbacks", [capture]) + router: Final = _rpm_tpm_router("lit-3058-async") + + response: Final = await router.acompletion( + model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong" + ) + total_tokens: Final = response.usage.total_tokens + headers: Final = _ratelimit_headers(response) + + assert total_tokens > 0 + assert headers["x-ratelimit-remaining-tokens"] == 1000 - total_tokens + assert headers["x-ratelimit-remaining-requests"] == 99 + assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1) + + await capture.next_event() + + assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("router_minute_pinned") +async def test_router_stream_counts_request_before_headers_and_tokens_once_on_completion( + monkeypatch: pytest.MonkeyPatch, +) -> None: + capture: Final = _RouterLoggingCapture("lit-3058-stream") + monkeypatch.setattr(litellm, "callbacks", [capture]) + router: Final = _rpm_tpm_router("lit-3058-stream") + stream: Final = await router.acompletion( + model="gpt-5-mini", + messages=[{"role": "user", "content": "hi"}], + mock_response="pong pong pong", + stream=True, + stream_options={"include_usage": True}, + ) + headers: Final = _ratelimit_headers(stream) + + assert headers["x-ratelimit-remaining-tokens"] == 1000 + assert headers["x-ratelimit-remaining-requests"] == 99 + assert await router.get_model_group_usage("gpt-5-mini") == (0, 1) + + chunks: Final = [chunk async for chunk in stream] + total_tokens: Final = chunks[-1].usage.total_tokens + await capture.next_event() + + assert total_tokens > 0 + assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1) + + +def test_router_validate_fallbacks_accepts_well_formed_and_rejects_malformed_entries() -> None: + router: Final = Router(model_list=[]) + + assert router.validate_fallbacks([{"gpt-5.5": ["gpt-5-mini"]}, {"gpt-5-mini": ["gpt-5.5"]}]) is None + with pytest.raises(ValueError, match="must have exactly one key"): + router.validate_fallbacks([{"primary": "fallback", "other": "fallback"}]) + with pytest.raises(ValueError, match="is not a dictionary"): + router.validate_fallbacks(["primary"]) + + +def _routing_deployment(deployment_id: str, extra_params: dict[str, float]) -> dict[str, object]: + return { + "model_name": "gpt", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "sk-fake", + "api_base": f"https://{deployment_id}.test/v1", + **extra_params, + }, + "model_info": {"id": deployment_id}, + } + + +def _routing_router( + strategy: RoutingStrategy | str, + a_params: dict[str, float] | None = None, + b_params: dict[str, float] | None = None, +) -> Router: + return Router( + model_list=[_routing_deployment("a", a_params or {}), _routing_deployment("b", b_params or {})], + routing_strategy=strategy, + disable_cooldowns=True, + num_retries=0, + ) + + +async def _selected_deployment_ids(router: Router) -> frozenset[str]: + deployments: Final = [ + await router.async_get_available_deployment(model="gpt", messages=list(ROUTING_MESSAGES), request_kwargs={}) + for _ in range(ROUTING_SELECTIONS) + ] + return frozenset(deployment["model_info"]["id"] for deployment in deployments) + + +def _enum_and_string(strategy: RoutingStrategy) -> list[RoutingStrategy | str]: + return [strategy, strategy.value] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.COST_BASED)) +async def test_router_cost_based_routing_selects_cheapest_deployment(strategy: RoutingStrategy | str) -> None: + router: Final = _routing_router(strategy, EXPENSIVE_COSTS, CHEAP_COSTS) + + assert await _selected_deployment_ids(router) == frozenset({"b"}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "strategy", + [*_enum_and_string(RoutingStrategy.USAGE_BASED_ROUTING), *_enum_and_string(RoutingStrategy.USAGE_BASED_ROUTING_V2)], +) +async def test_router_usage_based_routing_skips_deployment_over_tpm_limit(strategy: RoutingStrategy | str) -> None: + router: Final = _routing_router(strategy, {"tpm": 1}, {"tpm": 1_000_000}) + + assert await _selected_deployment_ids(router) == frozenset({"b"}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.LATENCY_BASED)) +async def test_router_latency_based_routing_avoids_deployment_that_timed_out( + strategy: RoutingStrategy | str, + respx_mock: respx.MockRouter, +) -> None: + router: Final = _routing_router(strategy) + timed_out: Final = respx_mock.post("https://a.test/v1/chat/completions").mock( + side_effect=httpx.ReadTimeout("upstream timed out") + ) + respx_mock.post("https://b.test/v1/chat/completions").respond(200, json=CHAT_RESPONSE) + + for _ in range(5): + if timed_out.called: + break + with contextlib.suppress(litellm.Timeout): + await router.acompletion(model="gpt", messages=list(ROUTING_MESSAGES), max_retries=0) + + assert timed_out.called + assert await _selected_deployment_ids(router) == frozenset({"b"}) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.LEAST_BUSY)) +async def test_router_least_busy_routing_avoids_deployment_with_request_in_flight( + strategy: RoutingStrategy | str, + respx_mock: respx.MockRouter, +) -> None: + router: Final = _routing_router(strategy) + upstream_hosts: Final = asyncio.Queue[str]() + release_upstream: Final = asyncio.Event() + + async def hold_request(request: httpx.Request) -> httpx.Response: + upstream_hosts.put_nowait(request.url.host.split(".")[0]) + await release_upstream.wait() + return httpx.Response(200, json=CHAT_RESPONSE) + + respx_mock.post(url__regex=r"https://[ab]\.test/v1/chat/completions").mock(side_effect=hold_request) + in_flight: Final = asyncio.create_task(router.acompletion(model="gpt", messages=list(ROUTING_MESSAGES))) + try: + busy_id: Final = await asyncio.wait_for(upstream_hosts.get(), timeout=EVENT_TIMEOUT_SECONDS) + selected: Final = await _selected_deployment_ids(router) + finally: + release_upstream.set() + await in_flight + + assert selected == frozenset({"a", "b"} - {busy_id}) + + +@pytest.mark.asyncio +async def test_router_simple_shuffle_ignores_cost_and_spreads_across_deployments() -> None: + router: Final = _routing_router("simple-shuffle", EXPENSIVE_COSTS, CHEAP_COSTS) + + assert await _selected_deployment_ids(router) == frozenset({"a", "b"}) + + +@pytest.mark.parametrize("strategy", _enum_and_string(RoutingStrategy.PROVIDER_BUDGET_LIMITING)) +def test_router_routing_strategy_init_accepts_provider_budget_strategy(strategy: RoutingStrategy | str) -> None: + router: Final = _routing_router(strategy) + + router.routing_strategy_init(routing_strategy=strategy, routing_strategy_args={}) + + assert router.get_settings()["routing_strategy"] == "provider-budget-routing" + + +def test_router_track_deployment_metrics_updates_observable_usage() -> None: + router: Final = Router( + model_list=[ + { + "model_name": "gpt-5-mini", + "litellm_params": {"model": "gpt-5-mini", "api_key": "sk-fake"}, + "model_info": {"id": "metrics-deployment"}, + } + ] + ) + deployment: Final = router.model_list[0] + + router._track_deployment_metrics(deployment=deployment, parent_otel_span=None) + router._track_deployment_metrics(deployment=deployment, parent_otel_span=None) + + assert router.cache.get_cache(key="metrics-deployment", local_only=True) == 2 + + +class _GatedIncrementCache(DualCache): + def __init__(self) -> None: + super().__init__(in_memory_cache=InMemoryCache()) + self.first_increment_started = asyncio.Event() + self.release_first_increment = asyncio.Event() + self.deployment_success_incremented = asyncio.Event() + self.increment_calls = 0 + + def increment_cache(self, key: str, value: int, local_only: bool = False, **kwargs: object) -> int: + result: Final = super().increment_cache(key=key, value=value, local_only=local_only, **kwargs) + if key.endswith(":successes"): + self.deployment_success_incremented.set() + return result + + async def async_increment_cache_pipeline( + self, + increment_list: list[RedisPipelineIncrementOperation], + local_only: bool = False, + parent_otel_span: object = None, + **kwargs: object, + ) -> list[float] | None: + self.increment_calls += 1 + if self.increment_calls == 1: + self.first_increment_started.set() + await self.release_first_increment.wait() + return await super().async_increment_cache_pipeline( + increment_list=increment_list, + local_only=local_only, + parent_otel_span=parent_otel_span, + **kwargs, + ) + + +class _UnavailableIncrementCache(DualCache): + def __init__(self) -> None: + super().__init__(in_memory_cache=InMemoryCache()) + self.first_increment_started = asyncio.Event() + self.release_first_increment = asyncio.Event() + self.deployment_success_incremented = asyncio.Event() + self.increment_calls = 0 + + def increment_cache(self, key: str, value: int, local_only: bool = False, **kwargs: object) -> int: + result: Final = super().increment_cache(key=key, value=value, local_only=local_only, **kwargs) + if key.endswith(":successes"): + self.deployment_success_incremented.set() + return result + + async def async_increment_cache_pipeline( + self, + increment_list: list[RedisPipelineIncrementOperation], + local_only: bool = False, + parent_otel_span: object = None, + **kwargs: object, + ) -> list[float] | None: + self.increment_calls += 1 + if self.increment_calls == 1: + self.first_increment_started.set() + await self.release_first_increment.wait() + raise RuntimeError("cache unavailable") + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("router_minute_pinned") +async def test_router_success_callback_during_pre_header_increment_does_not_double_count() -> None: + router: Final = _rpm_tpm_router("lit-3058-race") + cache: Final = _GatedIncrementCache() + router.cache = cache + request: Final = asyncio.create_task( + router.acompletion( + model="gpt-5-mini", + messages=[{"role": "user", "content": "hi"}], + mock_response="pong", + ) + ) + + await asyncio.wait_for(cache.first_increment_started.wait(), timeout=EVENT_TIMEOUT_SECONDS) + await asyncio.wait_for(cache.deployment_success_incremented.wait(), timeout=EVENT_TIMEOUT_SECONDS) + assert get_deployment_successes_for_current_minute(router, "lit-3058-race") == 1 + assert cache.increment_calls == 1 + cache.release_first_increment.set() + response: Final = await request + + assert await router.get_model_group_usage("gpt-5-mini") == (response.usage.total_tokens, 1) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("router_minute_pinned") +async def test_router_failed_pre_header_increment_clears_counted_tokens_stamp() -> None: + router: Final = _rpm_tpm_router("lit-3058-fail") + cache: Final = _UnavailableIncrementCache() + router.cache = cache + metadata: Final[dict[str, object]] = {} + request: Final = asyncio.create_task( + router.acompletion( + model="gpt-5-mini", + messages=[{"role": "user", "content": "hi"}], + mock_response="pong", + metadata=metadata, + ) + ) + + await asyncio.wait_for(cache.first_increment_started.wait(), timeout=EVENT_TIMEOUT_SECONDS) + await asyncio.wait_for(cache.deployment_success_incremented.wait(), timeout=EVENT_TIMEOUT_SECONDS) + assert metadata[ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY] == 30 + assert get_deployment_successes_for_current_minute(router, "lit-3058-fail") == 1 + assert cache.increment_calls == 1 + cache.release_first_increment.set() + response: Final = await request + + assert response.usage.total_tokens == 30 + assert ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY not in metadata + assert _ratelimit_headers(response)["x-ratelimit-remaining-requests"] == 100 + assert await router.get_model_group_usage("gpt-5-mini") == (None, None) + + +ASSISTANT_RESPONSE: Final = { + "object": "assistant", + "created_at": 1700000000, + "name": "offline", + "description": None, + "model": "gpt-4o-mini", + "instructions": "hello", + "tools": [], + "metadata": {}, +} + + +@pytest.mark.asyncio +async def test_router_assistants_endpoint_factory_invokes_provider( + respx_mock: respx.MockRouter, +) -> None: + router: Final = Router(model_list=[]) + route: Final = respx_mock.post("https://api.openai.com/v1/assistants").respond( + 200, json={**ASSISTANT_RESPONSE, "id": "asst-1"} + ) + + response: Final = await router._pass_through_assistants_endpoint_factory( + original_function=litellm.acreate_assistants, + custom_llm_provider="openai", + model="gpt-4o-mini", + api_key="sk-fake", + name="offline", + ) + + assert route.call_count == 1 + assert json.loads(route.calls[0].request.content) == {"model": "gpt-4o-mini", "name": "offline"} + assert response.id == "asst-1" + + +@pytest.mark.asyncio +async def test_router_factory_function_returns_invokable_assistants_wrapper( + respx_mock: respx.MockRouter, +) -> None: + router: Final = Router(model_list=[]) + route: Final = respx_mock.post("https://api.openai.com/v1/assistants").respond( + 200, json={**ASSISTANT_RESPONSE, "id": "asst-2"} + ) + wrapper: Final = router.factory_function(litellm.acreate_assistants, call_type="assistants") + + response: Final = await wrapper( + custom_llm_provider="openai", + model="gpt-4o-mini", + api_key="sk-fake", + name="offline", + ) + + assert route.call_count == 1 + assert json.loads(route.calls[0].request.content) == {"model": "gpt-4o-mini", "name": "offline"} + assert response.id == "asst-2" + + +@pytest.mark.asyncio +async def test_router_moderation_endpoint_factory_invokes_default_model( + respx_mock: respx.MockRouter, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("OPENAI_API_KEY", "sk-fake") + router: Final = Router(model_list=[]) + route: Final = respx_mock.post("https://api.openai.com/v1/moderations").respond( + 200, + json={ + "id": "modr-2", + "model": "omni-moderation-latest", + "results": [{"flagged": False, "categories": {}, "category_scores": {}}], + }, + ) + + response: Final = await router._pass_through_moderation_endpoint_factory( + original_function=litellm.amoderation, + custom_llm_provider="openai", + input="hello", + model=None, + api_key="sk-fake", + ) + + assert route.call_count == 1 + assert json.loads(route.calls[0].request.content) == {"input": "hello"} + assert response.id == "modr-2" diff --git a/tests/unit/test_router/test_router_vector_store_endpoints.py b/tests/unit/test_router/test_router_vector_store_endpoints.py new file mode 100644 index 00000000000..ad431c656f9 --- /dev/null +++ b/tests/unit/test_router/test_router_vector_store_endpoints.py @@ -0,0 +1,101 @@ +import json +from collections.abc import Iterator +from typing import Final + +import httpx +import litellm +import pytest +import respx +from litellm import Router + +VECTOR_STORES_URL: Final = "https://api.openai.com/v1/vector_stores" + + +@pytest.fixture(autouse=True) +def httpx_transport(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setenv("OPENAI_API_KEY", "sk-vector-store-test") + monkeypatch.delenv("OPENAI_BASE_URL", raising=False) + monkeypatch.delenv("OPENAI_API_BASE", raising=False) + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.asyncio +async def test_avector_store_update_sends_name_and_metadata_to_provider() -> None: + with respx.mock(assert_all_called=True) as mock: + route: Final = mock.post(f"{VECTOR_STORES_URL}/vs_test123").mock( + return_value=httpx.Response( + 200, + json={ + "id": "vs_test123", + "object": "vector_store", + "created_at": 1699061776, + "name": "Updated Name", + "metadata": {"key": "value"}, + "status": "completed", + }, + ) + ) + result: Final = await Router(model_list=[]).avector_store_update( + vector_store_id="vs_test123", + name="Updated Name", + metadata={"key": "value"}, + custom_llm_provider="openai", + ) + sent: Final = json.loads(route.calls.last.request.content) + assert route.call_count == 1 + assert sent["name"] == "Updated Name" + assert sent["metadata"] == {"key": "value"} + assert result["id"] == "vs_test123" + assert result["name"] == "Updated Name" + assert result["metadata"]["key"] == "value" + + +def test_vector_store_list_forwards_pagination_params_to_provider() -> None: + with respx.mock(assert_all_called=True) as mock: + route: Final = mock.get(VECTOR_STORES_URL).mock( + return_value=httpx.Response( + 200, + json={ + "object": "list", + "data": [{"id": f"vs_{index}", "object": "vector_store"} for index in range(5)], + "has_more": True, + "first_id": "vs_0", + "last_id": "vs_4", + }, + ) + ) + result: Final = Router(model_list=[]).vector_store_list( + limit=5, after="vs_previous", order="asc", custom_llm_provider="openai" + ) + params: Final = route.calls.last.request.url.params + assert params["limit"] == "5" + assert params["after"] == "vs_previous" + assert params["order"] == "asc" + assert result["has_more"] is True + assert len(result["data"]) == 5 + + +def test_vector_store_update_forwards_expires_after_to_provider() -> None: + expires_after: Final = {"anchor": "last_active_at", "days": 7} + with respx.mock(assert_all_called=True) as mock: + route: Final = mock.post(f"{VECTOR_STORES_URL}/vs_test123").mock( + return_value=httpx.Response( + 200, + json={ + "id": "vs_test123", + "object": "vector_store", + "expires_after": expires_after, + "expires_at": 1699668576, + }, + ) + ) + result: Final = Router(model_list=[]).vector_store_update( + vector_store_id="vs_test123", expires_after=expires_after, custom_llm_provider="openai" + ) + sent: Final = json.loads(route.calls.last.request.content) + assert sent["expires_after"] == expires_after + assert result["expires_after"]["days"] == 7 + assert result["expires_at"] == 1699668576 diff --git a/tests/unit/test_router_order_fallback.py b/tests/unit/test_router_order_fallback.py index a86338c625a..37859761858 100644 --- a/tests/unit/test_router_order_fallback.py +++ b/tests/unit/test_router_order_fallback.py @@ -11,6 +11,7 @@ from typing import Final, Optional import httpx import pytest +import respx from openai import AsyncOpenAI import litellm @@ -645,6 +646,171 @@ async def test_text_completion_order_fallback_hop_does_not_send_target_order_ups assert all("_target_order" not in body for body in upstream_bodies) + +_OPENAI_RESPONSES_URL: Final = "https://api.openai.com/v1/responses" +_MANTLE_RESPONSES_URL: Final = "https://bedrock-mantle.us-east-1.api.aws/openai/v1/responses" +_OVERLOADED_UPSTREAM: Final = {"error": {"message": "overloaded", "type": "server_error", "code": "server_error"}} + + +def _completed_response_body(response_id: str, model: str, text: str) -> dict[str, object]: + return { + "id": response_id, + "object": "response", + "created_at": 0, + "status": "completed", + "model": model, + "output": [ + { + "id": f"msg_{response_id}", + "type": "message", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": text, "annotations": []}], + } + ], + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + } + + +def _responses_history_with_order_1_reasoning() -> list[dict]: + return [ + {"type": "message", "role": "user", "content": "What is 17*23?"}, + { + "type": "reasoning", + "id": "rs_order1", + "encrypted_content": "gAAAAA-minted-by-order-1", + "summary": [{"type": "summary_text", "text": "multiply 17 by 23"}], + }, + {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "391"}]}, + {"type": "message", "role": "user", "content": "And 19*21?"}, + ] + + +def _responses_history_without_order_1_encrypted_reasoning() -> list[dict]: + return [ + {"type": "message", "role": "user", "content": "What is 17*23?"}, + {"type": "reasoning", "summary": [{"type": "summary_text", "text": "multiply 17 by 23"}]}, + {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "391"}]}, + {"type": "message", "role": "user", "content": "And 19*21?"}, + ] + + +def _openai_then_mantle_order_router() -> Router: + return Router( + model_list=[ + { + "model_name": "gpt-6-astra", + "litellm_params": {"model": "openai/gpt-6-astra", "api_key": "openai-key", "order": 1}, + "model_info": {"id": "openai-order-1"}, + }, + { + "model_name": "gpt-6-astra", + "litellm_params": { + "model": "bedrock_mantle/openai.gpt-6-astra", + "api_key": "mantle-bearer-token", + "aws_region_name": "us-east-1", + "order": 2, + }, + "model_info": {"id": "mantle-order-2"}, + }, + ], + num_retries=0, + ) + + +def _two_openai_orders_on_one_encryption_boundary_router() -> Router: + return Router( + model_list=[ + { + "model_name": "gpt-6-astra", + "litellm_params": { + "model": "openai/gpt-6-astra", + "api_base": "https://api.openai.com/v1", + "api_key": "openai-key", + "order": 1, + }, + "model_info": {"id": "openai-order-1"}, + }, + { + "model_name": "gpt-6-astra", + "litellm_params": { + "model": "openai/gpt-6-astra-mini", + "api_base": "https://api.openai.com/v1", + "api_key": "openai-key", + "order": 2, + }, + "model_info": {"id": "openai-order-2"}, + }, + ], + num_retries=0, + ) + + +@pytest.mark.asyncio +async def test_responses_order_fallback_hop_drops_the_encrypted_reasoning_the_next_provider_cannot_decrypt( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + openai_route: Final = respx_mock.post(_OPENAI_RESPONSES_URL).mock( + return_value=httpx.Response(500, json=_OVERLOADED_UPSTREAM) + ) + mantle_route: Final = respx_mock.post(_MANTLE_RESPONSES_URL).mock( + return_value=httpx.Response(200, json=_completed_response_body("resp_mantle", "openai.gpt-6-astra", "399")) + ) + + response = await _openai_then_mantle_order_router().aresponses( + model="gpt-6-astra", input=_responses_history_with_order_1_reasoning(), store=False + ) + + assert response._hidden_params["model_id"] == "mantle-order-2" + assert json.loads(openai_route.calls.last.request.read())["input"] == _responses_history_with_order_1_reasoning() + assert ( + json.loads(mantle_route.calls.last.request.read())["input"] + == _responses_history_without_order_1_encrypted_reasoning() + ) + + +@pytest.mark.asyncio +async def test_responses_order_fallback_hop_keeps_the_encrypted_reasoning_the_same_boundary_can_decrypt( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + openai_route: Final = respx_mock.post(_OPENAI_RESPONSES_URL).mock( + side_effect=[ + httpx.Response(500, json=_OVERLOADED_UPSTREAM), + httpx.Response(200, json=_completed_response_body("resp_order2", "gpt-6-astra-mini", "399")), + ] + ) + + response = await _two_openai_orders_on_one_encryption_boundary_router().aresponses( + model="gpt-6-astra", input=_responses_history_with_order_1_reasoning(), store=False + ) + + assert response._hidden_params["model_id"] == "openai-order-2" + assert [json.loads(call.request.read())["input"] for call in openai_route.calls] == [ + _responses_history_with_order_1_reasoning(), + _responses_history_with_order_1_reasoning(), + ] + + +def test_fallback_hop_reads_the_deployment_that_just_failed_from_the_metadata_bucket_it_writes(): + router: Final = _two_openai_orders_on_one_encryption_boundary_router() + order_2: Final = router.get_deployment(model_id="openai-order-2").model_dump(exclude_none=True) + hop_input: Final = _responses_history_with_order_1_reasoning() + hop_kwargs: Final = { + "model": "gpt-6-astra", + "input": hop_input, + "fallback_depth": 1, + "metadata": {"model_info": {"id": "openai-order-1"}}, + "litellm_metadata": {"previous_models": [{"deployment_id": None}]}, + } + + router._update_kwargs_with_deployment(deployment=order_2, kwargs=hop_kwargs) + + assert hop_input == _responses_history_with_order_1_reasoning() + assert hop_kwargs["metadata"]["model_info"]["id"] == "openai-order-2" + + def test_check_non_standard_fallback_format(): from litellm.router_utils.fallback_event_handlers import ( check_non_standard_fallback_format, diff --git a/tests/unit/test_ruff_strict_gate.py b/tests/unit/test_ruff_strict_gate.py index 8fa9a18cf53..852ebef7436 100644 --- a/tests/unit/test_ruff_strict_gate.py +++ b/tests/unit/test_ruff_strict_gate.py @@ -1,87 +1,27 @@ -import importlib.util import json import re import shutil import subprocess import sys from pathlib import Path +from typing import Final import pytest +import ruff_strict_gate as gate + if sys.version_info >= (3, 11): import tomllib else: import tomli as tomllib _REPO_ROOT = Path(__file__).resolve().parents[2] -_MODULE_PATH = _REPO_ROOT / "scripts" / "ruff_strict_gate.py" -_spec = importlib.util.spec_from_file_location("ruff_strict_gate", _MODULE_PATH) -gate = importlib.util.module_from_spec(_spec) -_spec.loader.exec_module(gate) Violation = gate.Violation _ENABLED_BY_RUFF_DEFAULTS = frozenset({"F401"}) -def rule(name, limit): - return {name: {"limit": limit}} - - -def test_under_ceiling_passes(): - assert gate.evaluate({"ANN001": 100}, {"ANN001": 100}, rule("ANN001", 110)) == [] - - -def test_ceiling_is_the_limit_boundary(): - budget = rule("ANN001", 110) - at = gate.evaluate({"ANN001": 110}, {"ANN001": 90}, budget) - over = gate.evaluate({"ANN001": 111}, {"ANN001": 90}, budget) - assert at == [] - assert [b.rule for b in over] == ["ANN001"] - assert over[0].cap == 110 - assert over[0].added == 21 - - -def test_over_ceiling_and_change_added_fails(): - breaches = gate.evaluate({"C901": 11}, {"C901": 9}, rule("C901", 10)) - assert [b.rule for b in breaches] == ["C901"] - assert breaches[0].added == 2 - - -def test_base_already_over_ceiling_change_added_nothing_is_not_blamed(): - # drift safety: base is over limit, this change leaves the count where it is - assert gate.evaluate({"C901": 15}, {"C901": 15}, rule("C901", 10)) == [] - - -def test_change_that_reduces_an_over_ceiling_rule_is_not_blamed(): - # still over limit, but moving the right direction - assert gate.evaluate({"C901": 14}, {"C901": 16}, rule("C901", 10)) == [] - - -def test_rules_are_independent(): - budget = {**rule("ANN001", 150), **rule("C901", 10)} - breaches = gate.evaluate( - {"ANN001": 130, "C901": 11}, {"ANN001": 100, "C901": 10}, budget - ) - assert [b.rule for b in breaches] == ["C901"] # ANN001 130 <= 150, C901 11 > 10 - - -def test_missing_rule_counts_as_zero(): - assert gate.evaluate({}, {}, rule("C901", 0)) == [] - - -def test_update_ratchets_limit_down_by_what_the_branch_fixed_never_up(): - budget = {**rule("ANN001", 150), **rule("C901", 10)} - # ANN001 fixed 20 (100 -> 80) so its limit falls 150 -> 130; C901 grew, so its - # limit holds flat at 10 (a fix must never loosen a ceiling). - current = {"ANN001": 80, "C901": 12} - base = {"ANN001": 100, "C901": 9} - assert gate.ratcheted_budget(budget, current, base) == { - "ANN001": {"limit": 130}, - "C901": {"limit": 10}, - } - - def test_parse_changed_lines_maps_added_lines_per_file(): diff = ( "+++ b/litellm/a.py\n" @@ -109,62 +49,6 @@ def test_parse_changed_lines_handles_single_and_ranged_hunks(hunk): assert gate.parse_changed_lines(f"+++ b/litellm/a.py\n{hunk}\n")["litellm/a.py"] -def test_over_ceiling_flags_only_counts_above_the_limit(): - budget = rule("C901", 10) - assert gate.over_ceiling({"C901": 10}, budget) == frozenset() - assert gate.over_ceiling({"C901": 11}, budget) == frozenset({"C901"}) - assert gate.over_ceiling({}, budget) == frozenset() - - -def test_over_ceiling_ignores_rules_missing_from_the_budget(): - assert gate.over_ceiling({"NEW99": 100}, rule("C901", 10)) == frozenset() - - -def test_over_ceiling_is_independent_across_rules(): - budget = {**rule("ANN001", 150), **rule("C901", 10)} - assert gate.over_ceiling({"ANN001": 130, "C901": 11}, budget) == frozenset({"C901"}) - - -def _git(cwd, *args): - proc = subprocess.run(["git", *args], cwd=cwd, capture_output=True, text=True) - assert proc.returncode == 0, proc.stderr - return proc.stdout.strip() - - -def _commit(cwd, name): - (cwd / name).write_text(name) - _git(cwd, "add", "-A") - _git(cwd, "commit", "-q", "-m", name) - return _git(cwd, "rev-parse", "HEAD") - - -def _branched_repo(tmp_path): - repo = tmp_path / "repo" - repo.mkdir() - _git(repo, "init", "-q", "-b", "main") - _git(repo, "config", "user.email", "gate@example.com") - _git(repo, "config", "user.name", "gate") - _git(repo, "config", "commit.gpgsign", "false") - branch_point = _commit(repo, "shared.txt") - _git(repo, "checkout", "-q", "-b", "feature") - _commit(repo, "feature.txt") - _git(repo, "checkout", "-q", "main") - base_tip = _commit(repo, "drift.txt") - _git(repo, "checkout", "-q", "feature") - return repo, branch_point, base_tip - - -def test_base_point_is_the_branch_point_when_no_merge_is_in_progress(tmp_path): - repo, branch_point, _ = _branched_repo(tmp_path) - assert gate.resolve_base_point("main", cwd=repo) == branch_point - - -def test_base_point_mid_merge_advances_to_the_merged_in_base_tip(tmp_path): - repo, _, base_tip = _branched_repo(tmp_path) - _git(repo, "merge", "--no-commit", "--no-ff", "main") - assert gate.resolve_base_point("main", cwd=repo) == base_tip - - def _lint_section(config_name: str) -> dict: return tomllib.loads((_REPO_ROOT / config_name).read_text())["lint"] @@ -189,10 +73,6 @@ def _selected_by_the_normal_config() -> frozenset: return frozenset(_lint_section("ruff.toml")["extend-select"]) | _ENABLED_BY_RUFF_DEFAULTS -def _budgeted_rules() -> frozenset: - return frozenset(json.loads((_REPO_ROOT / "ruff-strict-budget.json").read_text())) - - def _ruff_binary() -> str | None: beside_interpreter = Path(sys.executable).with_name("ruff") return str(beside_interpreter) if beside_interpreter.exists() else shutil.which("ruff") @@ -302,37 +182,6 @@ def test_every_base_owned_rule_is_external_or_selected_in_the_strict_config(all_ ) -def test_every_budgeted_rule_is_one_the_gate_actually_measures(): - selectors = tuple(_lint_section("ruff-strict.toml")["select"]) - unmeasured = frozenset(code for code in _budgeted_rules() if not code.startswith(selectors)) - assert unmeasured == frozenset(), ( - f"the gate never counts {sorted(unmeasured)}, so their ceilings are dead config that reads " - "as coverage. Either select them in ruff-strict.toml or drop them from the budget." - ) - - -@_needs_ruff -def test_every_strict_selected_rule_is_budgeted_or_hard_failed_by_the_base_config(all_ruff_rule_codes): - strict_enabled = frozenset( - code - for code in all_ruff_rule_codes - if code.startswith(tuple(_lint_section("ruff-strict.toml")["select"])) - ) - base_hard_failed = tuple(_lint_section("ruff.toml")["extend-select"]) - unpoliced = frozenset( - code - for code in strict_enabled - if code not in _budgeted_rules() - and not code.startswith(base_hard_failed) - and code not in _ENABLED_BY_RUFF_DEFAULTS - ) - assert unpoliced == frozenset(), ( - f"nothing enforces {sorted(unpoliced)}: the gate skips rules missing from the budget, and " - "the base config does not hard-fail them. Re-add a budget ceiling or graduate them into " - "ruff.toml's lint.extend-select." - ) - - @_needs_ruff def test_a_noqa_for_a_strict_gate_rule_survives_the_normal_ruff_run(): assert "RUF100" not in _ruff_output_for_noqa("ANN202") @@ -373,3 +222,21 @@ def test_a_graduated_rule_can_still_be_suppressed_without_tripping_unused_noqa() output = _ruff_output_for_source(suppressed) assert "UP006" not in output assert "RUF100" not in output + + +def test_editing_either_ruff_config_or_upgrading_ruff_rekeys_the_base_counts(tmp_path: Path) -> None: + strict: Final = tmp_path / "ruff-strict.toml" + base: Final = tmp_path / "ruff.toml" + strict.write_text("strict v1\n") + base.write_text("base v1\n") + + def artifact(version: str) -> str: + return gate.checker_identity(strict, base, lambda: version).artifact_name("abc123") + + before: Final = artifact("ruff 0.1.0") + assert artifact("ruff 0.1.0") == before + assert artifact("ruff 0.2.0") != before + strict.write_text("strict v2\n") + after_strict_edit: Final = artifact("ruff 0.1.0") + base.write_text("base v2\n") + assert len({before, after_strict_edit, artifact("ruff 0.1.0")}) == 3 diff --git a/tests/unit/test_service_logger.py b/tests/unit/test_service_logger.py index 3d74642a03e..99ba32ed955 100644 --- a/tests/unit/test_service_logger.py +++ b/tests/unit/test_service_logger.py @@ -7,10 +7,17 @@ is called without call_type in kwargs (e.g. from batch polling callbacks). import pytest from datetime import datetime +from typing import Final from unittest.mock import AsyncMock, patch +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + import litellm from litellm._service_logger import ServiceLogging +from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger +from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig from litellm.types.services import ServiceTypes @@ -316,3 +323,47 @@ async def test_only_redis_service_spans_carry_the_ambient_key_family(monkeypatch "redis.get router_session_pins": "router_session_pins", "batch_write_to_db _PROXY_track_cost_callback": None, } + + +def _in_memory_provider(exporter: InMemorySpanExporter) -> TracerProvider: + provider: Final = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + return provider + + +@pytest.mark.asyncio +async def test_generic_otel_logger_receives_service_event_once_beside_langfuse_otel( + monkeypatch: pytest.MonkeyPatch, +) -> None: + langfuse_exporter: Final = InMemorySpanExporter() + generic_exporter: Final = InMemorySpanExporter() + langfuse_provider: Final = _in_memory_provider(langfuse_exporter) + generic_provider: Final = _in_memory_provider(generic_exporter) + langfuse_logger: Final = LangfuseOtelLogger( + config=OpenTelemetryConfig(exporter="console", skip_set_global=True), tracer_provider=langfuse_provider + ) + generic_logger: Final = OpenTelemetry( + config=OpenTelemetryConfig(exporter="console", skip_set_global=True), tracer_provider=generic_provider + ) + monkeypatch.setattr(litellm, "service_callback", [langfuse_logger, generic_logger]) + parent: Final = generic_logger.tracer.start_span("parent") + + try: + await ServiceLogging().async_service_success_hook( + service=ServiceTypes.DB, + call_type="success", + duration=0.1, + parent_otel_span=parent, + start_time=0.0, + end_time=1.0, + ) + finally: + langfuse_provider.shutdown() + generic_provider.shutdown() + + service_spans: Final = [ + span for span in generic_exporter.get_finished_spans() if span.attributes.get("service") == ServiceTypes.DB.value + ] + assert len(service_spans) == 1 + assert service_spans[0].attributes.get("call_type") == "success" + assert langfuse_exporter.get_finished_spans() == () diff --git a/tests/unit/test_test_quality_gate.py b/tests/unit/test_test_quality_gate.py index 873e7b9b336..b98e539d4b0 100644 --- a/tests/unit/test_test_quality_gate.py +++ b/tests/unit/test_test_quality_gate.py @@ -1,12 +1,11 @@ """Tests for scripts/test_quality_gate.py. -The gate's whole value is that it blames a change only for what it adds and that a -limit can never rise. Both live in pure functions, so they are tested directly: -`evaluate` for the blame rule, `ratcheted_budget` for the one-way ratchet, and -`parse_changed_lines` for the diff scan that turns a breach into file:line. +The blame rule itself lives in scripts/lint_base_counts.py and is tested there; +what is tested here is `parse_changed_lines`, the diff scan that turns a breach into +file:line. The base scan spawns a worktree and a checker subprocess, so its cleanup +on termination is driven end to end. """ -import importlib.util import os import signal import subprocess @@ -15,77 +14,28 @@ import time from collections.abc import Callable from contextlib import suppress from pathlib import Path -from typing import NamedTuple +from typing import Final, NamedTuple + +import test_quality_gate as gate _REPO_ROOT = Path(__file__).resolve().parents[2] _MODULE_PATH = _REPO_ROOT / "scripts" / "test_quality_gate.py" -_spec = importlib.util.spec_from_file_location("test_quality_gate", _MODULE_PATH) -gate = importlib.util.module_from_spec(_spec) -# @dataclass(slots=True) rebuilds its class through sys.modules[__module__], so the -# module has to be registered before exec_module runs or Scope fails to construct. -sys.modules[_spec.name] = gate -_spec.loader.exec_module(gate) - -_BUDGET = {"TQ001": {"limit": 10}, "TQ003": {"limit": 5}} _SCAN_BASE = ( - "import importlib.util, pathlib, sys\n" - "spec = importlib.util.spec_from_file_location('test_quality_gate', sys.argv[1])\n" - "gate = importlib.util.module_from_spec(spec)\n" - "sys.modules[spec.name] = gate\n" - "spec.loader.exec_module(gate)\n" - "gate.base_counts('HEAD', repo_root=pathlib.Path(sys.argv[2]), checker=pathlib.Path(sys.argv[3]))\n" + "import pathlib, sys\n" + "import test_quality_gate as gate\n" + "gate.base_counts('HEAD', repo_root=pathlib.Path(sys.argv[1]), checker=pathlib.Path(sys.argv[2]))\n" ) _SCAN_BASE_WITH_SIGHUP_IGNORED = "import signal\nsignal.signal(signal.SIGHUP, signal.SIG_IGN)\n" + _SCAN_BASE -def test_a_rule_within_its_limit_is_not_a_breach(): - assert gate.evaluate({"TQ001": 10}, {"TQ001": 10}, _BUDGET) == () - - -def test_a_rule_over_its_limit_that_the_change_added_is_a_breach(): - breaches = gate.evaluate({"TQ001": 12}, {"TQ001": 10}, _BUDGET) - assert [(b.rule, b.total, b.cap, b.added) for b in breaches] == [("TQ001", 12, 10, 2)] - - -def test_drift_already_in_the_base_is_not_blamed_on_the_change(): - assert gate.evaluate({"TQ001": 14}, {"TQ001": 14}, _BUDGET) == () - - -def test_a_change_that_reduces_an_over_limit_rule_is_not_blamed(): - assert gate.evaluate({"TQ001": 13}, {"TQ001": 14}, _BUDGET) == () - - -def test_a_rule_absent_from_head_counts_as_zero(): - assert gate.evaluate({}, {}, _BUDGET) == () - - -def test_over_ceiling_names_only_the_rules_above_their_limit(): - assert gate.over_ceiling({"TQ001": 11, "TQ003": 5}, _BUDGET) == frozenset({"TQ001"}) - - -def test_over_ceiling_is_empty_when_everything_fits(): - assert gate.over_ceiling({"TQ001": 10, "TQ003": 4}, _BUDGET) == frozenset() - - -def test_ratchet_lowers_a_limit_by_what_the_branch_fixed(): - updated = gate.ratcheted_budget(_BUDGET, {"TQ001": 6}, {"TQ001": 10}) - assert updated["TQ001"]["limit"] == 6 - - -def test_ratchet_never_raises_a_limit_when_violations_grew(): - updated = gate.ratcheted_budget(_BUDGET, {"TQ001": 20}, {"TQ001": 10}) - assert updated["TQ001"]["limit"] == 10 - - -def test_ratchet_never_goes_below_zero(): - updated = gate.ratcheted_budget({"TQ001": {"limit": 2}}, {"TQ001": 0}, {"TQ001": 100}) - assert updated["TQ001"]["limit"] == 0 - - -def test_ratchet_lowers_a_rule_introduced_on_this_branch_like_any_other(): - updated = gate.ratcheted_budget(_BUDGET, {"TQ001": 4}, {"TQ001": 10}) - assert updated["TQ001"]["limit"] == 4 +def test_editing_the_checker_rekeys_the_base_counts(tmp_path: Path) -> None: + checker: Final = tmp_path / "check.py" + checker.write_text("print('v1')\n") + before: Final = gate.checker_identity(checker).artifact_name("abc123") + assert gate.checker_identity(checker).artifact_name("abc123") == before + checker.write_text("print('v2')\n") + assert gate.checker_identity(checker).artifact_name("abc123") != before def test_parse_changed_lines_groups_hunks_under_their_own_file(): @@ -132,16 +82,6 @@ def test_introduced_keeps_only_violations_on_changed_lines(): assert kept == (gate.Violation("tests/a.py", 3, "TQ001"),) -def test_the_shipped_budget_covers_every_rule_the_checker_can_emit(): - import json - - budget = json.loads((_REPO_ROOT / "test-quality-budget.json").read_text()) - assert set(budget) == { - "TQ001", "TQ002", "TQ003", "TQ004", "TQ005", "TQ006", "TQ007", "TQ009" - } - assert all(spec["limit"] >= 0 for spec in budget.values()) - - def _git(cwd: Path, *args: str) -> str: proc = subprocess.run(["git", *args], cwd=cwd, capture_output=True, text=True) assert proc.returncode == 0, proc.stderr @@ -204,7 +144,8 @@ def _base_scan_stalled_in_its_checker(tmp_path: Path, driver: str) -> _StalledSc temp_dir = tmp_path / "tmp" temp_dir.mkdir() scan = subprocess.Popen( - [sys.executable, "-c", driver, str(_MODULE_PATH), str(repo), str(slow_checker)], + [sys.executable, "-c", driver, str(repo), str(slow_checker)], + cwd=_MODULE_PATH.parent, env={**os.environ, "TMPDIR": str(temp_dir)}, ) if not _wait_until(scanning.exists, 30): diff --git a/tests/unit/test_type_check_gate.py b/tests/unit/test_type_check_gate.py index d104e4ca0c8..76003a3d296 100644 --- a/tests/unit/test_type_check_gate.py +++ b/tests/unit/test_type_check_gate.py @@ -1,14 +1,11 @@ import hashlib -import importlib.util import json -import os -import subprocess from pathlib import Path +from typing import Final -_MODULE_PATH = Path(__file__).resolve().parents[2] / "scripts" / "type_check_gate.py" -_spec = importlib.util.spec_from_file_location("type_check_gate", _MODULE_PATH) -gate = importlib.util.module_from_spec(_spec) -_spec.loader.exec_module(gate) +import pytest + +import type_check_gate as gate ROOT = gate.REPO_ROOT @@ -132,99 +129,6 @@ def test_run_basedpyright_fails_loudly_on_a_crash_exit_code(tmp_path): gate.run_basedpyright(cwd=tmp_path, env_dir=env_dir) -def test_at_or_under_ceiling_passes(): - budget = {"no-any-return": {"limit": 5}} - assert gate.evaluate({"no-any-return": 5}, {}, budget) == [] - - -def test_one_more_error_than_ceiling_fails(): - budget = {"no-any-return": {"limit": 5}} - assert gate.evaluate({"no-any-return": 6}, {}, budget) == [ - gate.Breach("no-any-return", 6, 5, 6) - ] - - -def test_limit_absorbs_increase_up_to_it_then_fails_past_it(): - budget = {"arg-type": {"limit": 10}} - assert gate.evaluate({"arg-type": 10}, {}, budget) == [] - assert gate.evaluate({"arg-type": 11}, {}, budget) == [ - gate.Breach("arg-type", 11, 10, 11) - ] - - -def test_unbudgeted_new_code_uses_default_limit(): - assert gate.evaluate({"brand-new": gate.DEFAULT_LIMIT}, {}, {}) == [] - assert gate.evaluate({"brand-new": gate.DEFAULT_LIMIT + 1}, {}, {}) == [ - gate.Breach( - "brand-new", - gate.DEFAULT_LIMIT + 1, - gate.DEFAULT_LIMIT, - gate.DEFAULT_LIMIT + 1, - ) - ] - - -def test_drift_already_over_cap_in_base_is_not_blamed_on_a_flat_change(): - # The bystander case: a rule sits over its limit because two earlier PRs - # summed past it. A PR that branches off that base and adds nothing must pass - # -- total > limit but total == base, so the `> base` guard spares it. - budget = {"arg-type": {"limit": 10}} - assert gate.evaluate({"arg-type": 12}, {"arg-type": 12}, budget) == [] - - -def test_change_that_grows_an_over_cap_rule_is_blamed_for_only_what_it_added(): - # Over limit AND above base: blamed, and `added` is the delta vs base, not the - # whole overage, so the message points at this change's contribution. - budget = {"arg-type": {"limit": 10}} - assert gate.evaluate({"arg-type": 14}, {"arg-type": 12}, budget) == [ - gate.Breach("arg-type", 14, 10, 2) - ] - - -def test_reducing_an_over_cap_rule_below_base_passes(): - budget = {"arg-type": {"limit": 10}} - assert gate.evaluate({"arg-type": 11}, {"arg-type": 12}, budget) == [] - - -def test_no_output_against_a_nonempty_budget_is_a_vacuous_run(): - # A crashed type checker emits nothing; the gate must not certify it as clean. - budget = {"no-untyped-def": {"limit": 4898}} - assert gate.is_vacuous_run({}, budget) is True - - -def test_genuine_zero_and_empty_budget_are_not_vacuous(): - assert gate.is_vacuous_run({}, {}) is False - assert gate.is_vacuous_run({}, {"no-untyped-def": {"limit": 0}}) is False - assert ( - gate.is_vacuous_run({"arg-type": 1}, {"arg-type": {"limit": 10}}) is False - ) - - -def test_update_ratchets_a_limit_down_by_what_the_branch_fixed(): - # A rule that dropped from 40 (branch point) to 30 (current) fixed 10, so its - # limit of 100 falls to 90 -- the granted headroom (60) is preserved, not the - # raw count. - budget = {"reportAny": {"limit": 100}} - assert gate.ratcheted_budget(budget, {"reportAny": 30}, {"reportAny": 40}) == { - "reportAny": {"limit": 90} - } - - -def test_update_never_raises_a_limit_when_a_rule_grows(): - # Adding violations must not loosen the ceiling; the limit holds flat. - budget = {"reportAny": {"limit": 100}} - assert gate.ratcheted_budget(budget, {"reportAny": 55}, {"reportAny": 40}) == { - "reportAny": {"limit": 100} - } - - -def test_update_clamps_a_limit_at_zero_never_negative(): - budget = {"reportAny": {"limit": 5}} - assert gate.ratcheted_budget(budget, {"reportAny": 0}, {"reportAny": 40}) == { - "reportAny": {"limit": 0} - } - - def test_malformed_basedpyright_json_exits_loudly_not_as_zero_errors(): import pytest @@ -238,35 +142,6 @@ def test_empty_basedpyright_payload_counts_zero(): assert gate.count_basedpyright("") == {} -def test_over_ceiling_flags_only_rules_above_their_limit(): - budget = {"reportAny": {"limit": 10}} - assert gate.over_ceiling({"reportAny": 10}, budget) == frozenset() - assert gate.over_ceiling({"reportAny": 11}, budget) == frozenset({"reportAny"}) - assert gate.over_ceiling({}, budget) == frozenset() - - -def test_over_ceiling_holds_unbudgeted_rules_to_the_default_limit(): - assert gate.over_ceiling({"brand-new": gate.DEFAULT_LIMIT}, {}) == frozenset() - assert gate.over_ceiling({"brand-new": gate.DEFAULT_LIMIT + 1}, {}) == frozenset( - {"brand-new"} - ) - - -def test_over_ceiling_is_independent_across_rules(): - budget = {"reportAny": {"limit": 10}, "reportArgumentType": {"limit": 5}} - assert gate.over_ceiling( - {"reportAny": 9, "reportArgumentType": 6}, budget - ) == frozenset({"reportArgumentType"}) - - -def test_cache_key_changes_with_base_point_and_each_fingerprint(): - key = gate.cache_key("abc", ("cfg", "lock")) - assert gate.cache_key("abc", ("cfg", "lock")) == key - assert gate.cache_key("def", ("cfg", "lock")) != key - assert gate.cache_key("abc", ("cfg2", "lock")) != key - assert gate.cache_key("abc", ("cfg", "lock2")) != key - - def test_fingerprints_carry_the_dependency_group_set(): # Counts measured under one group set must never be compared against # another's: the fingerprint difference re-keys every cache entry and @@ -350,381 +225,42 @@ def test_ensure_env_is_silent_when_the_env_already_exists(tmp_path, capsys): assert capsys.readouterr().err == "" -def test_cached_counts_round_trip(tmp_path): - path = gate.cache_path(tmp_path, "abc123", ("f1", "f2")) - gate.store_counts(tmp_path, path, "abc123", {"reportAny": 3, "reportCall": 1}) - assert gate.load_cached_counts(path) == {"reportAny": 3, "reportCall": 1} - - -def test_missing_corrupt_or_misshapen_cache_reads_as_none(tmp_path): - path = tmp_path / "cache.json" - assert gate.load_cached_counts(path) is None - path.write_text("{not json") - assert gate.load_cached_counts(path) is None - path.write_text(json.dumps(["counts"])) - assert gate.load_cached_counts(path) is None - path.write_text(json.dumps({"base_point": "abc"})) - assert gate.load_cached_counts(path) is None - path.write_text(json.dumps({"counts": {"reportAny": "three"}})) - assert gate.load_cached_counts(path) is None - path.write_text(json.dumps({"counts": {"reportAny": True}})) - assert gate.load_cached_counts(path) is None - - -def test_scratch_is_invisible_to_the_prune_glob(): - import fnmatch - - scratch = gate.scratch_path(gate.cache_path(Path("/c"), "abc", ("f",))) - assert not fnmatch.fnmatch(scratch.name, f"{gate.CACHE_FILE_PREFIX}*") - - -def test_store_prune_spares_a_concurrent_runs_in_flight_scratch(tmp_path): - foreign = gate.scratch_path(gate.cache_path(tmp_path, "other", ("f",))) - foreign.parent.mkdir(parents=True, exist_ok=True) - foreign.write_text("{}") - mine = gate.cache_path(tmp_path, "mine", ("f",)) - gate.store_counts(tmp_path, mine, "mine", {"reportAny": 1}) - assert foreign.exists() - assert gate.load_cached_counts(mine) == {"reportAny": 1} - - -def test_store_keeps_a_concurrent_worktrees_entry_for_another_branch_point(tmp_path): - old = gate.cache_path(tmp_path, "old", ("f",)) - gate.store_counts(tmp_path, old, "old", {"reportAny": 1}) - new = gate.cache_path(tmp_path, "new", ("f",)) - gate.store_counts(tmp_path, new, "new", {"reportAny": 2}) - assert gate.load_cached_counts(old) == {"reportAny": 1} - assert gate.load_cached_counts(new) == {"reportAny": 2} - - -def test_store_evicts_only_the_oldest_entries_beyond_the_cap(tmp_path): - aged = [ - gate.cache_path(tmp_path, f"base{i}", ("f",)) - for i in range(gate.CACHE_KEEP_ENTRIES) - ] - for age, path in enumerate(aged): - gate.store_counts(tmp_path, path, f"base{age}", {"reportAny": age}) - os.utime(path, (age, age)) - newest = gate.cache_path(tmp_path, "newest", ("f",)) - gate.store_counts(tmp_path, newest, "newest", {"reportAny": 99}) - assert not aged[0].exists() - assert all(path.exists() for path in aged[1:]) - assert gate.load_cached_counts(newest) == {"reportAny": 99} - - -def test_store_never_evicts_the_entry_it_just_wrote_even_on_mtime_ties(tmp_path): - others = [ - gate.cache_path(tmp_path, f"base{i}", ("f",)) - for i in range(gate.CACHE_KEEP_ENTRIES + 2) - ] - for path in others: - gate.store_counts(tmp_path, path, path.name, {"reportAny": 1}) - os.utime(path, (9_999_999_999, 9_999_999_999)) - mine = gate.cache_path(tmp_path, "mine", ("f",)) - gate.store_counts(tmp_path, mine, "mine", {"reportAny": 2}) - assert gate.load_cached_counts(mine) == {"reportAny": 2} - survivors = list(tmp_path.glob(f"{gate.CACHE_FILE_PREFIX}*.json")) - assert len(survivors) == gate.CACHE_KEEP_ENTRIES - - -def _no_fetch(ref): - return None - - -def _never(reason): - def callback(ref): - raise AssertionError(reason) - - return callback - - -def test_base_counts_cached_returns_the_hit_without_recomputing(tmp_path): - path = gate.cache_path(tmp_path, "abc123", gate.environment_fingerprints()) - gate.store_counts(tmp_path, path, "abc123", {"reportAny": 7}) - - assert gate.base_counts_cached( - "abc123", - cache_dir=tmp_path, - compute=_never("a cache hit must not re-run the base pass"), - fetch=_never("a cache hit must not reach for CI"), - ) == {"reportAny": 7} - - -def test_base_counts_cached_computes_once_then_hits(tmp_path): - calls = [] - - def fake(ref): - calls.append(ref) - return {"reportAny": 4} - - first = gate.base_counts_cached( - "abc123", cache_dir=tmp_path, compute=fake, fetch=_no_fetch - ) - second = gate.base_counts_cached( - "abc123", cache_dir=tmp_path, compute=fake, fetch=_no_fetch - ) - assert first == second == {"reportAny": 4} - assert calls == ["abc123"] - - -def test_an_empty_base_pass_is_never_cached(tmp_path): - calls = [] - - def crashed(ref): - calls.append(ref) - return {} - - assert ( - gate.base_counts_cached( - "abc123", cache_dir=tmp_path, compute=crashed, fetch=_no_fetch - ) - == {} - ) - assert ( - gate.base_counts_cached( - "abc123", cache_dir=tmp_path, compute=crashed, fetch=_no_fetch - ) - == {} - ) - assert calls == ["abc123", "abc123"] - assert list(tmp_path.iterdir()) == [] - - -def test_base_counts_cached_uses_fetched_counts_and_persists_them(tmp_path): - counts = gate.base_counts_cached( - "abc123", - cache_dir=tmp_path, - compute=_never("fetched counts must skip the local base pass"), - fetch=lambda ref: {"reportAny": 9}, - ) - assert counts == {"reportAny": 9} - path = gate.cache_path(tmp_path, "abc123", gate.environment_fingerprints()) - assert gate.load_cached_counts(path) == {"reportAny": 9} - assert gate.base_counts_cached( - "abc123", - cache_dir=tmp_path, - compute=_never("the persisted fetch must satisfy later runs"), - fetch=_never("the persisted fetch must satisfy later runs"), - ) == {"reportAny": 9} - - -def test_base_counts_cached_falls_back_to_compute_on_a_fetch_miss(tmp_path): - calls = [] - - def local(ref): - calls.append(ref) - return {"reportAny": 4} - - assert gate.base_counts_cached( - "abc123", cache_dir=tmp_path, compute=local, fetch=_no_fetch - ) == {"reportAny": 4} - assert calls == ["abc123"] - - -def test_base_counts_cached_treats_empty_fetched_counts_as_a_miss(tmp_path): - assert gate.base_counts_cached( - "abc123", - cache_dir=tmp_path, - compute=lambda ref: {"reportAny": 2}, - fetch=lambda ref: {}, - ) == {"reportAny": 2} - path = gate.cache_path(tmp_path, "abc123", gate.environment_fingerprints()) - assert gate.load_cached_counts(path) == {"reportAny": 2} - - -def test_origin_slug_parsing_supports_ssh_and_https_github_forms(): - assert gate.parse_origin_slug("git@github.com:BerriAI/litellm.git") == "BerriAI/litellm" - assert gate.parse_origin_slug("git@github.com:BerriAI/litellm") == "BerriAI/litellm" - assert gate.parse_origin_slug("https://github.com/BerriAI/litellm.git") == "BerriAI/litellm" - assert gate.parse_origin_slug("https://github.com/BerriAI/litellm") == "BerriAI/litellm" - assert gate.parse_origin_slug("https://github.com/BerriAI/litellm/") == "BerriAI/litellm" - - -def test_origin_slug_parsing_rejects_non_github_urls(): - assert gate.parse_origin_slug("https://gitlab.com/BerriAI/litellm.git") is None - assert gate.parse_origin_slug("git@bitbucket.org:BerriAI/litellm.git") is None - assert gate.parse_origin_slug("not a url") is None - assert gate.parse_origin_slug("") is None - - -def _artifact_zip(payload): - import io - import zipfile - - buffer = io.BytesIO() - with zipfile.ZipFile(buffer, "w") as archive: - archive.writestr("basedpyright-counts.json", json.dumps(payload)) - return buffer.getvalue() - - -def _gh_stub(listing, zip_bytes): - def gh_output(args): - if args[-1].startswith("repos/"): - return json.dumps(listing).encode() - return zip_bytes - - return gh_output - - -def _live_listing(): - return { - "artifacts": [ - {"expired": False, "archive_download_url": "https://api.github.com/x/zip"} - ] - } - - -def test_fetcher_returns_counts_from_a_matching_artifact(capsys): - payload = {"base_point": "abc123", "counts": {"reportAny": 3}} - fetched = gate.fetch_ci_base_counts( - "abc123", gh_output=_gh_stub(_live_listing(), _artifact_zip(payload)) - ) - assert fetched == {"reportAny": 3} - assert "fetched from CI artifact" in capsys.readouterr().err - - -def test_fetcher_rejects_an_artifact_for_a_different_base_point(): - payload = {"base_point": "someothersha", "counts": {"reportAny": 3}} - assert ( - gate.fetch_ci_base_counts( - "abc123", gh_output=_gh_stub(_live_listing(), _artifact_zip(payload)) - ) - is None - ) - - -def test_fetcher_rejects_empty_or_misshapen_artifact_counts(): - for counts in ({}, {"reportAny": "three"}, {"reportAny": True}): - payload = {"base_point": "abc123", "counts": counts} - assert ( - gate.fetch_ci_base_counts( - "abc123", gh_output=_gh_stub(_live_listing(), _artifact_zip(payload)) - ) - is None - ) - - -def test_fetcher_rejects_an_expired_artifact(): - listing = { - "artifacts": [ - {"expired": True, "archive_download_url": "https://api.github.com/x/zip"} - ] - } - payload = {"base_point": "abc123", "counts": {"reportAny": 3}} - assert ( - gate.fetch_ci_base_counts( - "abc123", gh_output=_gh_stub(listing, _artifact_zip(payload)) - ) - is None - ) - - -def test_fetcher_misses_when_no_artifact_is_published(): - assert ( - gate.fetch_ci_base_counts( - "abc123", gh_output=_gh_stub({"artifacts": []}, b"") - ) - is None - ) - - -def test_fetcher_misses_when_gh_is_unusable(capsys): - assert gate.fetch_ci_base_counts("abc123", gh_output=lambda args: None) is None - assert "computing base counts locally" in capsys.readouterr().err - - -def test_fetcher_misses_on_a_corrupt_artifact_archive(): - assert ( - gate.fetch_ci_base_counts( - "abc123", gh_output=_gh_stub(_live_listing(), b"not a zip") - ) - is None - ) - - -def test_emit_writes_the_artifact_json_named_by_the_head_key(tmp_path, capsys): - gate.cmd_emit_counts({"reportAny": 3, "aRule": 1}, tmp_path, "deadbeef") - key = gate.cache_key("deadbeef", gate.environment_fingerprints()) - path = tmp_path / f"basedpyright-counts-{key}.json" - assert json.loads(path.read_text()) == { - "base_point": "deadbeef", - "counts": {"aRule": 1, "reportAny": 3}, - } - summary = capsys.readouterr().out - assert "deadbeef" in summary - assert key in summary - assert "4" in summary - - -def test_emit_refuses_to_publish_empty_counts(tmp_path): - import pytest - +def test_changing_the_dependency_groups_rekeys_the_base_counts() -> None: + default: Final = gate.checker_identity().artifact_name("abc123") + assert gate.checker_identity().artifact_name("abc123") == default + assert gate.checker_identity(("proxy-dev",)).artifact_name("abc123") != default + + +@pytest.mark.parametrize("rule", ["reportAny", "reportExplicitAny"]) +def test_any_rules_may_grow_past_their_base_up_to_the_cap_and_no_further( + rule: str, capsys: pytest.CaptureFixture[str] +) -> None: + cap: Final = gate.ANY_CAPS[rule] + base: Final = {rule: cap - 5, "reportArgumentType": 3} + gate.judge({rule: cap, "reportArgumentType": 3}, base, "a" * 40) + assert "OK" in capsys.readouterr().out + with pytest.raises(SystemExit) as exit_info: + gate.judge({rule: cap + 1, "reportArgumentType": 3}, base, "a" * 40) + assert exit_info.value.code == 1 + assert f"BREACHED RULES: {rule} {cap + 1}/{cap} (+6)" in capsys.readouterr().out + + +def test_an_any_rule_already_over_its_cap_at_base_does_not_fail_a_bystander( + capsys: pytest.CaptureFixture[str], +) -> None: + over: Final = gate.ANY_CAPS["reportAny"] + 50 + gate.judge({"reportAny": over}, {"reportAny": over}, "a" * 40) + assert "OK" in capsys.readouterr().out + + +def test_rules_without_a_cap_may_not_grow_past_their_base(capsys: pytest.CaptureFixture[str]) -> None: with pytest.raises(SystemExit): - gate.cmd_emit_counts({}, tmp_path, "deadbeef") - assert list(tmp_path.iterdir()) == [] + gate.judge({"reportArgumentType": 4}, {"reportArgumentType": 3}, "a" * 40) + assert "BREACHED RULES: reportArgumentType 4/3 (+1)" in capsys.readouterr().out -def test_emitted_file_round_trips_through_the_fetch_validation(tmp_path): - gate.cmd_emit_counts({"reportAny": 3}, tmp_path, "deadbeef") - key = gate.cache_key("deadbeef", gate.environment_fingerprints()) - payload = json.loads((tmp_path / f"basedpyright-counts-{key}.json").read_text()) - assert gate.counts_for_base(payload, "deadbeef") == {"reportAny": 3} - assert gate.counts_for_base(payload, "someothersha") is None - - -def _git(cwd, *args): - proc = subprocess.run(["git", *args], cwd=cwd, capture_output=True, text=True) - assert proc.returncode == 0, proc.stderr - return proc.stdout.strip() - - -def _commit(cwd, name): - (cwd / name).write_text(name) - _git(cwd, "add", "-A") - _git(cwd, "commit", "-q", "-m", name) - return _git(cwd, "rev-parse", "HEAD") - - -def _init_repo(tmp_path): - repo = tmp_path / "repo" - repo.mkdir() - _git(repo, "init", "-q", "-b", "main") - _git(repo, "config", "user.email", "gate@example.com") - _git(repo, "config", "user.name", "gate") - _git(repo, "config", "commit.gpgsign", "false") - return repo - - -def _branched_repo(tmp_path): - repo = _init_repo(tmp_path) - branch_point = _commit(repo, "shared.txt") - _git(repo, "checkout", "-q", "-b", "feature") - _commit(repo, "feature.txt") - _git(repo, "checkout", "-q", "main") - base_tip = _commit(repo, "drift.txt") - _git(repo, "checkout", "-q", "feature") - return repo, branch_point, base_tip - - -def test_base_point_is_the_branch_point_when_no_merge_is_in_progress(tmp_path): - repo, branch_point, _ = _branched_repo(tmp_path) - assert gate.resolve_base_point("main", cwd=repo) == branch_point - - -def test_base_point_mid_merge_advances_to_the_merged_in_base_tip(tmp_path): - repo, _, base_tip = _branched_repo(tmp_path) - _git(repo, "merge", "--no-commit", "--no-ff", "main") - assert gate.resolve_base_point("main", cwd=repo) == base_tip - - -def test_base_point_mid_merge_of_an_older_side_branch_keeps_the_newer_branch_point(tmp_path): - repo = _init_repo(tmp_path) - _commit(repo, "shared.txt") - _git(repo, "checkout", "-q", "-b", "old-side") - _commit(repo, "old.txt") - _git(repo, "checkout", "-q", "main") - newer_point = _commit(repo, "drift.txt") - _git(repo, "checkout", "-q", "-b", "feature") - _commit(repo, "feature.txt") - _git(repo, "merge", "--no-commit", "--no-ff", "old-side") - assert gate.resolve_base_point("main", cwd=repo) == newer_point +def test_no_head_output_is_refused_as_vacuous_before_any_base_lookup(capsys: pytest.CaptureFixture[str]) -> None: + with pytest.raises(SystemExit) as exit_info: + gate.cmd_check({}, "irrelevant-base-ref") + assert exit_info.value.code == 1 + assert "vacuous" in capsys.readouterr().out diff --git a/tests/unit/test_type_discipline_gate.py b/tests/unit/test_type_discipline_gate.py index 1832668e7c3..e99a7be4359 100644 --- a/tests/unit/test_type_discipline_gate.py +++ b/tests/unit/test_type_discipline_gate.py @@ -1,106 +1,64 @@ """Tests for scripts/type_discipline_gate.py. -The gate's correctness lives in two pure functions: `over_ceiling` (which decides -whether the expensive base worktree scan is even needed) and `evaluate` (the -drift-safe breach check). Both are pinned here. +The gate compares each LIT rule's codebase count against the merge-base count, so +what is pinned here is the identity that keys those base counts and the diff scan +that turns a breach into file:line. """ -import importlib.util -import subprocess from pathlib import Path +from typing import Final -_MODULE_PATH = Path(__file__).resolve().parents[2] / "scripts" / "type_discipline_gate.py" -_spec = importlib.util.spec_from_file_location("type_discipline_gate", _MODULE_PATH) -gate = importlib.util.module_from_spec(_spec) -_spec.loader.exec_module(gate) +import type_discipline_gate as gate -def _budget(limit): - return {"LIT006": {"limit": limit}} +def test_editing_the_checker_rekeys_the_base_counts(tmp_path: Path) -> None: + checker: Final = tmp_path / "check.py" + checker.write_text("print('v1')\n") + before: Final = gate.checker_identity(checker).artifact_name("abc123") + assert gate.checker_identity(checker).artifact_name("abc123") == before + checker.write_text("print('v2')\n") + assert gate.checker_identity(checker).artifact_name("abc123") != before -def test_over_ceiling_flags_only_counts_above_the_limit(): - budget = _budget(12) - assert gate.over_ceiling({"LIT006": 12}, budget) == frozenset() # at limit - assert gate.over_ceiling({"LIT006": 13}, budget) == frozenset({"LIT006"}) # over limit - assert gate.over_ceiling({}, budget) == frozenset() # missing rule counts as zero +def test_parse_changed_lines_groups_hunks_under_their_own_file() -> None: + diff: Final = ( + "diff --git a/litellm/a.py b/litellm/a.py\n" + "--- a/litellm/a.py\n" + "+++ b/litellm/a.py\n" + "@@ -0,0 +3,2 @@\n" + "+one\n" + "+two\n" + "diff --git a/litellm/b.py b/litellm/b.py\n" + "--- a/litellm/b.py\n" + "+++ b/litellm/b.py\n" + "@@ -0,0 +10 @@\n" + "+only\n" + ) + changed: Final = gate.parse_changed_lines(diff) + assert changed["litellm/a.py"] == {3, 4} + assert changed["litellm/b.py"] == {10} -def test_over_ceiling_is_independent_across_rules(): - budget = {"LIT001": {"limit": 5}, "LIT006": {"limit": 10}} - assert gate.over_ceiling({"LIT001": 6, "LIT006": 10}, budget) == frozenset({"LIT001"}) +def test_parse_changed_lines_handles_several_hunks_in_one_file() -> None: + diff: Final = ( + "+++ b/litellm/a.py\n" + "@@ -0,0 +1,2 @@\n" + "+a\n" + "@@ -9,0 +20,1 @@\n" + "+b\n" + ) + assert gate.parse_changed_lines(diff)["litellm/a.py"] == {1, 2, 20} -def test_evaluate_blames_only_a_rule_over_limit_and_over_base(): - budget = _budget(10) - # over limit and grown vs base -> breach - assert [b.rule for b in gate.evaluate({"LIT006": 12}, {"LIT006": 9}, budget)] == ["LIT006"] - # over limit but flat vs base (pre-existing drift) -> not blamed - assert gate.evaluate({"LIT006": 12}, {"LIT006": 12}, budget) == [] - # within limit -> not blamed regardless of base - assert gate.evaluate({"LIT006": 10}, {"LIT006": 0}, budget) == [] +def test_parse_changed_lines_on_an_empty_diff_is_empty() -> None: + assert gate.parse_changed_lines("") == {} -def test_update_ratchets_limit_down_by_what_the_branch_fixed_never_up(): - budget = {"LIT001": {"limit": 100}, "LIT006": {"limit": 10}} - # LIT001 fixed 15 (60 -> 45) so its limit falls 100 -> 85; LIT006 grew, so its - # limit holds flat at 10. - current = {"LIT001": 45, "LIT006": 12} - base = {"LIT001": 60, "LIT006": 9} - assert gate.ratcheted_budget(budget, current, base) == { - "LIT001": {"limit": 85}, - "LIT006": {"limit": 10}, - } - - -def test_update_leaves_rules_seeded_on_this_branch_untouched(): - # A rule absent from the base budget was seeded with grandfathered headroom on - # this branch; the base tree predates the rule (e.g. no Final annotations yet), - # so ratcheting against it would collapse the deliberate headroom. - budget = {"LIT001": {"limit": 100}, "LIT010": {"limit": 24600}} - current = {"LIT001": 45, "LIT010": 16400} - base = {"LIT001": 60, "LIT010": 40000} - assert gate.ratcheted_budget(budget, current, base, frozenset({"LIT010"})) == { - "LIT001": {"limit": 85}, - "LIT010": {"limit": 24600}, - } - - -def _git(cwd, *args): - proc = subprocess.run(["git", *args], cwd=cwd, capture_output=True, text=True) - assert proc.returncode == 0, proc.stderr - return proc.stdout.strip() - - -def _commit(cwd, name): - (cwd / name).write_text(name) - _git(cwd, "add", "-A") - _git(cwd, "commit", "-q", "-m", name) - return _git(cwd, "rev-parse", "HEAD") - - -def _branched_repo(tmp_path): - repo = tmp_path / "repo" - repo.mkdir() - _git(repo, "init", "-q", "-b", "main") - _git(repo, "config", "user.email", "gate@example.com") - _git(repo, "config", "user.name", "gate") - _git(repo, "config", "commit.gpgsign", "false") - branch_point = _commit(repo, "shared.txt") - _git(repo, "checkout", "-q", "-b", "feature") - _commit(repo, "feature.txt") - _git(repo, "checkout", "-q", "main") - base_tip = _commit(repo, "drift.txt") - _git(repo, "checkout", "-q", "feature") - return repo, branch_point, base_tip - - -def test_base_point_is_the_branch_point_when_no_merge_is_in_progress(tmp_path): - repo, branch_point, _ = _branched_repo(tmp_path) - assert gate.resolve_base_point("main", cwd=repo) == branch_point - - -def test_base_point_mid_merge_advances_to_the_merged_in_base_tip(tmp_path): - repo, _, base_tip = _branched_repo(tmp_path) - _git(repo, "merge", "--no-commit", "--no-ff", "main") - assert gate.resolve_base_point("main", cwd=repo) == base_tip +def test_introduced_keeps_only_violations_on_changed_lines() -> None: + violations: Final = ( + gate.Violation("litellm/a.py", 3, "LIT006"), + gate.Violation("litellm/a.py", 99, "LIT006"), + gate.Violation("litellm/b.py", 3, "LIT001"), + ) + kept: Final = gate.introduced(violations, {"litellm/a.py": {3}}) + assert kept == [gate.Violation("litellm/a.py", 3, "LIT006")] diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 7551a310977..58fd2250d2a 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -37,6 +37,10 @@ from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.humanloop import HumanloopLogger +from litellm.llms.base_llm.audio_transcription.transformation import BaseAudioTranscriptionConfig +from litellm.integrations.langfuse.langfuse_prompt_management import LangfusePromptManagement +from litellm.litellm_core_utils import litellm_logging from litellm.litellm_core_utils.duration_parser import ( _extract_from_regex, duration_in_seconds, @@ -50,7 +54,7 @@ from litellm.proxy.utils import is_valid_api_key from litellm.types.caching import CachingSupportedCallTypes from litellm.types.integrations.custom_logger import HEADROOM_CONVERTED_STREAM_KEY from litellm.types.llms.openai import ResponsesAPIResponse -from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams +from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams, LiteLLM_Params from litellm.types.utils import ( ADDRESSED_RESPONSE_ID_FIELD, CallTypes, @@ -74,6 +78,7 @@ from litellm.types.utils import ( from litellm.types.videos.main import VideoObject from litellm.utils import ( _invalidate_model_cost_lowercase_map, + add_custom_logger_callback_to_specific_event, check_valid_key, CustomStreamWrapper, filter_out_litellm_params, @@ -325,6 +330,25 @@ def test_supports_function_calling_unknown_github_alias_returns_false(): assert litellm.utils.supports_function_calling(model="github/non-existent-model-for-capability-check") is False +@pytest.mark.parametrize(("audio_input", "audio_output"), [(True, False), (False, True)]) +def test_supports_audio_output_reads_its_own_cost_map_flag( + local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch, audio_input: bool, audio_output: bool +) -> None: + model: Final = "openai/audio-flags-disagree-model" + monkeypatch.setitem( + litellm.model_cost, + model, + { + "litellm_provider": "openai", + "mode": "chat", + "supports_audio_input": audio_input, + "supports_audio_output": audio_output, + }, + ) + assert litellm.supports_audio_input(model) is audio_input + assert litellm.supports_audio_output(model) is audio_output + + def test_get_optional_params_image_gen(): from litellm.llms.azure.image_generation import AzureGPTImageGenerationConfig @@ -996,6 +1020,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "supports_web_search": {"type": "boolean"}, "supports_bedrock_runtime_chat_completions_tools_with_reasoning": {"type": "boolean"}, "supports_bedrock_runtime_chat_completions_response_format": {"type": "boolean"}, + "supports_bedrock_runtime_chat_completions_inline_reasoning": {"type": "boolean"}, "supports_url_context": {"type": "boolean"}, "supports_multimodal": {"type": "boolean"}, "uses_embed_content": {"type": "boolean"}, @@ -6783,6 +6808,7 @@ def setup_and_teardown(): MODEL: Final = "anthropic/claude-haiku-4-5" + @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") def test_validate_tool_choice_none(): """Test that None is returned as-is.""" @@ -6927,6 +6953,7 @@ _SCALAR_DEFAULTS = { "api_key": getattr(litellm, "api_key", None), } + @pytest.fixture(scope="module") def setup_and_teardown_local_testing(): """ @@ -8767,6 +8794,407 @@ def test_get_valid_models_from_provider(): assert "gpt-5-mini" in valid_models +@pytest.mark.parametrize( + ("provider", "api_base", "api_key", "model_id"), + [ + ("anthropic", "https://anthropic.models.test", "anthropic-test-key", "claude-test-model"), + ("xai", "https://xai.models.test", "xai-test-key", "grok-test-model"), + ], +) +def test_get_valid_models_discovers_provider_models_from_http( + provider: str, + api_base: str, + api_key: str, + model_id: str, +) -> None: + response_body: Final = { + "data": [ + { + "id": model_id, + "type": "model", + "display_name": "Test model", + "created_at": "2024-01-01T00:00:00Z", + } + ], + "has_more": False, + "first_id": model_id, + "last_id": model_id, + } + models_url: Final = f"{api_base}/v1/models" + upstream: Final[respx.MockRouter] + + with respx.mock(assert_all_called=True) as upstream: + model_list_route: Final = upstream.get(models_url).respond(200, json=response_body) + + discovered_models: Final = get_valid_models( + check_provider_endpoint=True, + custom_llm_provider=provider, + api_key=api_key, + api_base=api_base, + ) + + assert discovered_models == [f"{provider}/{model_id}"] + assert model_list_route.called + assert len(upstream.calls) == 1 + + +def test_check_valid_key_returns_false_for_http_unauthorized() -> None: + response_body: Final = { + "error": { + "message": "Invalid API key", + "type": "invalid_request_error", + "param": None, + "code": "invalid_api_key", + } + } + upstream: Final[respx.MockRouter] + + with respx.mock(assert_all_called=True) as upstream: + invalid_key_route: Final = upstream.post("https://api.openai.com/v1/chat/completions").respond( + 401, + json=response_body, + ) + + valid_key: Final = check_valid_key(model="gpt-5-mini", api_key="invalid-test-key") + + assert valid_key is False + assert invalid_key_route.called + assert upstream.calls.last.request.headers["Authorization"] == "Bearer invalid-test-key" + + +def test_check_valid_key_returns_true_for_successful_completion() -> None: + response_body: Final = { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 0, + "model": "gpt-5-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } + upstream: Final[respx.MockRouter] + + with respx.mock(assert_all_called=True) as upstream: + valid_key_route: Final = upstream.post("https://api.openai.com/v1/chat/completions").respond( + 200, + json=response_body, + ) + + valid_key: Final = check_valid_key(model="gpt-5-mini", api_key="valid-test-key") + + assert valid_key is True + assert valid_key_route.called + assert upstream.calls.last.request.headers["Authorization"] == "Bearer valid-test-key" + + +def test_function_to_dict_parses_numpy_docstring_schema() -> None: + pytest.importorskip("numpydoc") + + def get_current_weather(location: str, unit: str) -> str: + """Get the current weather in a given location + + Parameters + ---------- + location : str + The city and state, e.g. San Francisco, CA + unit : {'celsius', 'fahrenheit'} + Temperature unit + + Returns + ------- + str + A sentence indicating the weather + """ + return f"Weather for {location} in {unit}" + + schema: Final = litellm.utils.function_to_dict(get_current_weather) + + assert schema["name"] == "get_current_weather" + assert schema["description"] == "Get the current weather in a given location" + assert schema["parameters"]["type"] == "object" + assert schema["parameters"]["properties"]["location"] == { + "type": "string", + "description": "The city and state, e.g. San Francisco, CA", + } + assert schema["parameters"]["properties"]["unit"]["type"] == "string" + assert schema["parameters"]["properties"]["unit"]["description"] == "Temperature unit" + assert schema["parameters"]["required"] == ["location", "unit"] + + +def test_duration_in_seconds_one_month_uses_the_fixed_calendar_interval( + monkeypatch: pytest.MonkeyPatch, +) -> None: + fixed_start: Final = datetime(2025, 2, 15, 12, 0, 0, 123456) + fixed_timestamp: Final = fixed_start.timestamp() + monkeypatch.setattr( + "litellm.litellm_core_utils.duration_parser.time_module.time", + lambda: fixed_timestamp, + ) + + assert duration_in_seconds("1mo") == 28 * 24 * 60 * 60 + + +def test_prompt_caching_image_check_uses_default_image_dimensions() -> None: + image_bytes: Final = base64.b64decode( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/x+AAwMCAO+ip1sAAAAASUVORK5CYII=" + ) + image_url: Final = "https://93.184.216.34/test.png" + messages: Final = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is in this image?"}, + { + "type": "image_url", + "image_url": {"url": image_url, "detail": "high"}, + }, + ], + } + ] + + with respx.mock(assert_all_called=False) as upstream: + image_route: Final = upstream.get(image_url).respond(200, content=image_bytes) + cacheable: Final = is_prompt_caching_valid_prompt( + model="gpt-4o-mini", + messages=messages, + custom_llm_provider="openai", + min_token_count=100_000, + ) + + assert cacheable is False + assert image_route.called is False + assert len(upstream.calls) == 0 + + +def test_get_valid_models_discovers_fireworks_models_from_http( + monkeypatch: pytest.MonkeyPatch, +) -> None: + api_base: Final = "https://fireworks.models.test/v1" + api_key: Final = "fireworks-test-key" + account_id: Final = "fireworks-test-account" + model_name: Final = "accounts/fireworks/models/llama-test-model" + models_url: Final = f"https://fireworks.models.test/v1/accounts/{account_id}/models" + monkeypatch.setenv("FIREWORKS_ACCOUNT_ID", account_id) + upstream: Final[respx.MockRouter] + + with respx.mock(assert_all_called=True) as upstream: + model_list_route: Final = upstream.get(models_url).respond( + 200, + json={"models": [{"name": model_name}]}, + ) + + discovered_models: Final = get_valid_models( + check_provider_endpoint=True, + custom_llm_provider="fireworks_ai", + api_key=api_key, + api_base=api_base, + ) + + assert discovered_models == [f"fireworks_ai/{model_name}"] + assert model_list_route.called + assert upstream.calls.last.request.headers["Authorization"] == f"Bearer {api_key}" + + +def test_get_valid_models_returns_static_fireworks_models_without_endpoint_check( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm, "check_provider_endpoint", False) + monkeypatch.setenv("FIREWORKS_AI_API_KEY", "fireworks-test-key") + expected_models: Final = litellm.models_by_provider["fireworks_ai"] + + actual_models: Final = get_valid_models( + check_provider_endpoint=False, + custom_llm_provider="fireworks_ai", + ) + env_inferred_models: Final = get_valid_models() + + assert set(actual_models) == expected_models + assert actual_models + assert expected_models <= set(env_inferred_models) + + +def test_get_valid_models_uses_the_litellm_params_anthropic_api_key() -> None: + model_id: Final = "claude-test-model" + models_url: Final = "https://api.anthropic.com/v1/models" + response_body: Final = { + "data": [ + { + "id": model_id, + "type": "model", + "display_name": "Test Claude", + "created_at": "2024-01-01T00:00:00Z", + } + ], + "has_more": False, + "first_id": model_id, + "last_id": model_id, + } + upstream: Final[respx.MockRouter] + + def response_for_api_key(request: httpx.Request) -> httpx.Response: + if request.headers["x-api-key"] == "bad-test-key": + return httpx.Response(401, json={"error": {"message": "invalid key"}}, request=request) + return httpx.Response(200, json=response_body, request=request) + + with respx.mock(assert_all_called=True) as upstream: + model_list_route: Final = upstream.get(models_url).mock(side_effect=response_for_api_key) + + bad_key_models: Final = get_valid_models( + check_provider_endpoint=True, + custom_llm_provider="anthropic", + litellm_params=LiteLLM_Params(model="anthropic/*", api_key="bad-test-key"), + ) + good_key_models: Final = get_valid_models( + check_provider_endpoint=True, + custom_llm_provider="anthropic", + litellm_params=LiteLLM_Params(model="anthropic/*", api_key="good-test-key"), + ) + + assert bad_key_models == [] + assert good_key_models == [f"anthropic/{model_id}"] + assert model_list_route.called + assert [call.request.headers["x-api-key"] for call in upstream.calls] == [ + "bad-test-key", + "good-test-key", + ] + + +def test_add_custom_logger_to_success_callback_registers_once( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "_async_success_callback", []) + monkeypatch.setattr(litellm, "_async_failure_callback", []) + monkeypatch.setattr(litellm, "success_callback", []) + monkeypatch.setattr(litellm, "failure_callback", []) + monkeypatch.setattr(litellm_logging, "_in_memory_loggers", []) + + add_custom_logger_callback_to_specific_event("langfuse", "success") + + assert len(litellm.success_callback) == 1 + assert isinstance(litellm.success_callback[0], LangfusePromptManagement) + assert len(litellm._async_success_callback) == 1 + assert isinstance(litellm._async_success_callback[0], LangfusePromptManagement) + assert litellm.failure_callback == [] + assert litellm._async_failure_callback == [] + + +@pytest.mark.parametrize( + "registered_lists", + [ + ("success_callback", "_async_success_callback"), + ("success_callback",), + ], +) +def test_add_custom_logger_callback_does_not_duplicate_existing_success_logger( + monkeypatch: pytest.MonkeyPatch, + registered_lists: tuple[str, ...], +) -> None: + logger: Final = HumanloopLogger() + async_success_callbacks: Final = [logger] if "_async_success_callback" in registered_lists else [] + success_callbacks: Final = [logger] if "success_callback" in registered_lists else [] + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "_async_success_callback", async_success_callbacks) + monkeypatch.setattr(litellm, "_async_failure_callback", []) + monkeypatch.setattr(litellm, "success_callback", success_callbacks) + monkeypatch.setattr(litellm, "failure_callback", []) + monkeypatch.setattr(litellm_logging, "_in_memory_loggers", []) + + add_custom_logger_callback_to_specific_event("humanloop", "success") + + assert sum(type(callback) is HumanloopLogger for callback in litellm.success_callback) == int( + "success_callback" in registered_lists + ) + assert sum(type(callback) is HumanloopLogger for callback in litellm._async_success_callback) == int( + "_async_success_callback" in registered_lists + ) + assert litellm.failure_callback == [] + assert litellm._async_failure_callback == [] + + +@pytest.mark.parametrize( + ("registered_lists", "expected_async_success_callback_count"), + [ + (("success_callback", "_async_success_callback"), 1), + (("success_callback",), 0), + ], +) +@pytest.mark.asyncio +async def test_acompletion_does_not_duplicate_a_logger_already_in_success_callbacks( + monkeypatch: pytest.MonkeyPatch, + registered_lists: tuple[str, ...], + expected_async_success_callback_count: int, +) -> None: + logger: Final = HumanloopLogger() + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "input_callback", []) + monkeypatch.setattr(litellm, "success_callback", [logger]) + monkeypatch.setattr(litellm, "failure_callback", []) + monkeypatch.setattr(litellm, "_async_input_callback", []) + monkeypatch.setattr( + litellm, "_async_success_callback", [logger] if "_async_success_callback" in registered_lists else [] + ) + monkeypatch.setattr(litellm, "_async_failure_callback", []) + + await litellm.acompletion( + model="gpt-5-mini", + messages=[{"role": "user", "content": "callback registration"}], + mock_response="ok", + ) + + assert litellm.success_callback == [logger] + assert litellm._async_success_callback == [logger] * expected_async_success_callback_count + + +@pytest.mark.asyncio +async def test_custom_logger_in_global_callbacks_registers_once_across_completion_calls( + monkeypatch: pytest.MonkeyPatch, +) -> None: + logger: Final = HumanloopLogger() + monkeypatch.setattr(litellm, "callbacks", [logger]) + monkeypatch.setattr(litellm, "input_callback", []) + monkeypatch.setattr(litellm, "success_callback", []) + monkeypatch.setattr(litellm, "failure_callback", []) + monkeypatch.setattr(litellm, "_async_input_callback", []) + monkeypatch.setattr(litellm, "_async_success_callback", []) + monkeypatch.setattr(litellm, "_async_failure_callback", []) + + for _ in range(11): + await litellm.acompletion( + model="gpt-5-mini", + messages=[{"role": "user", "content": "callback registration"}], + mock_response="ok", + ) + + assert litellm.callbacks == [logger] + assert litellm.input_callback == [logger] + assert litellm.success_callback == [logger] + assert litellm.failure_callback == [logger] + assert litellm._async_input_callback == [] + assert litellm._async_success_callback == [logger] + assert litellm._async_failure_callback == [logger] + + +def test_get_provider_audio_transcription_config_resolves_for_every_provider() -> None: + configs: Final = { + provider: ProviderConfigManager.get_provider_audio_transcription_config(model="whisper-1", provider=provider) + for provider in LlmProviders + } + unexpected: Final = { + provider: config + for provider, config in configs.items() + if config is not None and not isinstance(config, BaseAudioTranscriptionConfig) + } + + assert unexpected == {} + assert isinstance(configs[LlmProviders.OPENAI], litellm.OpenAIWhisperAudioTranscriptionConfig) + + def test_get_valid_models_from_provider_cache_invalidation(monkeypatch): """ Test that get_valid_models returns the correct models for a given provider diff --git a/tests/unit/test_version.py b/tests/unit/test_version.py new file mode 100644 index 00000000000..30b1f259f1e --- /dev/null +++ b/tests/unit/test_version.py @@ -0,0 +1,40 @@ +import runpy +from pathlib import Path +from typing import Final +from unittest.mock import patch + +import importlib_metadata +import pytest + + +@pytest.mark.parametrize( + ("installed", "expected"), + [ + ({"litellm": "1.2.3"}, "1.2.3"), + ({"litellm-core": "2.3.4"}, "2.3.4"), + ({}, "unknown"), + ], +) +def test_version_uses_installed_distribution(installed: dict[str, str], expected: str) -> None: + def lookup(name: str) -> str: + if name not in installed: + raise importlib_metadata.PackageNotFoundError(name) + return installed[name] + + with patch("importlib_metadata.version", side_effect=lookup): + result: Final = runpy.run_path(str(Path(__file__).resolve().parents[2] / "litellm/_version.py")) + assert result["version"] == expected + + +@pytest.mark.parametrize("core_version", ["1.2.3", "2.3.4"]) +def test_version_rejects_overlapping_distributions(core_version: str) -> None: + installed: Final = {"litellm": "1.2.3", "litellm-core": core_version} + with patch("importlib_metadata.version", side_effect=installed.__getitem__): + with pytest.raises(RuntimeError, match=r"litellm and litellm-core.*separate environments"): + runpy.run_path(str(Path(__file__).resolve().parents[2] / "litellm/_version.py")) + + +def test_version_handles_unreadable_metadata() -> None: + with patch("importlib_metadata.version", side_effect=ValueError("Invalid metadata")): + result: Final = runpy.run_path(str(Path(__file__).resolve().parents[2] / "litellm/_version.py")) + assert result["version"] == "unknown" diff --git a/tests/unit/types/llms/test_types_llms_openai.py b/tests/unit/types/llms/test_types_llms_openai.py index e59643f6509..871e4cc8cfb 100644 --- a/tests/unit/types/llms/test_types_llms_openai.py +++ b/tests/unit/types/llms/test_types_llms_openai.py @@ -552,6 +552,17 @@ class TestOpenAIFileObjectBatchGuardrailSerialization: original = self._file_object(litellm_batch_guardrail=self._report()) assert OpenAIFileObject(**original.model_dump()) == original + def test_details_fallback_marker_is_omitted_when_unset_and_round_trips_when_set(self): + from litellm.types.llms.openai import OpenAIFileObject + + without_marker = self._file_object() + assert "litellm_details_fallback" not in without_marker.model_dump() + assert "litellm_details_fallback" not in without_marker.model_dump_json() + + with_marker = self._file_object(litellm_details_fallback=True) + assert with_marker.model_dump()["litellm_details_fallback"] is True + assert OpenAIFileObject.model_validate_json(with_marker.model_dump_json()) == with_marker + def test_serialization_json_schema_still_describes_the_model(self): """A return annotation on the wrap serializer would collapse this to a bare object.""" from litellm.types.llms.openai import OpenAIFileObject diff --git a/tests/unit/types/test_guardrails_case_normalization.py b/tests/unit/types/test_guardrails_case_normalization.py index 26c1d395320..8c8192be3be 100644 --- a/tests/unit/types/test_guardrails_case_normalization.py +++ b/tests/unit/types/test_guardrails_case_normalization.py @@ -2,12 +2,19 @@ Test case normalization in LitellmParams for all guardrail types """ -from typing import Literal +import logging +from typing import Final, Literal import pytest from pydantic import ValidationError -from litellm.types.guardrails import BaseLitellmParams, LitellmParams +from litellm.types.guardrails import ( + BaseLitellmParams, + LitellmParams, + runtime_stream_scope, + stored_stream_scope, + with_tolerated_stream_scope, +) class TestLitellmParamsCaseNormalization: @@ -184,3 +191,80 @@ class TestSensitiveDataRoutingValidation: on_sensitive_data="BLOCK", ) assert params.on_sensitive_data == "block" + + +class TestStreamScopeValidation: + def test_scalar_is_case_normalized(self): + params = LitellmParams(guardrail="bedrock", mode="post_call", stream_scope="Streaming") + assert params.stream_scope == "streaming" + + def test_map_keys_and_values_are_normalized(self): + params = LitellmParams( + guardrail="bedrock", + mode=["pre_call", "post_call"], + stream_scope={"Pre_Call": "Both", "POST_CALL": "Non_Streaming"}, + ) + assert params.stream_scope == {"pre_call": "both", "post_call": "non_streaming"} + + def test_invalid_scalar_is_rejected(self): + with pytest.raises(ValidationError, match="stream_scope must be one of"): + LitellmParams(guardrail="bedrock", mode="post_call", stream_scope="chunks") + + def test_invalid_map_key_is_rejected(self): + with pytest.raises(ValidationError, match="stream_scope keys must be guardrail modes"): + LitellmParams(guardrail="bedrock", mode="post_call", stream_scope={"not_a_mode": "both"}) + + def test_invalid_map_value_is_rejected(self): + with pytest.raises(ValidationError, match="stream_scope must be one of"): + LitellmParams(guardrail="bedrock", mode="post_call", stream_scope={"post_call": "sometimes"}) + + def test_runtime_stream_scope_normalizes_direct_constructor_maps(self): + default, by_hook = runtime_stream_scope({"Pre_Call": "streaming"}) + assert default == "both" + assert dict(by_hook) == {"pre_call": "streaming"} + + def test_runtime_stream_scope_rejects_invalid_direct_input(self): + with pytest.raises(ValueError, match="stream_scope must be one of"): + runtime_stream_scope("chunks") + + @pytest.mark.parametrize( + "value, expected", + [ + ("streaming", "streaming"), + ("non_streaming", "non_streaming"), + ("both", "both"), + ({"pre_call": "streaming"}, {"pre_call": "streaming"}), + ("sometimes", None), + ({"pre_call": "sometimes"}, None), + ], + ) + def test_stored_stream_scope_tolerates_invalid_values( + self, + value: object, + expected: object, + caplog: pytest.LogCaptureFixture, + ) -> None: + caplog.set_level(logging.WARNING) + + result: Final = stored_stream_scope(value) + + assert result == expected + if expected is None: + assert f"Ignoring invalid stored stream_scope value of type {type(value).__name__}" in caplog.text + assert "sometimes" not in caplog.text + + def test_tolerated_stream_scope_rewrites_only_the_scope_field(self) -> None: + params: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "stream_scope": "sometimes", + } + + tolerated: Final = with_tolerated_stream_scope(params) + + assert tolerated == { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "stream_scope": None, + } + assert params["stream_scope"] == "sometimes" diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index 8d163731a51..f8e2befe237 100644 --- a/tests/unit/types/test_litellm_params.py +++ b/tests/unit/types/test_litellm_params.py @@ -85,6 +85,7 @@ CONNECTION_NAMES: Final = ( "litellm_credential_name", "configurable_clientside_auth_params", "use_xai_oauth", + "fireworks_forward_user_id", "aws_batch_role_arn", "s3_bucket_name", "s3_region_name", diff --git a/tests/unit/types/test_openai_decisions.py b/tests/unit/types/test_openai_decisions.py deleted file mode 100644 index d9387d644d5..00000000000 --- a/tests/unit/types/test_openai_decisions.py +++ /dev/null @@ -1,150 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping -from typing import Final - -import pytest -from pydantic import TypeAdapter, ValidationError - -from litellm.types.openai_decisions import ( - ChoiceAnswer, - DecisionsRequestBody, - DecisionsResponse, - RefusalAnswer, - ScoreQuestion, -) - -_REQUEST: Final[Mapping[str, object]] = { - "input": [ - { - "role": "user", - "content": [ - {"type": "input_text", "text": "Is this receipt a valid business expense?"}, - {"type": "input_image", "image_url": "https://example.com/receipt.png", "detail": "high"}, - ], - } - ], - "questions": [ - {"type": "predicate", "name": "is_expense", "instructions": "Is this a business expense?"}, - { - "type": "choice", - "name": "approve", - "instructions": "Should this be approved?", - "choices": [{"value": True, "description": "approve"}, {"value": False, "description": "reject"}], - }, - { - "type": "score", - "name": "risk", - "instructions": "How risky is this expense?", - "levels": [{"label": "low"}, {"label": "high", "description": "needs a manager"}], - }, - ], - "safety_identifier": "user-123", -} -_RESPONSE: Final[Mapping[str, object]] = { - "model": "gpt-6-luna", - "answers": [ - {"type": "predicate", "name": "is_expense", "probability": 0.92}, - { - "type": "choice", - "name": "approve", - "choice": True, - "probabilities": [{"value": True, "probability": 0.7}, {"value": False, "probability": 0.3}], - "confidence": 0.7, - }, - {"type": "refusal", "name": "risk"}, - ], - "usage": { - "input_tokens": 120, - "input_tokens_details": {"cached_tokens": 100, "cache_write_tokens": 0}, - "output_tokens": 12, - "output_tokens_details": {"reasoning_tokens": 4}, - "total_tokens": 132, - }, -} -_REQUEST_ADAPTER: Final[TypeAdapter[DecisionsRequestBody]] = TypeAdapter(DecisionsRequestBody) -_RESPONSE_ADAPTER: Final[TypeAdapter[DecisionsResponse]] = TypeAdapter(DecisionsResponse) - - -def test_the_documented_request_round_trips_with_its_boolean_choices_and_image_part() -> None: - request: Final = _REQUEST_ADAPTER.validate_python(_REQUEST) - - assert request.model_dump(mode="json", exclude_none=True) == _REQUEST - assert isinstance(request.questions[2], ScoreQuestion) - - -def test_the_documented_response_keeps_answer_order_refusals_and_token_details() -> None: - response: Final = _RESPONSE_ADAPTER.validate_python(_RESPONSE) - - assert response.model_dump(mode="json") == _RESPONSE - assert isinstance(response.answers[1], ChoiceAnswer) - assert isinstance(response.answers[2], RefusalAnswer) - - -_OFF_SPEC: Final[tuple[tuple[str, object], ...]] = ( - ("questions", [{"type": "noul", "name": "q", "instructions": "x"}]), - ("questions", [{"type": "choice", "name": "q", "instructions": "x", "choices": [{"value": 1}]}]), - ("questions", [{"type": "score", "name": "q", "levels": [{"label": "low"}]}]), - ("input", {"state": "not an OpenAI input"}), -) - - -@pytest.mark.parametrize(("field", "value"), _OFF_SPEC) -def test_requests_off_the_spec_are_rejected(field: str, value: object) -> None: - with pytest.raises(ValidationError): - _REQUEST_ADAPTER.validate_python({**_REQUEST, field: value}) - - -_EMPTY_COLLECTIONS: Final[tuple[tuple[str, list[object]], ...]] = ( - ("questions", []), - ("questions", [{"type": "choice", "name": "q", "instructions": "x", "choices": []}]), - ("questions", [{"type": "score", "name": "q", "instructions": "x", "levels": []}]), -) - - -@pytest.mark.parametrize(("field", "value"), _EMPTY_COLLECTIONS) -def test_empty_collections_are_left_for_the_provider_to_judge(field: str, value: list[object]) -> None: - request: Final = _REQUEST_ADAPTER.validate_python({**_REQUEST, field: value}) - - assert request.model_dump(mode="json", exclude_none=True)[field] == value - - -def test_choice_values_keep_their_type_so_a_string_true_and_a_boolean_true_stay_distinct() -> None: - answer: Final = { - "type": "choice", - "name": "approve", - "choice": "true", - "probabilities": [{"value": "true", "probability": 0.6}, {"value": True, "probability": 0.4}], - "confidence": 0.6, - } - - response: Final = _RESPONSE_ADAPTER.validate_python({**_RESPONSE, "answers": [answer]}) - - assert response.model_dump(mode="json")["answers"] == [answer] - with pytest.raises(ValidationError): - _RESPONSE_ADAPTER.validate_python({**_RESPONSE, "answers": [{**answer, "choice": 1}]}) - - -_OFF_SPEC_RESPONSE: Final[tuple[str, ...]] = ("model", "usage") - - -@pytest.mark.parametrize("field", _OFF_SPEC_RESPONSE) -def test_responses_missing_a_required_field_are_rejected(field: str) -> None: - with pytest.raises(ValidationError): - _RESPONSE_ADAPTER.validate_python({k: v for k, v in _RESPONSE.items() if k != field}) - - -def test_usage_without_token_details_is_rejected() -> None: - usage: Final = {"input_tokens": 120, "output_tokens": 12, "total_tokens": 132} - - with pytest.raises(ValidationError): - _RESPONSE_ADAPTER.validate_python({**_RESPONSE, "usage": usage}) - - -def test_hidden_params_live_outside_the_wire_body() -> None: - response: Final = _RESPONSE_ADAPTER.validate_python(_RESPONSE) - - response.set_hidden_params({"custom_llm_provider": "openai"}) - - assert response.hidden_params == {"custom_llm_provider": "openai"} - assert "_hidden_params" not in response.model_dump(mode="json") diff --git a/type-discipline-budget.json b/type-discipline-budget.json deleted file mode 100644 index 11bc9722d81..00000000000 --- a/type-discipline-budget.json +++ /dev/null @@ -1,41 +0,0 @@ -{ - "LIT001": { - "limit": 22174 - }, - "LIT003": { - "limit": 261 - }, - "LIT004": { - "limit": 38 - }, - "LIT005": { - "limit": 0 - }, - "LIT006": { - "limit": 1035 - }, - "LIT007": { - "limit": 0 - }, - "LIT008": { - "limit": 945 - }, - "LIT009": { - "limit": 0 - }, - "LIT010": { - "limit": 16398 - }, - "LIT011": { - "limit": 5504 - }, - "LIT012": { - "limit": 4486 - }, - "LIT013": { - "limit": 0 - }, - "LIT014": { - "limit": 369 - } -} diff --git a/ui/litellm-dashboard/public/assets/logos/typesafe.png b/ui/litellm-dashboard/public/assets/logos/typesafe.png new file mode 100644 index 00000000000..45107c7d09d Binary files /dev/null and b/ui/litellm-dashboard/public/assets/logos/typesafe.png differ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailReadOnlyDetails.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailReadOnlyDetails.tsx new file mode 100644 index 00000000000..5fe88b4870c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailReadOnlyDetails.tsx @@ -0,0 +1,68 @@ +import { Badge } from "@/components/ui/badge"; +import { GuardrailModeRows } from "./GuardrailModeDisplay"; +import { GuardrailStreamScopeDetail } from "./StreamScopeFields"; +import ToolPermissionRulesEditor, { type ToolPermissionConfig } from "./tool_permission/ToolPermissionRulesEditor"; + +export const GuardrailReadOnlyDetails = ({ + guardrailId, + guardrailName, + displayName, + litellmParams, + streamScope, + defaultOn, + piiEntityCount, + createdAt, + updatedAt, + showToolPermission, + toolPermissionConfig, +}: { + guardrailId: string; + guardrailName: string; + displayName: string; + litellmParams: { mode?: unknown; logging_only_scope?: string | null }; + streamScope: unknown; + defaultOn: boolean | undefined; + piiEntityCount: number; + createdAt: string; + updatedAt: string; + showToolPermission: boolean; + toolPermissionConfig: ToolPermissionConfig; +}) => ( +
+
+

Guardrail ID

+
{guardrailId}
+
+
+

Guardrail Name

+
{guardrailName || "Unnamed Guardrail"}
+
+
+

Provider

+
{displayName}
+
+ + +
+

Default On

+ {defaultOn ? "Yes" : "No"} +
+ {piiEntityCount > 0 && ( +
+

PII Protection

+
+ {piiEntityCount} PII entities configured +
+
+ )} +
+

Created At

+
{createdAt}
+
+
+

Last Updated

+
{updatedAt}
+
+ {showToolPermission && } +
+); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/StreamScopeFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/StreamScopeFields.tsx new file mode 100644 index 00000000000..1a8f55ce477 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/StreamScopeFields.tsx @@ -0,0 +1,95 @@ +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { GuardrailField, labelWithHint, type GuardrailFormControl } from "./GuardrailFormField"; +import { + STREAM_SCOPE_OPTIONS, + formatGuardrailStreamScope, + type GuardrailStreamScope, + isGuardrailStreamScope, +} from "./guardrail_info_helpers"; + +const STREAM_SCOPE_ITEMS = STREAM_SCOPE_OPTIONS.map((option) => ({ + label: option.label, + value: option.value, +})); + +const REQUEST_SHAPE_HINT = + "Run this guardrail on streaming requests, non-streaming requests, or both, for each selected mode."; + +export const StreamScopeFields = ({ + modes, + value, + onChange, +}: { + modes: string[]; + value: Record; + onChange: (next: Record) => void; +}) => { + if (modes.length === 0) return null; + + return ( +
+ {modes.map((mode) => { + const selected = value[mode] ?? "both"; + return ( +
+ + +
+ ); + })} +
+ ); +}; + +export const StreamScopeFormField = ({ control, modes }: { control: GuardrailFormControl; modes: string[] }) => ( + + {({ value, onChange }) => ( + | undefined) ?? {}} + onChange={onChange} + /> + )} + +); + +export const GuardrailStreamScopeCaption = ({ raw }: { raw: unknown }) => { + const label = formatGuardrailStreamScope(raw); + if (!label) return null; + return

{label}

; +}; + +export const GuardrailStreamScopeDetail = ({ raw }: { raw: unknown }) => { + const label = formatGuardrailStreamScope(raw); + if (!label) return null; + return ( +
+

Request shape

+
{label}
+
+ ); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx index 230be13ef33..9628c284bd2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx @@ -89,6 +89,23 @@ describe("AddGuardrailForm create payload characterization", () => { }); }); + it("sends stream_scope when a mode is restricted to streaming requests", async () => { + const user = userEvent.setup({ delay: null }); + renderForm(); + + await user.type(await screen.findByLabelText("Guardrail Name"), "my-bedrock"); + await pickProvider(user, "Bedrock Guardrail"); + await chooseSelectOption(user, screen.getByLabelText("pre_call applies to"), "Streaming only"); + await user.type(await screen.findByPlaceholderText("The guardrail id on Bedrock"), "gr-123"); + await user.click(screen.getByRole("button", { name: "Next" })); + await user.click(await screen.findByRole("button", { name: "Create Guardrail" })); + + await waitFor(() => expect(networking.createGuardrailCall).toHaveBeenCalledTimes(1)); + expect(payload()).toMatchObject({ + litellm_params: { stream_scope: "streaming" }, + }); + }); + it("switches mode from the seeded string to an array once the user touches the multi select", async () => { const user = userEvent.setup({ delay: null }); renderForm(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx index 43a542289f6..098245fa34b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx @@ -23,11 +23,14 @@ import { shouldRenderContentFilterConfigSettings, shouldRenderLLMJudgeFields, shouldRenderPIIConfigSettings, - supportsDirectionalLoggingOnlyScope, + streamScopePayload, toModeArray, + type GuardrailStreamScope, + supportsDirectionalLoggingOnlyScope, type LoggingOnlyScope, type LoggingOnlyScopeChoice, } from "./guardrail_info_helpers"; +import { StreamScopeFormField } from "./StreamScopeFields"; import { Logo } from "@/components/molecules/logo/Logo"; import { MultiSelect } from "@/components/shared/MultiSelect"; import { FieldGroup } from "@/components/ui/field"; @@ -172,6 +175,7 @@ const INITIAL_VALUES: GuardrailFormValues = { logging_only_scope_choice: "default", skip_system_message_choice: "inherit", skip_tool_message_choice: "inherit", + stream_scope_by_mode: {}, }; const ALWAYS_ON_ITEMS = [ @@ -466,6 +470,14 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a guardrail_info: {}, }; + const streamScope = streamScopePayload( + toModeArray(values.mode), + (values.stream_scope_by_mode as Record | undefined) ?? {}, + ); + if (streamScope !== undefined) { + guardrailData.litellm_params.stream_scope = streamScope; + } + const skipForCreate = choiceToSkipSystemForCreate(asSkipChoice(values.skip_system_message_choice)); if (skipForCreate !== undefined) { guardrailData.litellm_params.skip_system_message_in_guardrail = skipForCreate; @@ -773,6 +785,8 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a )} + + { expect(screen.getByRole("button", { name: /update guardrail/i })).toBeInTheDocument(); }); + it("should open without crashing for a tag-scoped mode and show it read-only", async () => { + renderModal({ + editData: { + guardrail_id: "g-tag", + guardrail_name: "tag-mode-guardrail", + litellm_params: { + mode: { tags: { "team-a": "pre_call" }, default: "post_call" }, + default_on: true, + custom_code: "def apply_guardrail(): pass", + }, + }, + }); + + expect(await screen.findByText("Edit Custom Guardrail")).toBeInTheDocument(); + const modeInput = screen.getByLabelText("Mode (tag-scoped)"); + expect(modeInput).toBeDisabled(); + expect(modeInput).toHaveValue("post_call, pre_call (tag-based)"); + expect(screen.getByText("Mode (tag-scoped, read-only)")).toBeInTheDocument(); + expect(screen.queryByText(/applies to/)).not.toBeInTheDocument(); + }); + + it("should omit mode and stream_scope from the update payload for a tag-scoped guardrail", async () => { + const user = userEvent.setup(); + renderModal({ + editData: { + guardrail_id: "g-tag", + guardrail_name: "tag-mode-guardrail", + litellm_params: { + mode: { tags: { "team-a": "pre_call" }, default: "post_call" }, + default_on: true, + custom_code: "def apply_guardrail(): pass", + }, + }, + }); + + const nameInput = await screen.findByDisplayValue("tag-mode-guardrail"); + await user.clear(nameInput); + await user.type(nameInput, "renamed-guardrail"); + await user.click(screen.getByRole("button", { name: /update guardrail/i })); + + await waitFor(() => { + expect(mockUpdate).toHaveBeenCalledTimes(1); + }); + const [token, guardrailId, payload] = mockUpdate.mock.calls[0] as [ + string, + string, + Record>, + ]; + expect(token).toBe("test-token"); + expect(guardrailId).toBe("g-tag"); + expect(payload.guardrail_name).toBe("renamed-guardrail"); + expect(payload.litellm_params.custom_code).toBe("def apply_guardrail(): pass"); + expect(payload.litellm_params).not.toHaveProperty("mode"); + expect(payload.litellm_params).not.toHaveProperty("stream_scope"); + }); + it("should keep save disabled until a guardrail name is entered", async () => { const user = userEvent.setup(); renderModal(); @@ -111,8 +167,7 @@ describe("CustomCodeModal", () => { expect(await screen.findByDisplayValue(/async def apply_guardrail/)).toBeInTheDocument(); - const comboboxes = screen.getAllByRole("combobox"); - await user.click(comboboxes[comboboxes.length - 1]); + await user.click(screen.getByRole("combobox", { name: "Template" })); const options = await screen.findAllByText("Block SSN"); await user.click(options[options.length - 1]); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx index 9deef4aeee3..c89540596bd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx @@ -37,152 +37,30 @@ import { import { Switch } from "@/components/ui/switch"; import { Textarea } from "@/components/ui/textarea"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; - -// Code templates -const CODE_TEMPLATES = { - empty: { - name: "Empty Template", - code: `async def apply_guardrail(inputs, request_data, input_type): - # inputs: {texts, images, tools, tool_calls, structured_messages, model} - # request_data: {model, user_id, team_id, end_user_id, metadata} - # input_type: "request" or "response" - return allow()`, - }, - blockSSN: { - name: "Block SSN", - code: `def apply_guardrail(inputs, request_data, input_type): - for text in inputs["texts"]: - if regex_match(text, r"\\d{3}-\\d{2}-\\d{4}"): - return block("SSN detected") - return allow()`, - }, - redactEmail: { - name: "Redact Emails", - code: `def apply_guardrail(inputs, request_data, input_type): - pattern = r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\\.[a-zA-Z]{2,}" - modified = [] - for text in inputs["texts"]: - modified.append(regex_replace(text, pattern, "[EMAIL REDACTED]")) - return modify(texts=modified)`, - }, - blockSQL: { - name: "Block SQL Injection", - code: `def apply_guardrail(inputs, request_data, input_type): - if input_type != "request": - return allow() - for text in inputs["texts"]: - if contains_code_language(text, ["sql"]): - return block("SQL code not allowed") - return allow()`, - }, - validateJSON: { - name: "Validate JSON", - code: `def apply_guardrail(inputs, request_data, input_type): - if input_type != "response": - return allow() - - schema = {"type": "object", "required": ["name", "value"]} - - for text in inputs["texts"]: - obj = json_parse(text) - if obj is None: - return block("Invalid JSON response") - if not json_schema_valid(obj, schema): - return block("Response missing required fields") - return allow()`, - }, - externalAPI: { - name: "External API Check (async)", - code: `async def apply_guardrail(inputs, request_data, input_type): - # Call an external moderation API (async for non-blocking) - for text in inputs["texts"]: - response = await http_post( - "https://api.example.com/moderate", - body={"text": text, "user_id": request_data["user_id"]}, - headers={"Authorization": "Bearer YOUR_API_KEY"}, - timeout=10 - ) - - if not response["success"]: - # API call failed, allow by default or block - return allow() - - if response["body"].get("flagged"): - return block(response["body"].get("reason", "Content flagged")) - - return allow()`, - }, -}; - -// Available primitives organized by category -const PRIMITIVES = { - "Return Values": [ - { name: "allow()", desc: "Let request/response through" }, - { name: "block(reason)", desc: "Reject with message" }, - { name: "flag(reason, metadata={})", desc: "Let through, record a non-blocking violation" }, - { name: "modify(texts=[], images=[], tool_calls=[])", desc: "Transform content" }, - ], - "HTTP Requests (async)": [ - { name: "await http_request(url, method, headers, body)", desc: "Make async HTTP request" }, - { name: "await http_get(url, headers)", desc: "Async GET request" }, - { name: "await http_post(url, body, headers)", desc: "Async POST request" }, - ], - "Regex Functions": [ - { name: "regex_match(text, pattern)", desc: "Returns True if pattern found" }, - { name: "regex_replace(text, pattern, replacement)", desc: "Replace all matches" }, - { name: "regex_find_all(text, pattern)", desc: "Return list of matches" }, - ], - "JSON Functions": [ - { name: "json_parse(text)", desc: "Parse JSON string, returns None on error" }, - { name: "json_stringify(obj)", desc: "Convert to JSON string" }, - { name: "json_schema_valid(obj, schema)", desc: "Validate against JSON schema" }, - ], - "URL Functions": [ - { name: "extract_urls(text)", desc: "Extract all URLs from text" }, - { name: "is_valid_url(url)", desc: "Check if URL is valid" }, - { name: "all_urls_valid(text)", desc: "Check all URLs in text are valid" }, - ], - "Code Detection": [ - { name: "detect_code(text)", desc: "Returns True if code detected" }, - { name: "detect_code_languages(text)", desc: "Returns list of detected languages" }, - { name: 'contains_code_language(text, ["sql"])', desc: "Check for specific languages" }, - ], - "Text Utilities": [ - { name: "contains(text, substring)", desc: "Check if substring exists" }, - { name: "contains_any(text, [substr1, substr2])", desc: "Check if any substring exists" }, - { name: "word_count(text)", desc: "Count words" }, - { name: "char_count(text)", desc: "Count characters" }, - { name: "lower(text) / upper(text) / trim(text)", desc: "String transforms" }, - ], -}; - -const MODE_OPTIONS = [ - { value: "pre_call", label: "pre_call (Request)" }, - { value: "post_call", label: "post_call (Response)" }, - { value: "during_call", label: "during_call (Parallel)" }, - { value: "logging_only", label: "logging_only" }, - { value: "pre_mcp_call", label: "pre_mcp_call (Before MCP Tool Call)" }, - { value: "post_mcp_call", label: "post_mcp_call (After MCP Tool Call)" }, - { value: "during_mcp_call", label: "during_mcp_call (During MCP Tool Call)" }, -]; - -const TEMPLATE_ITEMS = Object.entries(CODE_TEMPLATES).map(([key, template]) => ({ - value: key, - label: template.name, -})); - -type ModeOption = (typeof MODE_OPTIONS)[number]; - -const MODE_OPTION_BY_VALUE: Record = Object.fromEntries( - MODE_OPTIONS.map((option) => [option.value, option]), -); +import { StreamScopeFields } from "../StreamScopeFields"; +import { + formatGuardrailMode, + streamScopeByModeFromConfig, + streamScopeForUpdate, + streamScopePayload, + type GuardrailStreamScope, +} from "../guardrail_info_helpers"; +import { + CODE_TEMPLATES, + MODE_OPTION_BY_VALUE, + MODE_OPTIONS, + PRIMITIVES, + TEMPLATE_ITEMS, + type ModeOption, +} from "./custom_code_catalog"; // Data for editing an existing guardrail + export interface EditGuardrailData { guardrail_id: string; guardrail_name: string; litellm_params: { - mode?: string | string[]; + mode?: string | string[] | Record; default_on?: boolean; custom_code?: string; logging_only_scope?: LoggingOnlyScope | null; @@ -204,6 +82,7 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS const isEditMode = !!editData; const [guardrailName, setGuardrailName] = useState(""); const [mode, setMode] = useState(["pre_call"]); + const [streamScopeByMode, setStreamScopeByMode] = useState>({}); const [loggingOnlyScopeChoice, setLoggingOnlyScopeChoice] = useState("default"); const [defaultOn, setDefaultOn] = useState(false); const [selectedTemplate, setSelectedTemplate] = useState("empty"); @@ -315,11 +194,14 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS setCode(CODE_TEMPLATES[templateKey as keyof typeof CODE_TEMPLATES].code); }; - // Normalize mode from API (string or string[]) to string[] - const normalizeMode = (m: string | string[] | undefined): string[] => { + // Normalize mode from API (string or string[]) to string[]. + // A tag-scoped mode dict ({ tags, default }) is managed outside this editor, so it + // contributes no editable modes and is displayed read-only instead. + const normalizeMode = (m: string | string[] | Record | undefined): string[] => { if (m === undefined || m === null) return ["pre_call"]; if (Array.isArray(m)) return m.length ? m : ["pre_call"]; - return [m]; + if (typeof m === "string") return [m]; + return []; }; // Reset form when modal opens or editData changes @@ -329,6 +211,12 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS // Edit mode: populate with existing data setGuardrailName(editData.guardrail_name || ""); setMode(normalizeMode(editData.litellm_params?.mode)); + setStreamScopeByMode( + streamScopeByModeFromConfig( + editData.litellm_params?.stream_scope, + normalizeMode(editData.litellm_params?.mode), + ), + ); setLoggingOnlyScopeChoice(loggingOnlyScopeToChoice(editData.litellm_params?.logging_only_scope)); setDefaultOn(editData.litellm_params?.default_on || false); setCode(editData.litellm_params?.custom_code || CODE_TEMPLATES.empty.code); @@ -337,6 +225,7 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS // Create mode: reset to defaults setGuardrailName(""); setMode(["pre_call"]); + setStreamScopeByMode({}); setLoggingOnlyScopeChoice("default"); setDefaultOn(false); setSelectedTemplate("empty"); @@ -411,11 +300,21 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS if (defaultOn !== editData.litellm_params?.default_on) { updateData.litellm_params.default_on = defaultOn; } + const nextStreamScope = streamScopeForUpdate( + mode, + streamScopeByMode, + editData.litellm_params?.stream_scope, + existingMode, + ); + if (nextStreamScope !== undefined) { + updateData.litellm_params.stream_scope = nextStreamScope; + } await updateGuardrailCall(accessToken, editData.guardrail_id, updateData); toast.success("Custom code guardrail updated successfully"); } else { // Create new guardrail + const streamScope = streamScopePayload(mode, streamScopeByMode); const guardrailData = { guardrail_name: guardrailName, litellm_params: { @@ -423,6 +322,7 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS mode: mode, default_on: defaultOn, custom_code: code, + ...(streamScope !== undefined ? { stream_scope: streamScope } : {}), ...getCustomCodeLoggingOnlyScopeCreate(mode, loggingOnlyScopeChoice), }, guardrail_info: {}, @@ -511,6 +411,11 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS const lineCount = code.split("\n").length; const selectedModeOptions = mode.map((value) => MODE_OPTION_BY_VALUE[value]).filter(Boolean); + const rawEditMode = editData?.litellm_params?.mode; + const tagScopedModeLabel = + rawEditMode !== null && typeof rawEditMode === "object" && !Array.isArray(rawEditMode) + ? formatGuardrailMode(rawEditMode) || "-" + : null; return ( !open && onClose()}> @@ -533,32 +438,38 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS />
- - setMode(options.map((option) => option.value))} - multiple - > - } className="w-full"> - {selectedModeOptions.map((option) => ( - - {option.label} - - ))} - - - - No matching modes - - {(option: ModeOption) => ( - + + {tagScopedModeLabel ? ( + + ) : ( + setMode(options.map((option) => option.value))} + multiple + > + } className="w-full"> + {selectedModeOptions.map((option) => ( + {option.label} - - )} - - - + + ))} + + + + No matching modes + + {(option: ModeOption) => ( + + {option.label} + + )} + + + + )}
{mode.includes("logging_only") && ( @@ -600,6 +511,11 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS + {mode.length > 0 && ( +
+ +
+ )} {/* Main Content */}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/custom_code_catalog.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/custom_code_catalog.ts new file mode 100644 index 00000000000..aad3df181a7 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/custom_code_catalog.ts @@ -0,0 +1,136 @@ +export const CODE_TEMPLATES = { + empty: { + name: "Empty Template", + code: `async def apply_guardrail(inputs, request_data, input_type): + # inputs: {texts, images, tools, tool_calls, structured_messages, model} + # request_data: {model, user_id, team_id, end_user_id, metadata} + # input_type: "request" or "response" + return allow()`, + }, + blockSSN: { + name: "Block SSN", + code: `def apply_guardrail(inputs, request_data, input_type): + for text in inputs["texts"]: + if regex_match(text, r"\\d{3}-\\d{2}-\\d{4}"): + return block("SSN detected") + return allow()`, + }, + redactEmail: { + name: "Redact Emails", + code: `def apply_guardrail(inputs, request_data, input_type): + pattern = r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\\.[a-zA-Z]{2,}" + modified = [] + for text in inputs["texts"]: + modified.append(regex_replace(text, pattern, "[EMAIL REDACTED]")) + return modify(texts=modified)`, + }, + blockSQL: { + name: "Block SQL Injection", + code: `def apply_guardrail(inputs, request_data, input_type): + if input_type != "request": + return allow() + for text in inputs["texts"]: + if contains_code_language(text, ["sql"]): + return block("SQL code not allowed") + return allow()`, + }, + validateJSON: { + name: "Validate JSON", + code: `def apply_guardrail(inputs, request_data, input_type): + if input_type != "response": + return allow() + + schema = {"type": "object", "required": ["name", "value"]} + + for text in inputs["texts"]: + obj = json_parse(text) + if obj is None: + return block("Invalid JSON response") + if not json_schema_valid(obj, schema): + return block("Response missing required fields") + return allow()`, + }, + externalAPI: { + name: "External API Check (async)", + code: `async def apply_guardrail(inputs, request_data, input_type): + # Call an external moderation API (async for non-blocking) + for text in inputs["texts"]: + response = await http_post( + "https://api.example.com/moderate", + body={"text": text, "user_id": request_data["user_id"]}, + headers={"Authorization": "Bearer YOUR_API_KEY"}, + timeout=10 + ) + + if not response["success"]: + # API call failed, allow by default or block + return allow() + + if response["body"].get("flagged"): + return block(response["body"].get("reason", "Content flagged")) + + return allow()`, + }, +}; + +export const PRIMITIVES = { + "Return Values": [ + { name: "allow()", desc: "Let request/response through" }, + { name: "block(reason)", desc: "Reject with message" }, + { name: "flag(reason, metadata={})", desc: "Let through, record a non-blocking violation" }, + { name: "modify(texts=[], images=[], tool_calls=[])", desc: "Transform content" }, + ], + "HTTP Requests (async)": [ + { name: "await http_request(url, method, headers, body)", desc: "Make async HTTP request" }, + { name: "await http_get(url, headers)", desc: "Async GET request" }, + { name: "await http_post(url, body, headers)", desc: "Async POST request" }, + ], + "Regex Functions": [ + { name: "regex_match(text, pattern)", desc: "Returns True if pattern found" }, + { name: "regex_replace(text, pattern, replacement)", desc: "Replace all matches" }, + { name: "regex_find_all(text, pattern)", desc: "Return list of matches" }, + ], + "JSON Functions": [ + { name: "json_parse(text)", desc: "Parse JSON string, returns None on error" }, + { name: "json_stringify(obj)", desc: "Convert to JSON string" }, + { name: "json_schema_valid(obj, schema)", desc: "Validate against JSON schema" }, + ], + "URL Functions": [ + { name: "extract_urls(text)", desc: "Extract all URLs from text" }, + { name: "is_valid_url(url)", desc: "Check if URL is valid" }, + { name: "all_urls_valid(text)", desc: "Check all URLs in text are valid" }, + ], + "Code Detection": [ + { name: "detect_code(text)", desc: "Returns True if code detected" }, + { name: "detect_code_languages(text)", desc: "Returns list of detected languages" }, + { name: 'contains_code_language(text, ["sql"])', desc: "Check for specific languages" }, + ], + "Text Utilities": [ + { name: "contains(text, substring)", desc: "Check if substring exists" }, + { name: "contains_any(text, [substr1, substr2])", desc: "Check if any substring exists" }, + { name: "word_count(text)", desc: "Count words" }, + { name: "char_count(text)", desc: "Count characters" }, + { name: "lower(text) / upper(text) / trim(text)", desc: "String transforms" }, + ], +}; + +export const MODE_OPTIONS = [ + { value: "pre_call", label: "pre_call (Request)" }, + { value: "post_call", label: "post_call (Response)" }, + { value: "during_call", label: "during_call (Parallel)" }, + { value: "logging_only", label: "logging_only" }, + { value: "pre_mcp_call", label: "pre_mcp_call (Before MCP Tool Call)" }, + { value: "post_mcp_call", label: "post_mcp_call (After MCP Tool Call)" }, + { value: "during_mcp_call", label: "during_mcp_call (During MCP Tool Call)" }, +]; + +export const TEMPLATE_ITEMS = Object.entries(CODE_TEMPLATES).map(([key, template]) => ({ + value: key, + label: template.name, +})); + +export type ModeOption = (typeof MODE_OPTIONS)[number]; + +export const MODE_OPTION_BY_VALUE: Record = Object.fromEntries( + MODE_OPTIONS.map((option) => [option.value, option]), +); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx index 73d9fb8123e..c67758c0d98 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx @@ -33,7 +33,8 @@ import { } from "./GuardrailFormField"; import ContentFilterManager, { formatContentFilterDataForAPI } from "./content_filter/ContentFilterManager"; import CustomCodeModal, { EditGuardrailData } from "./custom_code/CustomCodeModal"; -import { GuardrailModeCard, GuardrailModeRows } from "./GuardrailModeDisplay"; +import { GuardrailModeCard } from "./GuardrailModeDisplay"; +import { GuardrailReadOnlyDetails } from "./GuardrailReadOnlyDetails"; import { getLoggingOnlyScopeUpdate, getGuardrailLogoAndName, @@ -41,10 +42,15 @@ import { loggingOnlyScopeToChoice, skipSystemMessageToChoice, skipToolMessageToChoice, + streamScopeByModeFromConfig, + streamScopeForUpdate, supportsDirectionalLoggingOnlyScope, + toModeArray, type SkipSystemMessageChoice, type SkipToolMessageChoice, + type GuardrailStreamScope, } from "./guardrail_info_helpers"; +import { GuardrailStreamScopeCaption, StreamScopeFormField } from "./StreamScopeFields"; import GuardrailOptionalParams from "./guardrail_optional_params"; import GuardrailProviderFields from "./guardrail_provider_fields"; import PiiConfiguration from "./pii_configuration"; @@ -236,6 +242,13 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, "skip_tool_message_choice", skipToolMessageToChoice(guardrailData.litellm_params?.skip_tool_message_in_guardrail), ); + form.setValue( + "stream_scope_by_mode", + streamScopeByModeFromConfig( + guardrailData.litellm_params?.stream_scope, + toModeArray(guardrailData.litellm_params?.mode), + ), + ); form.setValue( "guardrail_info", guardrailData.guardrail_info ? JSON.stringify(guardrailData.guardrail_info, null, 2) : "", @@ -304,6 +317,16 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, updateData.litellm_params.default_on = values.default_on; } + const modes = toModeArray(guardrailData.litellm_params?.mode); + const nextStreamScope = streamScopeForUpdate( + modes, + (values.stream_scope_by_mode as Record | undefined) ?? {}, + guardrailData.litellm_params?.stream_scope, + ); + if (nextStreamScope !== undefined) { + updateData.litellm_params.stream_scope = nextStreamScope; + } + const prevSkipChoice = skipSystemMessageToChoice(guardrailData.litellm_params?.skip_system_message_in_guardrail); const nextSkipChoice = values.skip_system_message_choice as SkipSystemMessageChoice | undefined; if (nextSkipChoice !== undefined && nextSkipChoice !== prevSkipChoice) { @@ -566,6 +589,7 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, +

Created At

@@ -723,6 +747,11 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, )} + + = ({ guardrailId, onClose, ) : ( -
-
-

Guardrail ID

-
{guardrailData.guardrail_id}
-
-
-

Guardrail Name

-
{guardrailData.guardrail_name || "Unnamed Guardrail"}
-
-
-

Provider

-
{displayName}
-
- -
-

Default On

- - {guardrailData.litellm_params?.default_on ? "Yes" : "No"} - -
- - {guardrailData.litellm_params?.pii_entities_config && - Object.keys(guardrailData.litellm_params.pii_entities_config).length > 0 && ( -
-

PII Protection

-
- - {Object.keys(guardrailData.litellm_params.pii_entities_config).length} PII entities - configured - -
-
- )} - -
-

Created At

-
{formatDate(guardrailData.created_at)}
-
-
-

Last Updated

-
{formatDate(guardrailData.updated_at)}
-
- - {guardrailData.litellm_params?.guardrail === "tool_permission" && ( - - )} -
+ )}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx index 96ade4b8c1b..d149ef98328 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx @@ -15,6 +15,11 @@ import { skipToolMessageToChoice, choiceToSkipToolForCreate, formatGuardrailMode, + formatGuardrailStreamScope, + streamScopeByModeFromConfig, + streamScopeForMode, + streamScopeForUpdate, + streamScopePayload, loggingOnlyScopeToChoice, choiceToLoggingOnlyScope, getLoggingOnlyScopeUpdate, @@ -369,4 +374,56 @@ describe("guardrail_info_helpers", () => { expect(choiceToSkipToolForCreate("no")).toBe(false); }); }); + + describe("stream_scope helpers", () => { + it("treats omitted config as both for every mode", () => { + expect(streamScopeForMode(undefined, "pre_call")).toBe("both"); + expect(streamScopeByModeFromConfig(undefined, ["pre_call", "post_call"])).toEqual({ + pre_call: "both", + post_call: "both", + }); + }); + + it("applies a scalar to every selected mode and omits both-only payloads", () => { + expect(streamScopeForMode("streaming", "post_call")).toBe("streaming"); + expect(streamScopePayload(["pre_call", "post_call"], { pre_call: "both", post_call: "both" })).toBeUndefined(); + expect(streamScopePayload(["pre_call", "post_call"], { pre_call: "streaming", post_call: "streaming" })).toBe( + "streaming", + ); + }); + + it("keeps a mixed map instead of collapsing it to a scalar", () => { + expect(streamScopePayload(["pre_call", "post_call"], { pre_call: "both", post_call: "streaming" })).toEqual({ + post_call: "streaming", + }); + }); + + it("formats scalar and per-mode stream scopes for display", () => { + expect(formatGuardrailStreamScope(undefined)).toBe(""); + expect(formatGuardrailStreamScope("both")).toBe("Streaming and non-streaming"); + expect(formatGuardrailStreamScope("non_streaming")).toBe("Non-streaming only"); + expect(formatGuardrailStreamScope({ post_call: "streaming", pre_call: "both" })).toBe( + "post_call: Streaming only, pre_call: Streaming and non-streaming", + ); + }); + + it("emits both on update only when a prior restriction is cleared", () => { + expect(streamScopeForUpdate(["post_call"], { post_call: "streaming" }, undefined)).toBe("streaming"); + expect(streamScopeForUpdate(["post_call"], { post_call: "both" }, "streaming")).toBe("both"); + expect(streamScopeForUpdate(["post_call"], { post_call: "both" }, undefined)).toBeUndefined(); + }); + + it("keeps stored restrictions for modes outside the current selection", () => { + expect( + streamScopeForUpdate( + ["pre_call"], + { pre_call: "streaming" }, + { pre_call: "streaming", post_call: "non_streaming" }, + ), + ).toBeUndefined(); + expect( + streamScopeForUpdate(["pre_call"], { pre_call: "both" }, { pre_call: "streaming", post_call: "non_streaming" }), + ).toEqual({ post_call: "non_streaming" }); + }); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index f55148b7315..a898f3b9abd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -343,3 +343,86 @@ export function choiceToSkipToolForCreate(choice: SkipToolMessageChoice | undefi if (choice === "no") return false; return undefined; } + +export const GUARDRAIL_STREAM_SCOPES = ["both", "streaming", "non_streaming"] as const; +export type GuardrailStreamScope = (typeof GUARDRAIL_STREAM_SCOPES)[number]; + +export const STREAM_SCOPE_OPTIONS: { value: GuardrailStreamScope; label: string }[] = [ + { value: "both", label: "Streaming and non-streaming" }, + { value: "streaming", label: "Streaming only" }, + { value: "non_streaming", label: "Non-streaming only" }, +]; + +export const isGuardrailStreamScope = (value: unknown): value is GuardrailStreamScope => + value === "both" || value === "streaming" || value === "non_streaming"; + +export const streamScopeForMode = (raw: unknown, mode: string): GuardrailStreamScope => { + if (isGuardrailStreamScope(raw)) return raw; + if (raw !== null && typeof raw === "object" && !Array.isArray(raw)) { + const value = (raw as Record)[mode]; + if (isGuardrailStreamScope(value)) return value; + } + return "both"; +}; + +export const streamScopeByModeFromConfig = (raw: unknown, modes: string[]): Record => + Object.fromEntries(modes.map((mode) => [mode, streamScopeForMode(raw, mode)])); + +export const streamScopePayload = ( + modes: string[], + scopes: Record, +): GuardrailStreamScope | Record | undefined => { + const perMode: Record = Object.fromEntries( + modes.map((mode) => [mode, scopes[mode] ?? "both"]), + ); + const values = Object.values(perMode); + if (values.length === 0 || values.every((scope) => scope === "both")) return undefined; + const unique = new Set(values); + if (unique.size === 1) return values[0]; + return Object.fromEntries(Object.entries(perMode).filter((entry) => entry[1] !== "both")); +}; + +export const formatGuardrailStreamScope = (raw: unknown): string => { + if (isGuardrailStreamScope(raw)) { + return STREAM_SCOPE_OPTIONS.find((option) => option.value === raw)?.label ?? raw; + } + if (raw !== null && typeof raw === "object" && !Array.isArray(raw)) { + const entries = Object.entries(raw as Record).filter( + (entry): entry is [string, GuardrailStreamScope] => isGuardrailStreamScope(entry[1]), + ); + if (entries.length === 0) return ""; + return entries.map(([mode, scope]) => `${mode}: ${formatGuardrailStreamScope(scope)}`).join(", "); + } + return ""; +}; + +export const streamScopeForUpdate = ( + modes: string[], + nextByMode: Record, + previousRaw: unknown, + previousModes: string[] = modes, +): GuardrailStreamScope | Record | undefined => { + const previousMap: Record = + previousRaw !== null && typeof previousRaw === "object" && !Array.isArray(previousRaw) + ? (previousRaw as Record) + : {}; + const preserved: Record = Object.fromEntries( + Object.entries(previousMap).filter( + (entry): entry is [string, GuardrailStreamScope] => !modes.includes(entry[0]) && isGuardrailStreamScope(entry[1]), + ), + ); + const nextModes = [...modes, ...Object.keys(preserved)]; + const nextStreamScope = streamScopePayload(nextModes, { + ...preserved, + ...Object.fromEntries(modes.map((mode) => [mode, nextByMode[mode] ?? "both"])), + }); + const previousCompareModes = Object.keys(previousMap).length > 0 ? Object.keys(previousMap) : previousModes; + const previousStreamScope = streamScopePayload( + previousCompareModes, + streamScopeByModeFromConfig(previousRaw, previousCompareModes), + ); + if (JSON.stringify(nextStreamScope ?? "both") === JSON.stringify(previousStreamScope ?? "both")) { + return undefined; + } + return nextStreamScope ?? "both"; +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts index 7e7089e685f..32dac9c2f65 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts @@ -110,13 +110,14 @@ export const useKeys = ( page: number, pageSize: number, options: KeyListCallOptions = {}, + queryOptions: { enabled?: boolean } = {}, ): UseQueryResult => { const { accessToken } = useAuthorized(); return useQuery({ queryKey: keyKeys.list({ page, limit: pageSize, ...options }), queryFn: async () => await keyListCall(accessToken!, page, pageSize, options), - enabled: Boolean(accessToken), + enabled: Boolean(accessToken) && (queryOptions.enabled ?? true), staleTime: 30000, // 30 seconds placeholderData: keepPreviousData, }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModelCostMap.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModelCostMap.test.ts index f79ca33bc5d..68983f5146e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModelCostMap.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModelCostMap.test.ts @@ -141,4 +141,23 @@ describe("useModelCostMap", () => { expect(result.current).toHaveProperty("isSuccess"); expect(result.current).toHaveProperty("error"); }); + + it("fetches the catalog-only map under its own cache entry when catalogOnly is set", async () => { + (modelCostMap as any).mockImplementation(async (catalogOnly: boolean) => + catalogOnly ? { catalog: { litellm_provider: "bedrock" } } : mockModelCostData, + ); + + const { result: live } = renderHook(() => useModelCostMap(), { wrapper }); + const { result: catalog } = renderHook(() => useModelCostMap(true, true), { wrapper }); + + await waitFor(() => { + expect(live.current.isSuccess).toBe(true); + expect(catalog.current.isSuccess).toBe(true); + }); + + expect(live.current.data).toEqual(mockModelCostData); + expect(catalog.current.data).toEqual({ catalog: { litellm_provider: "bedrock" } }); + expect(modelCostMap).toHaveBeenCalledWith(false); + expect(modelCostMap).toHaveBeenCalledWith(true); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModelCostMap.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModelCostMap.ts index d9824b4753e..cf97a504e1e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModelCostMap.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModelCostMap.ts @@ -2,13 +2,13 @@ import { modelCostMap } from "@/components/networking"; import { useQuery } from "@tanstack/react-query"; import { createQueryKeys } from "../common/queryKeysFactory"; -const modelCostMapKeys = createQueryKeys("modelCostMap"); +export const modelCostMapKeys = createQueryKeys("modelCostMap"); -export const useModelCostMap = (enabled = true) => { +export const useModelCostMap = (enabled = true, catalogOnly = false) => { return useQuery>({ enabled, - queryKey: modelCostMapKeys.list({}), - queryFn: async () => await modelCostMap(), + queryKey: modelCostMapKeys.list(catalogOnly ? { filters: { catalog_only: "true" } } : {}), + queryFn: async () => await modelCostMap(catalogOnly), staleTime: 60 * 1000, // 1 minute gcTime: 60 * 1000, // 1 minute }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts index 6489bc2171d..09a0f214da3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts @@ -1,5 +1,5 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { renderHook, waitFor } from "@testing-library/react"; +import { act, renderHook, waitFor } from "@testing-library/react"; import React, { ReactNode } from "react"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { @@ -11,6 +11,7 @@ import { useAutoRouters, useInfiniteModelInfo, useModelHub, + useModelAccessGroupNames, useModelsInfo, usePlainChatModelGroups, useSelectedTeamModels, @@ -565,6 +566,127 @@ describe("useUserModels", () => { }); }); +describe("useModelAccessGroupNames", () => { + let queryClient: QueryClient; + + beforeEach(() => { + queryClient = new QueryClient({ + defaultOptions: { + queries: { + retry: false, + }, + }, + }); + + vi.clearAllMocks(); + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId: "test-user-id", + userRole: "Admin", + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); + }); + + const wrapper = ({ children }: { children: ReactNode }) => + React.createElement(QueryClientProvider, { client: queryClient }, children); + + it("fetches and returns the caller's model access group names", async () => { + vi.mocked(modelAvailableCall).mockResolvedValue({ + data: [ + { id: "repro-access-group", object: "model", created: 0, owned_by: "litellm" }, + { id: "another-access-group", object: "model", created: 0, owned_by: "litellm" }, + ], + }); + + const { result } = renderHook(() => useModelAccessGroupNames(), { wrapper }); + + await waitFor(() => { + expect(result.current).toEqual(new Set(["repro-access-group", "another-access-group"])); + }); + + expect(modelAvailableCall).toHaveBeenCalledWith( + "test-access-token", + "test-user-id", + "Admin", + false, + null, + true, + true, + ); + }); + + it("returns undefined while the access-group lookup is pending", () => { + vi.mocked(modelAvailableCall).mockReturnValue(new Promise(() => undefined)); + + const { result } = renderHook(() => useModelAccessGroupNames(), { wrapper }); + + expect(result.current).toBeUndefined(); + }); + + it("returns undefined until authorization is ready", () => { + const unauthorizedContext = { + accessToken: null, + userId: null, + userRole: null, + token: null, + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }; + mockUseAuthorized.mockReturnValue(unauthorizedContext); + + const { result } = renderHook(() => useModelAccessGroupNames(), { wrapper }); + + expect(result.current).toBeUndefined(); + expect(modelAvailableCall).not.toHaveBeenCalled(); + }); + + it("returns an empty set when the access-group lookup fails", async () => { + vi.mocked(modelAvailableCall).mockRejectedValue(new Error("lookup failed")); + + const { result } = renderHook(() => useModelAccessGroupNames(), { wrapper }); + + await waitFor(() => expect(result.current).toBeDefined()); + expect(result.current?.size).toBe(0); + }); + + it("keeps cached access-group names after a failed refetch", async () => { + vi.mocked(modelAvailableCall).mockResolvedValueOnce({ + data: [{ id: "repro-access-group", object: "model", created: 0, owned_by: "litellm" }], + }); + + const { result } = renderHook(() => useModelAccessGroupNames(), { wrapper }); + + await waitFor(() => { + expect(result.current?.has("repro-access-group")).toBe(true); + }); + + const queryKey = queryClient + .getQueryCache() + .getAll() + .find((query) => { + return query.queryKey[0] === "modelAccessGroupNames"; + })?.queryKey; + expect(queryKey).toBeDefined(); + if (!queryKey) throw new Error("The access-group query was not created"); + + vi.mocked(modelAvailableCall).mockRejectedValueOnce(new Error("refetch failed")); + await act(async () => { + await queryClient.refetchQueries({ queryKey }); + }); + + await waitFor(() => { + expect(queryClient.getQueryCache().find({ queryKey })?.state.status).toBe("error"); + }); + expect(result.current).toEqual(new Set(["repro-access-group"])); + }); +}); + describe("useSelectedTeamModels", () => { let queryClient: QueryClient; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts index b5cc329d4c9..8bc2a50b4d8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts @@ -1,4 +1,5 @@ import { useQuery, useInfiniteQuery, useQueryClient, UseQueryResult } from "@tanstack/react-query"; +import { useMemo } from "react"; import { createQueryKeys } from "../common/queryKeysFactory"; import { modelInfoCall, modelHubCall, modelAvailableCall } from "@/components/networking"; import useAuthorized from "../useAuthorized"; @@ -29,6 +30,7 @@ const allProxyModelsKeys = createQueryKeys("allProxyModels"); const selectedTeamModelsKeys = createQueryKeys("selectedTeamModels"); const infiniteModelKeys = createQueryKeys("infiniteModels"); const userModelsKeys = createQueryKeys("userModels"); +const modelAccessGroupNameKeys = createQueryKeys("modelAccessGroupNames"); export const useModelsInfo = ( page: number = 1, @@ -265,6 +267,31 @@ export const useUserModels = (): UseQueryResult => { }); }; +export const useModelAccessGroupNames = (): ReadonlySet | undefined => { + const { accessToken, userId, userRole } = useAuthorized(); + const { data, isError } = useQuery({ + queryKey: modelAccessGroupNameKeys.list({}), + queryFn: async () => { + const response: AllProxyModelsResponse = await modelAvailableCall( + accessToken!, + userId!, + userRole!, + false, + null, + true, + true, + ); + return response.data.map((model) => model.id); + }, + enabled: Boolean(accessToken && userId && userRole), + }); + return useMemo(() => { + if (data !== undefined) return new Set(data); + if (isError) return new Set(); + return undefined; + }, [data, isError]); +}; + export const useSelectedTeamModels = (teamID: string | null) => { const { accessToken, userId, userRole } = useAuthorized(); return useQuery({ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts index 7d1d035b4d4..2f46e548a7b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts @@ -118,13 +118,14 @@ export const useTeamsTable = ( }; export const teamKeys = createQueryKeys("teams"); -export const useTeams = (): UseQueryResult => { +export const useTeams = (queryOptions: { enabled?: boolean } = {}): UseQueryResult => { const { accessToken, userId, userRole } = useAuthorized(); - return useQuery({ + const teamsQueryOptions = { queryKey: teamKeys.list({}), queryFn: async () => await fetchTeams(accessToken!, userId, userRole, null), - enabled: Boolean(accessToken), - }); + enabled: Boolean(accessToken) && queryOptions.enabled !== false, + }; + return useQuery(teamsQueryOptions); }; const ALL_TEAMS_PAGE_SIZE = 100; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx index 4762dd2fed1..81d290cdfe4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx @@ -23,6 +23,7 @@ import { Sheet, SheetContent, SheetTitle, SheetTrigger } from "@/components/ui/s import { Button } from "@/components/ui/button"; import { Menu } from "lucide-react"; import { useMediaQuery } from "usehooks-ts"; +import { CommandPaletteProvider } from "@/components/CommandPalette/CommandPaletteProvider"; const pluginApiClient = createApiClient({ getBaseUrl: () => getProxyBaseUrl() ?? "" }); @@ -168,26 +169,28 @@ function DashboardShell({ children }: { children: React.ReactNode }) { setMobileNavigationKey(null)} /> -
- - } - > - - - } - /> - - - - - - -
{children}
-
+ +
+ + } + > + + + } + /> + + + + + + +
{children}
+
+
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.test.tsx index 7cd418bc176..b2beb5beeed 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.test.tsx @@ -1,17 +1,39 @@ /* @vitest-environment jsdom */ -import { render, screen } from "@testing-library/react"; +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { fireEvent, render, screen } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; +import { modelCostMapKeys } from "../../hooks/models/useModelCostMap"; import PriceDataManagementTab from "./PriceDataManagementTab"; -vi.mock("@/components/price_data_reload", () => ({ default: () =>
reload
})); -vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => ({ accessToken: "sk-test" }) })); -vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({ - useModelCostMap: () => ({ refetch: vi.fn() }), +vi.mock("@/components/price_data_reload", () => ({ + default: ({ onReloadSuccess }: { onReloadSuccess: () => void }) => , })); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => ({ accessToken: "sk-test" }) })); + +const renderTab = (queryClient: QueryClient) => + render( + + + , + ); describe("PriceDataManagementTab", () => { it("renders its content standalone, without a tab-panel ancestor", () => { - render(); + renderTab(new QueryClient()); expect(screen.getByText("Price Data Management")).toBeInTheDocument(); }); + + it("invalidates both the live and the catalog-only cost map after a reload", () => { + const queryClient = new QueryClient(); + const liveKey = modelCostMapKeys.list({}); + const catalogKey = modelCostMapKeys.list({ filters: { catalog_only: "true" } }); + queryClient.setQueryData(liveKey, {}); + queryClient.setQueryData(catalogKey, {}); + renderTab(queryClient); + + fireEvent.click(screen.getByRole("button", { name: "reload" })); + + expect(queryClient.getQueryState(liveKey)?.isInvalidated).toBe(true); + expect(queryClient.getQueryState(catalogKey)?.isInvalidated).toBe(true); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.tsx index 126f6970c2e..ce6f8e77d38 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.tsx @@ -1,11 +1,12 @@ import PriceDataReload from "@/components/price_data_reload"; import React from "react"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -import { useModelCostMap } from "../../hooks/models/useModelCostMap"; +import { useQueryClient } from "@tanstack/react-query"; +import { modelCostMapKeys } from "../../hooks/models/useModelCostMap"; const PriceDataManagementTab = () => { const { accessToken } = useAuthorized(); - const { refetch: refetchModelCostMap } = useModelCostMap(); + const queryClient = useQueryClient(); return (
@@ -19,7 +20,7 @@ const PriceDataManagementTab = () => { { - refetchModelCostMap(); + queryClient.invalidateQueries({ queryKey: modelCostMapKeys.all }); }} buttonText="Reload Price Data" size="middle" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.integration.test.tsx index 865cc808627..86b5659d570 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.integration.test.tsx @@ -84,6 +84,7 @@ vi.mock("@/app/(dashboard)/hooks/uiSettings/useUISettings", () => ({ vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({ useAllProxyModels: vi.fn(() => ({ data: { data: [] }, isLoading: false })), + useModelAccessGroupNames: vi.fn(() => new Set()), })); vi.mock("@/app/(dashboard)/hooks/teams/useTeams", () => ({ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx index efba26734ff..844ac9f44e2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx @@ -1,4 +1,5 @@ -import { renderWithProviders, screen, waitFor } from "../../../../../tests/test-utils"; +import { useSyncExternalStore } from "react"; +import { act, renderWithProviders, screen, waitFor } from "../../../../../tests/test-utils"; import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import AddModelPanel from "./AddModelPanel"; @@ -22,7 +23,11 @@ vi.mock("@/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled", () => usePtuCostAttributionEnabled: () => mockPtuEnabled(), })); -vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({ useModelCostMap: () => ({ data: {} }) })); +const mockUseModelCostMap = vi.fn(); + +vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({ + useModelCostMap: (...args: unknown[]) => mockUseModelCostMap(...args), +})); vi.mock("@/app/(dashboard)/hooks/credentials/useCredentials", () => ({ useCredentials: () => ({ data: { credentials: [] } }), @@ -60,6 +65,13 @@ vi.mock("@/app/(dashboard)/hooks/providers/useProviderFields", () => ({ { key: "api_base", label: "API Base", field_type: "text", required: false }, ], }, + { + provider: "Anthropic", + provider_display_name: "Anthropic", + litellm_provider: "anthropic", + default_model_placeholder: "claude-3-opus", + credential_fields: [{ key: "api_key", label: "API Key", field_type: "password", required: false }], + }, ], isLoading: false, error: null, @@ -70,6 +82,26 @@ vi.mock("@/components/vector_store_management/VectorStoreSelector", () => ({ default: () =>
, })); +type Catalog = Record; + +const createCatalogFeed = () => { + let current: Catalog | undefined; + const listeners = new Set<() => void>(); + const subscribe = (listener: () => void) => { + listeners.add(listener); + return () => { + listeners.delete(listener); + }; + }; + return { + useData: () => ({ data: useSyncExternalStore(subscribe, () => current) }), + publish: (next: Catalog) => { + current = next; + listeners.forEach((listener) => listener()); + }, + }; +}; + const lastCreatedModel = () => modelCreateCall.mock.calls.at(-1)?.[1]; const PROXY_ADMIN = { @@ -138,10 +170,93 @@ const setup = async () => { describe("AddModelPanel submit payload contract", () => { beforeEach(() => { vi.clearAllMocks(); + mockUseModelCostMap.mockReturnValue({ data: {} }); mockPtuEnabled.mockReturnValue(false); mockAuthorized.mockReturnValue(PROXY_ADMIN); }); + it("lists catalog models in the model picker, not entries registered at runtime for deployments", async () => { + mockUseModelCostMap.mockImplementation((_enabled: boolean, catalogOnly: boolean) => ({ + data: catalogOnly + ? { "gpt-4o-2024-08-06": { litellm_provider: "openai" } } + : { + "gpt-4o-2024-08-06": { litellm_provider: "openai" }, + "openai-gpt-4o-deployment-id": { litellm_provider: "openai" }, + }, + })); + const { user } = await setup(); + await user.click(screen.getByRole("combobox", { name: /provider/i })); + await user.click(await screen.findByText("OpenAI")); + await user.click(await screen.findByPlaceholderText("Select models")); + + expect(await screen.findByText("gpt-4o-2024-08-06")).toBeInTheDocument(); + expect(screen.queryByText("openai-gpt-4o-deployment-id")).not.toBeInTheDocument(); + }); + + it("offers the catalog models once the catalog arrives after the provider was picked", async () => { + const catalog = createCatalogFeed(); + mockUseModelCostMap.mockImplementation(() => catalog.useData()); + const { user } = await setup(); + await user.click(screen.getByRole("combobox", { name: /provider/i })); + await user.click(await screen.findByText("OpenAI")); + expect(await screen.findByPlaceholderText("gpt-3.5-turbo")).toBeInTheDocument(); + + act(() => catalog.publish({ "gpt-4o-2024-08-06": { litellm_provider: "openai" } })); + await user.click(await screen.findByPlaceholderText("Select models")); + + expect(await screen.findByText("gpt-4o-2024-08-06")).toBeInTheDocument(); + }); + + it("keeps a model name typed before the catalog arrives instead of swapping the field under the user", async () => { + const catalog = createCatalogFeed(); + mockUseModelCostMap.mockImplementation(() => catalog.useData()); + const { user } = await setup(); + await user.click(screen.getByRole("combobox", { name: /provider/i })); + await user.click(await screen.findByText("OpenAI")); + await user.type(await screen.findByPlaceholderText("gpt-3.5-turbo"), "my-fine-tune"); + + act(() => catalog.publish({ "gpt-4o-2024-08-06": { litellm_provider: "openai" } })); + + await waitFor(() => expect(screen.getByPlaceholderText("gpt-3.5-turbo")).toHaveValue("my-fine-tune")); + expect(screen.queryByPlaceholderText("Select models")).not.toBeInTheDocument(); + }); + + it("offers the catalog models when a name typed before the catalog arrived was cleared again", async () => { + const catalog = createCatalogFeed(); + mockUseModelCostMap.mockImplementation(() => catalog.useData()); + const { user } = await setup(); + await user.click(screen.getByRole("combobox", { name: /provider/i })); + await user.click(await screen.findByText("OpenAI")); + const typed = await screen.findByPlaceholderText("gpt-3.5-turbo"); + await user.type(typed, "my-fine-tune"); + await user.clear(typed); + + act(() => catalog.publish({ "gpt-4o-2024-08-06": { litellm_provider: "openai" } })); + await user.click(await screen.findByPlaceholderText("Select models")); + + expect(await screen.findByText("gpt-4o-2024-08-06")).toBeInTheDocument(); + }); + + it("swaps the offered models when the provider changes while the catalog stays the same", async () => { + const catalog = createCatalogFeed(); + catalog.publish({ + "gpt-4o-2024-08-06": { litellm_provider: "openai" }, + "claude-sonnet-4-5": { litellm_provider: "anthropic" }, + }); + mockUseModelCostMap.mockImplementation(() => catalog.useData()); + const { user } = await setup(); + const provider = screen.getByRole("combobox", { name: /provider/i }); + await user.click(provider); + await user.click(await screen.findByText("OpenAI")); + await user.clear(provider); + await user.type(provider, "Anthropic"); + await user.click(await screen.findByText("Anthropic")); + await user.click(await screen.findByPlaceholderText("Select models")); + + expect(await screen.findByText("claude-sonnet-4-5")).toBeInTheDocument(); + expect(screen.queryByText("gpt-4o-2024-08-06")).not.toBeInTheDocument(); + }); + it("sends only the always-mounted fields while Advanced Settings stays closed", async () => { const { fillRequired, submit } = await setup(); await fillRequired(); @@ -295,6 +410,7 @@ describe("AddModelPanel submit payload contract", () => { describe("AddModelPanel empty-string skip", () => { beforeEach(() => { vi.clearAllMocks(); + mockUseModelCostMap.mockReturnValue({ data: {} }); mockPtuEnabled.mockReturnValue(false); mockAuthorized.mockReturnValue(PROXY_ADMIN); }); @@ -376,6 +492,7 @@ describe("AddModelPanel validation gates", () => { describe("AddModelPanel behaviours the removed Advanced Settings form instance never drove", () => { beforeEach(() => { vi.clearAllMocks(); + mockUseModelCostMap.mockReturnValue({ data: {} }); mockPtuEnabled.mockReturnValue(false); mockAuthorized.mockReturnValue(PROXY_ADMIN); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx index 1cdce04d07c..9fbd138bda4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx @@ -1,7 +1,7 @@ "use client"; -import { useState } from "react"; -import { useForm } from "react-hook-form"; +import { useMemo, useState } from "react"; +import { useForm, useWatch } from "react-hook-form"; import { useQueryClient } from "@tanstack/react-query"; import AddModelForm from "@/components/add_model/AddModelForm"; import { handleAddModelSubmit } from "@/components/add_model/handle_add_model_submit"; @@ -23,11 +23,15 @@ export default function AddModelPanel() { const form = useForm({ mode: "onChange", defaultValues: INITIAL_VALUES }); const registry = useMountRegistry(); const queryClient = useQueryClient(); - const { data: modelCostMapData } = useModelCostMap(); + const { data: modelCostMapData } = useModelCostMap(true, true); const { data: credentialsResponse } = useCredentials(); const { data: teams } = useTeams(); const [selectedProvider, setSelectedProvider] = useState(Providers.Anthropic); - const [providerModels, setProviderModels] = useState([]); + const pickedProvider = useWatch({ control: form.control, name: "custom_llm_provider" }); + const providerModels = useMemo( + () => (typeof pickedProvider === "string" ? getProviderModels(pickedProvider, modelCostMapData) : []), + [pickedProvider, modelCostMapData], + ); const [showAdvancedSettings, setShowAdvancedSettings] = useState(false); const refresh = () => queryClient.invalidateQueries({ queryKey: ["models", "list"] }); @@ -57,9 +61,6 @@ export default function AddModelPanel() { selectedProvider={selectedProvider} setSelectedProvider={setSelectedProvider} providerModels={providerModels} - setProviderModelsFn={(provider) => - setProviderModels(provider === null ? [] : getProviderModels(provider, modelCostMapData)) - } getPlaceholder={getPlaceholder} showAdvancedSettings={showAdvancedSettings} setShowAdvancedSettings={setShowAdvancedSettings} diff --git a/ui/litellm-dashboard/src/components/CommandPalette/CommandPalette.integration.test.tsx b/ui/litellm-dashboard/src/components/CommandPalette/CommandPalette.integration.test.tsx new file mode 100644 index 00000000000..c722194d496 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CommandPalette/CommandPalette.integration.test.tsx @@ -0,0 +1,442 @@ +import { act, fireEvent, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; +import type { KeyResponse } from "@/components/key_team_helpers/key_list"; +import { writeStorage } from "@/lib/storage"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; +import { keyDetailHref } from "@/utils/entityLinks"; +import { uiHref } from "@/utils/uiHref"; +import { renderWithProviders, testQueryClient } from "../../../tests/test-utils"; +import { COMMAND_PALETTE_HINT_KEY } from "./CommandPaletteHint"; +import { CommandPaletteProvider } from "./CommandPaletteProvider"; +import { CommandPaletteTrigger } from "./CommandPaletteTrigger"; + +const mocks = vi.hoisted(() => ({ + pathname: "/api-keys", + push: vi.fn<(href: string) => void>(), + useKeys: vi.fn(), + teamAlias: "platform-team", + isPlaceholderData: false, +})); + +vi.mock("next/navigation", () => ({ + usePathname: () => mocks.pathname, + useRouter: () => ({ push: mocks.push }), +})); + +vi.mock("@/app/(dashboard)/hooks/keys/useKeys", () => ({ + useKeys: mocks.useKeys, +})); +vi.mock("@/app/(dashboard)/hooks/teams/useTeams", () => ({ + useTeams: () => ({ + data: [{ team_id: "team-1", team_alias: mocks.teamAlias, members_with_roles: [] }], + }), +})); +vi.mock("@/app/(dashboard)/hooks/useIsOrgAdmin", () => ({ + default: () => false, +})); +vi.mock("@/app/(dashboard)/hooks/uiSettings/useUISettings", () => ({ + useUISettings: () => ({ data: { values: {} } }), +})); + +const HIGH_VOLUME_KEY = keyFixture("key-prod-0", "high-volume-prod", "sk-...4zoA", "team-1"); +const PROD_KEY = keyFixture("key-prod-1", "prod-backend", "sk-...AHeA", "team-1"); +const STAGING_KEY = keyFixture("key-staging-1", "staging-agent", "sk-...stgB", "unknown-team"); +const KEY_FIXTURES = [HIGH_VOLUME_KEY, PROD_KEY, STAGING_KEY]; +const UI_SETTINGS_RESPONSE = { + server_root_path: "", + proxy_base_url: null, + admin_ui_disabled: false, + auto_redirect_to_sso: false, + sso_configured: false, + is_control_plane: false, + workers: [], +}; +const PROD_KEY_LIST_OPTIONS = { search: "prod", sortBy: "created_at", sortOrder: "desc", expand: "user" }; +const cachedKeyResults = new Map(); + +function keyFixture(token: string, alias: string, keyName: string, teamId: string | null = null): KeyResponse { + return { + token, + token_id: token, + key_name: keyName, + key_alias: alias, + spend: 2.5, + team_id: teamId, + team_alias: "", + } as KeyResponse; +} + +function sessionCookie() { + const encode = (value: object) => + btoa(JSON.stringify(value)).replaceAll("=", "").replaceAll("+", "-").replaceAll("/", "_"); + const claims = { + key: "sk-session-test", + user_id: "test-admin", + user_role: "proxy_admin", + premium_user: true, + auth_header_name: "X-Gateway-Session", + exp: Date.now() / 1000 + 3600, + }; + document.cookie = `token=${encode({ alg: "none" })}.${encode(claims)}.test; Path=/`; +} + +function renderPalette() { + return renderWithProviders( + + + , + ); +} + +beforeEach(() => { + testQueryClient.clear(); + window.localStorage.clear(); + cachedKeyResults.clear(); + vi.clearAllMocks(); + mocks.pathname = "/api-keys"; + mocks.teamAlias = "platform-team"; + mocks.isPlaceholderData = false; + sessionCookie(); + vi.stubGlobal( + "fetch", + vi.fn(async (input: RequestInfo | URL) => + String(input).includes("/v2/team/list") + ? Response.json({ + teams: [{ team_id: "team-1", team_alias: "platform-team" }], + total_pages: 1, + }) + : Response.json(UI_SETTINGS_RESPONSE), + ), + ); + vi.mocked(useKeys).mockImplementation((_page, _pageSize, options) => { + const query = options.search?.toLowerCase() ?? ""; + const keys = + cachedKeyResults.get(query) ?? + KEY_FIXTURES.filter( + (key) => !query || key.key_alias.toLowerCase().includes(query) || key.key_name.toLowerCase().includes(query), + ); + cachedKeyResults.set(query, keys); + return { + data: { keys, total_count: keys.length, current_page: 1, total_pages: 1 }, + isFetching: false, + isError: false, + isPlaceholderData: mocks.isPlaceholderData, + } as ReturnType; + }); +}); + +afterEach(() => { + vi.useRealTimers(); + vi.restoreAllMocks(); + testQueryClient.clear(); + vi.unstubAllGlobals(); + document.cookie = "token=; Max-Age=0; Path=/"; +}); + +describe("CommandPalette integration", () => { + it("shows the discovery hint on the keys route", async () => { + renderPalette(); + + expect(await screen.findByText("to search keys")).toBeVisible(); + }); + + it("opens the palette from the hint and persists that it was seen", async () => { + const user = userEvent.setup(); + renderPalette(); + + await user.click(screen.getByRole("button", { name: /to search keys/ })); + + expect(await screen.findByRole("combobox", { name: "Search" })).toBeVisible(); + expect(screen.queryByText("to search keys")).not.toBeInTheDocument(); + expect(window.localStorage.getItem(COMMAND_PALETTE_HINT_KEY.name)).toBe("true"); + }); + + it("hides and persists the discovery hint when Ctrl+K opens the palette", async () => { + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + + expect(await screen.findByRole("combobox", { name: "Search" })).toBeVisible(); + expect(screen.queryByText("to search keys")).not.toBeInTheDocument(); + expect(window.localStorage.getItem(COMMAND_PALETTE_HINT_KEY.name)).toBe("true"); + }); + + it("dismisses the discovery hint without opening the palette", async () => { + const user = userEvent.setup(); + renderPalette(); + + await user.click(screen.getByRole("button", { name: "Dismiss search hint" })); + + expect(screen.queryByRole("dialog", { name: "Command palette" })).not.toBeInTheDocument(); + expect(screen.queryByText("to search keys")).not.toBeInTheDocument(); + expect(window.localStorage.getItem(COMMAND_PALETTE_HINT_KEY.name)).toBe("true"); + }); + + it("hides the discovery hint in memory when storage cannot persist its dismissal", async () => { + const user = userEvent.setup(); + vi.spyOn(Storage.prototype, "setItem").mockImplementation(() => { + throw new Error("Storage unavailable"); + }); + renderPalette(); + + await user.click(await screen.findByRole("button", { name: "Dismiss search hint" })); + + expect(screen.queryByText("to search keys")).not.toBeInTheDocument(); + expect(window.localStorage.getItem(COMMAND_PALETTE_HINT_KEY.name)).toBeNull(); + }); + + it("keeps the discovery hint hidden after closing the palette when storage is unavailable", async () => { + const user = userEvent.setup(); + vi.spyOn(Storage.prototype, "setItem").mockImplementation(() => { + throw new Error("Storage unavailable"); + }); + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + expect(screen.queryByText("to search keys")).not.toBeInTheDocument(); + + await user.keyboard("{Escape}"); + + expect(screen.queryByText("to search keys")).not.toBeInTheDocument(); + expect(window.localStorage.getItem(COMMAND_PALETTE_HINT_KEY.name)).toBeNull(); + }); + + it("does not show the discovery hint when it has already been seen", () => { + writeStorage(COMMAND_PALETTE_HINT_KEY, true); + renderPalette(); + + expect(screen.queryByText("to search keys")).not.toBeInTheDocument(); + }); + + it("shows global-search copy outside the keys routes", async () => { + mocks.pathname = "/logs"; + renderPalette(); + + expect(await screen.findByText("to search or jump to a page")).toBeVisible(); + }); + + it("opens and focuses with Ctrl+K, rejects extra modifiers, and toggles closed", async () => { + renderPalette(); + fireEvent.keyDown(document, { key: "k", ctrlKey: true, shiftKey: true }); + fireEvent.keyDown(document, { key: "k", ctrlKey: true, altKey: true }); + expect(screen.queryByRole("dialog", { name: "Command palette" })).not.toBeInTheDocument(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + expect(input).toHaveFocus(); + fireEvent.change(input, { target: { value: "stale" } }); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + expect(screen.queryByRole("dialog", { name: "Command palette" })).not.toBeInTheDocument(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const reopenedInput = await screen.findByRole("combobox", { name: "Search" }); + expect(reopenedInput).toHaveValue(""); + }); + + it("searches keys, shows key details, and opens the selected key", async () => { + const user = userEvent.setup(); + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + expect(input).toHaveAttribute("placeholder", "Search keys by alias or ID…"); + expect(screen.getAllByText("Virtual Keys")).toHaveLength(2); + + await user.type(input, "prod"); + await waitFor(() => { + expect(vi.mocked(useKeys)).toHaveBeenLastCalledWith(1, 8, expect.objectContaining(PROD_KEY_LIST_OPTIONS), { + enabled: true, + }); + }); + const firstKeyOption = await screen.findByRole("option", { name: /high-volume-prod/ }); + const keyOption = await screen.findByRole("option", { name: /prod-backend/ }); + expect(firstKeyOption).toHaveAttribute("aria-selected", "true"); + expect(keyOption).toHaveTextContent("platform-team"); + expect(keyOption).toHaveTextContent("$2.50"); + expect(screen.queryByRole("option", { name: /staging-agent/ })).not.toBeInTheDocument(); + + fireEvent.keyDown(input, { key: "ArrowDown" }); + expect(keyOption).toHaveAttribute("aria-selected", "true"); + fireEvent.keyDown(input, { key: "Enter" }); + expect(mocks.push).toHaveBeenCalledWith(keyDetailHref(PROD_KEY.token)); + expect(screen.queryByRole("dialog", { name: "Command palette" })).not.toBeInTheDocument(); + }); + + it("opens the first matching key when Enter is pressed without navigating", async () => { + const user = userEvent.setup(); + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + await user.type(input, "prod"); + const firstKeyOption = await screen.findByRole("option", { name: /high-volume-prod/ }); + expect(firstKeyOption).toHaveAttribute("aria-selected", "true"); + + fireEvent.keyDown(input, { key: "Enter" }); + expect(mocks.push).toHaveBeenCalledWith(keyDetailHref(HIGH_VOLUME_KEY.token)); + }); + + it("does not activate stale key results while the new query is debouncing", () => { + vi.useFakeTimers(); + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = screen.getByRole("combobox", { name: "Search" }); + fireEvent.change(input, { target: { value: "prod" } }); + vi.advanceTimersByTime(DEBOUNCE_WAIT_MS / 2); + + expect(screen.getByText("Searching…")).toBeVisible(); + expect(screen.queryByRole("option", { name: /high-volume-prod/ })).not.toBeInTheDocument(); + expect(screen.getByRole("option", { name: /Filter the keys table/ })).toBeVisible(); + + fireEvent.keyDown(input, { key: "Enter" }); + + expect(mocks.push).toHaveBeenCalledWith(`${uiHref("api-keys")}?key_search=prod`); + expect(mocks.push).not.toHaveBeenCalledWith(keyDetailHref(HIGH_VOLUME_KEY.token)); + expect(mocks.push).not.toHaveBeenCalledWith(keyDetailHref(PROD_KEY.token)); + }); + + it("does not show placeholder key rows after the query debounce completes", () => { + vi.useFakeTimers(); + mocks.isPlaceholderData = true; + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = screen.getByRole("combobox", { name: "Search" }); + fireEvent.change(input, { target: { value: "prod" } }); + act(() => { + vi.advanceTimersByTime(DEBOUNCE_WAIT_MS); + }); + + expect(screen.getByText("Searching…")).toBeVisible(); + expect(screen.queryByRole("option", { name: /high-volume-prod/ })).not.toBeInTheDocument(); + }); + + it("keeps the same key selected when its subtitle changes", async () => { + const { rerender } = renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + const secondKeyOption = await screen.findByRole("option", { name: /prod-backend/ }); + + fireEvent.keyDown(input, { key: "ArrowDown" }); + expect(secondKeyOption).toHaveAttribute("aria-selected", "true"); + + mocks.teamAlias = "renamed-platform-team"; + rerender( + + + , + ); + + const updatedSecondKeyOption = await screen.findByRole("option", { name: /prod-backend/ }); + expect(updatedSecondKeyOption).toHaveTextContent("renamed-platform-team"); + expect(updatedSecondKeyOption).toHaveAttribute("aria-selected", "true"); + expect(screen.getByRole("option", { name: /high-volume-prod/ })).toHaveAttribute("aria-selected", "false"); + }); + + it("omits an unknown team ID from the key subtitle", async () => { + const user = userEvent.setup(); + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + await user.type(input, "staging"); + const stagingOption = await screen.findByRole("option", { name: /staging-agent/ }); + expect(stagingOption).not.toHaveTextContent("unknown-team"); + }); + + it("keeps only non-empty groups and the filter action when no keys match", async () => { + const user = userEvent.setup(); + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + await user.type(input, "no-match"); + + expect(await screen.findByText("No keys match “no-match”")).toBeVisible(); + expect(screen.getByRole("group", { name: "Keys" })).toHaveTextContent("No keys match “no-match”"); + expect(screen.getByRole("option", { name: /Filter the keys table/ })).toBeVisible(); + expect(screen.queryByRole("group", { name: "Pages" })).not.toBeInTheDocument(); + }); + + it("shows one global no-results message when the query has no matches", async () => { + const user = userEvent.setup(); + mocks.pathname = "/teams"; + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + await user.type(input, "no-match"); + + expect(await screen.findByText("No results for “no-match”")).toBeVisible(); + expect(screen.queryByRole("group")).not.toBeInTheDocument(); + expect(screen.queryByRole("option", { name: /Search virtual keys/ })).not.toBeInTheDocument(); + }); + + it("opens the virtual-key table with the current search applied", async () => { + const user = userEvent.setup(); + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + await user.type(input, "prod"); + await waitFor(() => { + expect(vi.mocked(useKeys)).toHaveBeenLastCalledWith(1, 8, expect.objectContaining({ search: "prod" }), { + enabled: true, + }); + }); + + fireEvent.keyDown(input, { key: "ArrowDown" }); + fireEvent.keyDown(input, { key: "ArrowDown" }); + fireEvent.keyDown(input, { key: "Enter" }); + expect(mocks.push).toHaveBeenCalledWith(`${uiHref("api-keys")}?key_search=prod`); + }); + + it("opens a matching global page from the teams route", async () => { + const user = userEvent.setup(); + mocks.pathname = "/teams"; + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + expect(input).toHaveAttribute("placeholder", "Search pages and actions…"); + await user.type(input, "logs"); + + const logsOption = await screen.findByRole("option", { name: /^Logs/ }); + expect(logsOption).toBeVisible(); + fireEvent.keyDown(input, { key: "Enter" }); + expect(mocks.push).toHaveBeenCalledWith(uiHref("logs")); + }); + + it("clears the global query when switching to key search", async () => { + const user = userEvent.setup(); + mocks.pathname = "/teams"; + renderPalette(); + + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + await user.type(input, "keys"); + + expect(await screen.findByRole("option", { name: /Search virtual keys/ })).toBeVisible(); + fireEvent.keyDown(input, { key: "Enter" }); + + const keysInput = await screen.findByRole("combobox", { name: "Search" }); + expect(keysInput).toHaveValue(""); + expect(await screen.findByText("Recent keys")).toBeVisible(); + }); + + it("switches from an empty keys scope to global search on Backspace", async () => { + renderPalette(); + fireEvent.keyDown(document, { key: "k", ctrlKey: true }); + const input = await screen.findByRole("combobox", { name: "Search" }); + expect(input).toHaveAttribute("placeholder", "Search keys by alias or ID…"); + + fireEvent.keyDown(input, { key: "Backspace" }); + expect(input).toHaveAttribute("placeholder", "Search pages and actions…"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CommandPalette/CommandPalette.tsx b/ui/litellm-dashboard/src/components/CommandPalette/CommandPalette.tsx new file mode 100644 index 00000000000..c8a1d8d0f6c --- /dev/null +++ b/ui/litellm-dashboard/src/components/CommandPalette/CommandPalette.tsx @@ -0,0 +1,319 @@ +"use client"; + +import { useEffect, useMemo, useState, type KeyboardEvent as ReactKeyboardEvent, type ReactNode } from "react"; +import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; +import { FileText, KeyRound, LoaderCircle, Search } from "lucide-react"; +import { usePathname, useRouter } from "next/navigation"; +import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; +import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; +import { Dialog, DialogContent, DialogTitle } from "@/components/ui/dialog"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; +import { keyDetailHref } from "@/utils/entityLinks"; +import { routeSegmentForPathname, uiHref } from "@/utils/uiHref"; +import type { KeyResponse } from "@/components/key_team_helpers/key_list"; +import { CommandPaletteResultRow, type CommandPaletteRowData } from "./CommandPaletteResultRow"; +import { useVisibleMenuGroups } from "./useVisibleMenuGroups"; +import { flattenNavItems, matchNavItems, paletteScopeForRoute, type PaletteNavItem, type PaletteScope } from "./utils"; + +type PaletteEntry = CommandPaletteRowData & + ( + | { kind: "key"; key: KeyResponse } + | { kind: "filter" } + | { kind: "search-keys" } + | { kind: "page"; page: PaletteNavItem } + ); + +interface PaletteGroup { + label: string; + entries: PaletteEntry[]; + emptyMessage?: string; + isLoading?: boolean; +} + +const shortcutChipClass = + "rounded-sm border border-border border-b-2 bg-muted px-[3px] font-mono text-muted-foreground"; + +export function CommandPalette({ open, setOpen }: { open: boolean; setOpen: (open: boolean) => void }) { + if (!open) return null; + return ; +} + +function CommandPaletteDialog({ setOpen }: { setOpen: (open: boolean) => void }) { + const pathname = usePathname(); + const router = useRouter(); + const [query, setQuery] = useState(""); + const [activeSelection, setActiveSelection] = useState<{ + id: string; + query: string; + scope: PaletteScope; + } | null>(null); + const [scope, setScope] = useState(() => paletteScopeForRoute(routeSegmentForPathname(pathname))); + const [debouncedQuery] = useDebouncedValue(query, { wait: DEBOUNCE_WAIT_MS }); + const trimmedQuery = query.trim(); + const debouncedSearch = debouncedQuery.trim(); + const { data: teams } = useTeams({ enabled: false }); + const teamAliases = useMemo( + () => new Map((teams ?? []).map((team) => [team.team_id, team.team_alias] as const)), + [teams], + ); + const keyListOptions = { + search: debouncedSearch || undefined, + sortBy: "created_at", + sortOrder: "desc", + expand: "user", + }; + const keyResults = useKeys(1, 8, keyListOptions, { enabled: scope === "keys" }); + const isPendingQuery = trimmedQuery !== debouncedSearch; + const hasPlaceholderKeyResults = keyResults.isPlaceholderData; + const isInitialKeyLoad = keyResults.isFetching && !keyResults.data; + const isLoadingKeys = scope === "keys" && (isPendingQuery || hasPlaceholderKeyResults || isInitialKeyLoad); + const visibleGroups = useVisibleMenuGroups(); + const navItems = useMemo(() => flattenNavItems(visibleGroups), [visibleGroups]); + const navMatches = useMemo( + () => (scope === "global" || trimmedQuery ? matchNavItems(navItems, query) : []), + [navItems, query, scope, trimmedQuery], + ); + const groups = useMemo(() => { + const pages: PaletteEntry[] = navMatches.map((page) => ({ + id: `page:${page.route}`, + kind: "page", + title: page.label, + subtitle: page.section, + icon: page.icon ?? , + typeLabel: "Page", + page, + })); + + if (scope === "global") { + const searchKeysAction = !trimmedQuery || "search virtual keys".includes(trimmedQuery.toLowerCase()); + return [ + ...(searchKeysAction + ? [ + { + label: "Actions", + entries: [ + { + id: "action:search-keys", + kind: "search-keys" as const, + title: "Search virtual keys", + subtitle: "Find a key by alias or ID", + icon: , + typeLabel: "Action" as const, + }, + ], + }, + ] + : []), + ...(pages.length > 0 ? [{ label: "Pages", entries: pages }] : []), + ]; + } + + const keyEntries: PaletteEntry[] = (isLoadingKeys ? [] : keyResults.data?.keys ?? []).map((key) => { + const teamAlias = key.team_alias || (key.team_id ? teamAliases.get(key.team_id) : undefined); + const spend = typeof key.spend === "number" && Number.isFinite(key.spend) ? `$${key.spend.toFixed(2)}` : null; + const details = [teamAlias, spend].filter((detail): detail is string => Boolean(detail)); + const subtitle: ReactNode = ( + <> + {key.key_name} + {details.length > 0 && · {details.join(" · ")}} + + ); + return { + id: `key:${key.token}`, + kind: "key", + title: key.key_alias || "Unnamed key", + subtitle, + icon: , + typeLabel: "Key", + key, + }; + }); + const keysGroup: PaletteGroup = { + label: trimmedQuery ? "Keys" : "Recent keys", + entries: keyEntries, + emptyMessage: trimmedQuery ? `No keys match “${trimmedQuery}”` : "No recent keys", + isLoading: isLoadingKeys, + }; + + if (!trimmedQuery) return [keysGroup]; + + const filterEntry: PaletteEntry = { + id: "action:filter-keys", + kind: "filter", + title: `Filter the keys table for “${trimmedQuery}”`, + subtitle: "Open Virtual Keys with this search applied", + icon: , + typeLabel: "Action", + }; + + return [ + keysGroup, + { label: "Actions", entries: [filterEntry] }, + ...(pages.length > 0 ? [{ label: "Pages", entries: pages }] : []), + ]; + }, [isLoadingKeys, keyResults.data?.keys, navMatches, scope, teamAliases, trimmedQuery]); + const selectableEntries = useMemo(() => groups.flatMap((group) => group.entries), [groups]); + const activeEntryId = activeSelection?.query === query && activeSelection.scope === scope ? activeSelection.id : null; + const selectedIndex = selectableEntries.findIndex((entry) => entry.id === activeEntryId); + const activeIndex = selectableEntries.length === 0 ? -1 : Math.max(0, selectedIndex); + const activeOptionId = activeIndex >= 0 ? `command-palette-option-${activeIndex}` : undefined; + const hasKeyError = scope === "keys" && keyResults.isError && keyResults.data === undefined; + + useEffect(() => { + if (activeOptionId) document.getElementById(activeOptionId)?.scrollIntoView?.({ block: "nearest" }); + }, [activeOptionId]); + + const activate = (entry: PaletteEntry) => { + switch (entry.kind) { + case "key": + router.push(keyDetailHref(entry.key.token)); + setOpen(false); + return; + case "filter": + router.push(`${uiHref("api-keys")}?key_search=${encodeURIComponent(trimmedQuery)}`); + setOpen(false); + return; + case "search-keys": + setQuery(""); + setScope("keys"); + return; + case "page": + router.push(uiHref(entry.page.route)); + setOpen(false); + return; + } + }; + + const handleInputKeyDown = (event: ReactKeyboardEvent) => { + if (event.key === "Backspace" && scope === "keys" && query === "") { + event.preventDefault(); + setScope("global"); + return; + } + + if (event.key === "Enter") { + event.preventDefault(); + const activeEntry = selectableEntries[activeIndex]; + if (activeEntry) activate(activeEntry); + return; + } + + if (event.key !== "ArrowDown" && event.key !== "ArrowUp") return; + event.preventDefault(); + if (selectableEntries.length === 0) return; + const currentIndex = activeIndex < 0 ? 0 : activeIndex; + const nextIndex = + event.key === "ArrowDown" + ? (currentIndex + 1) % selectableEntries.length + : (currentIndex - 1 + selectableEntries.length) % selectableEntries.length; + setActiveSelection({ id: selectableEntries[nextIndex].id, query, scope }); + }; + + const renderGroupContents = (group: PaletteGroup, firstOptionIndex: number): ReactNode => { + if (group.isLoading) { + return ( +
+ + Searching… +
+ ); + } + if (group.entries.length === 0) + return
{group.emptyMessage}
; + return group.entries.map((entry, entryIndex) => { + const currentIndex = firstOptionIndex + entryIndex; + return ( + activate(entry)} + onHover={() => setActiveSelection({ id: entry.id, query, scope })} + /> + ); + }); + }; + + return ( + + + Command palette +
+ + {scope === "keys" && ( + + + Virtual Keys + + )} + { + setQuery(event.currentTarget.value); + }} + onKeyDown={handleInputKeyDown} + /> +
+
+ {hasKeyError && ( +
+ Unable to search keys +
+ )} + {!hasKeyError && scope === "global" && trimmedQuery && groups.length === 0 && ( +
No results for “{trimmedQuery}”
+ )} + {!hasKeyError && + groups.map((group, groupIndex) => { + const firstOptionIndex = groups + .slice(0, groupIndex) + .reduce((count, previousGroup) => count + previousGroup.entries.length, 0); + return ( +
+
+ {group.label} +
+ {renderGroupContents(group, firstOptionIndex)} +
+ ); + })} +
+
+
+ LiteLLM + {scope === "keys" ? "Virtual Keys" : "All pages"} +
+
+ + ↵ Open + + + ↑↓ Navigate + + + esc Close + +
+
+
+
+ ); +} diff --git a/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteHint.tsx b/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteHint.tsx new file mode 100644 index 00000000000..0a4c3fcda18 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteHint.tsx @@ -0,0 +1,44 @@ +"use client"; + +import { Search, X } from "lucide-react"; +import { z } from "zod"; +import { storageKey } from "@/lib/storage"; +import type { PaletteScope } from "./utils"; +import { useCommandPalette } from "./CommandPaletteProvider"; +import { useClientMounted, useShortcutLabel } from "./useShortcutLabel"; + +export const COMMAND_PALETTE_HINT_KEY = storageKey("local", "litellmCommandPaletteHintSeen", z.boolean(), false); + +interface CommandPaletteHintProps { + scope: PaletteScope; + onDismiss: () => void; +} + +export function CommandPaletteHint({ scope, onDismiss }: CommandPaletteHintProps) { + const { setOpen } = useCommandPalette(); + const shortcut = useShortcutLabel(); + const mounted = useClientMounted(); + + if (!mounted) return null; + + return ( +
+ + +
+ ); +} diff --git a/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteProvider.tsx b/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteProvider.tsx new file mode 100644 index 00000000000..b2ec2625c45 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteProvider.tsx @@ -0,0 +1,65 @@ +"use client"; + +import { createContext, useCallback, useContext, useEffect, useMemo, useState, type ReactNode } from "react"; +import { usePathname } from "next/navigation"; +import { useStoredValue } from "@/lib/storage"; +import { routeSegmentForPathname } from "@/utils/uiHref"; +import { CommandPalette } from "./CommandPalette"; +import { CommandPaletteHint, COMMAND_PALETTE_HINT_KEY } from "./CommandPaletteHint"; +import { paletteScopeForRoute } from "./utils"; + +interface CommandPaletteContextValue { + open: boolean; + setOpen: (open: boolean) => void; + toggle: () => void; +} + +const CommandPaletteContext = createContext(null); + +export function useCommandPalette(): CommandPaletteContextValue { + const context = useContext(CommandPaletteContext); + if (!context) throw new Error("useCommandPalette must be used within CommandPaletteProvider"); + return context; +} + +export function CommandPaletteProvider({ children }: { children: ReactNode }) { + const [open, setOpenState] = useState(false); + const [hintSeen, setHintSeen] = useStoredValue(COMMAND_PALETTE_HINT_KEY); + const [hintDismissedInMemory, setHintDismissedInMemory] = useState(false); + const pathname = usePathname(); + const scope = paletteScopeForRoute(routeSegmentForPathname(pathname)); + const markHintSeen = useCallback(() => { + setHintDismissedInMemory(true); + setHintSeen(true); + }, [setHintSeen]); + const setOpen = useCallback( + (nextOpen: boolean) => { + if (nextOpen) markHintSeen(); + setOpenState(nextOpen); + }, + [markHintSeen], + ); + const toggle = useCallback(() => setOpen(!open), [open, setOpen]); + const value = useMemo(() => ({ open, setOpen, toggle }), [open, setOpen, toggle]); + const shouldShowHint = !hintSeen && !hintDismissedInMemory && !open; + + useEffect(() => { + const handleKeyDown = (event: KeyboardEvent) => { + if ((event.metaKey || event.ctrlKey) && event.key.toLowerCase() === "k" && !event.shiftKey && !event.altKey) { + event.preventDefault(); + toggle(); + } + }; + + document.addEventListener("keydown", handleKeyDown); + return () => document.removeEventListener("keydown", handleKeyDown); + }, [toggle]); + + return ( + + {children} + + {shouldShowHint && } + + ); +} diff --git a/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteResultRow.tsx b/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteResultRow.tsx new file mode 100644 index 00000000000..b4d8e2b3e74 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteResultRow.tsx @@ -0,0 +1,50 @@ +"use client"; + +import type { ReactNode } from "react"; +import { cn } from "@/lib/cva.config"; + +export interface CommandPaletteRowData { + id: string; + title: string; + subtitle?: ReactNode; + icon: ReactNode; + typeLabel: "Key" | "Page" | "Action"; +} + +export function CommandPaletteResultRow({ + item, + optionId, + active, + onActivate, + onHover, +}: { + item: CommandPaletteRowData; + optionId: string; + active: boolean; + onActivate: () => void; + onHover: () => void; +}) { + return ( + + ); +} diff --git a/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteTrigger.tsx b/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteTrigger.tsx new file mode 100644 index 00000000000..bd594800736 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CommandPalette/CommandPaletteTrigger.tsx @@ -0,0 +1,50 @@ +"use client"; + +import { Search } from "lucide-react"; +import { usePathname } from "next/navigation"; +import { Button } from "@/components/ui/button"; +import { routeSegmentForPathname } from "@/utils/uiHref"; +import { paletteScopeForRoute } from "./utils"; +import { useCommandPalette } from "./CommandPaletteProvider"; +import { useShortcutLabel } from "./useShortcutLabel"; + +export function CommandPaletteTrigger({ mobile = false }: { mobile?: boolean }) { + const pathname = usePathname(); + const { open, setOpen } = useCommandPalette(); + const shortcut = useShortcutLabel(); + const isKeysRoute = paletteScopeForRoute(routeSegmentForPathname(pathname)) === "keys"; + const label = isKeysRoute ? "Search keys…" : "Search or jump to…"; + + if (mobile) { + return ( + + ); + } + + return ( + + ); +} diff --git a/ui/litellm-dashboard/src/components/CommandPalette/useShortcutLabel.ts b/ui/litellm-dashboard/src/components/CommandPalette/useShortcutLabel.ts new file mode 100644 index 00000000000..78e1865c90e --- /dev/null +++ b/ui/litellm-dashboard/src/components/CommandPalette/useShortcutLabel.ts @@ -0,0 +1,20 @@ +"use client"; + +import { useSyncExternalStore } from "react"; +import { shortcutLabel } from "./utils"; + +const subscribeToShortcut = () => () => {}; +const getClientShortcut = () => + typeof navigator === "undefined" ? "Ctrl K" : shortcutLabel(navigator.platform, navigator.userAgent); +const getServerShortcut = () => "Ctrl K"; +const subscribeToMount = () => () => {}; +const getClientMounted = () => true; +const getServerMounted = () => false; + +export function useShortcutLabel(): string { + return useSyncExternalStore(subscribeToShortcut, getClientShortcut, getServerShortcut); +} + +export function useClientMounted(): boolean { + return useSyncExternalStore(subscribeToMount, getClientMounted, getServerMounted); +} diff --git a/ui/litellm-dashboard/src/components/CommandPalette/useVisibleMenuGroups.ts b/ui/litellm-dashboard/src/components/CommandPalette/useVisibleMenuGroups.ts new file mode 100644 index 00000000000..8c3b9382934 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CommandPalette/useVisibleMenuGroups.ts @@ -0,0 +1,34 @@ +"use client"; + +import { useMemo } from "react"; +import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; +import useIsOrgAdmin from "@/app/(dashboard)/hooks/useIsOrgAdmin"; +import { isUserTeamAdminForAnyTeam } from "@/utils/roles"; +import { menuGroups, visibleMenuGroups, type MenuVisibilityContext } from "@/components/leftnav"; + +export const useVisibleMenuGroups = () => { + const { userId, userRole, isViewOnly } = useAuthorized(); + const isOrgAdmin = useIsOrgAdmin(); + const { data: teams } = useTeams({ enabled: false }); + const { data: settings } = useUISettings(); + const values = settings?.values; + const isTeamAdmin = useMemo(() => isUserTeamAdminForAnyTeam(teams ?? null, userId ?? ""), [teams, userId]); + + return useMemo(() => { + const context: MenuVisibilityContext = { + userRole, + isViewOnly, + isOrgAdmin, + isTeamAdmin, + enabledPagesInternalUsers: values?.enabled_ui_pages_internal_users ?? null, + enableProjectsUI: Boolean(values?.enable_projects_ui), + disableAgentsForInternalUsers: Boolean(values?.disable_agents_for_internal_users), + allowAgentsForTeamAdmins: Boolean(values?.allow_agents_for_team_admins), + disableVectorStoresForInternalUsers: Boolean(values?.disable_vector_stores_for_internal_users), + allowVectorStoresForTeamAdmins: Boolean(values?.allow_vector_stores_for_team_admins), + }; + return visibleMenuGroups(menuGroups, context); + }, [isOrgAdmin, isTeamAdmin, isViewOnly, userRole, values]); +}; diff --git a/ui/litellm-dashboard/src/components/CommandPalette/utils.test.ts b/ui/litellm-dashboard/src/components/CommandPalette/utils.test.ts new file mode 100644 index 00000000000..5bb61756f80 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CommandPalette/utils.test.ts @@ -0,0 +1,91 @@ +import { describe, expect, it } from "vitest"; +import { visibleMenuGroups, type MenuGroup, menuGroups } from "@/components/leftnav"; +import { flattenNavItems, matchNavItems, paletteScopeForRoute, shortcutLabel, type PaletteNavItem } from "./utils"; + +describe("paletteScopeForRoute", () => { + it.each([ + ["api-keys", "keys"], + ["", "keys"], + ["teams", "global"], + ] as const)("uses the expected palette scope for %s", (routeSegment, expected) => { + expect(paletteScopeForRoute(routeSegment)).toBe(expected); + }); +}); + +describe("flattenNavItems", () => { + it("flattens visible leaves and skips external links", () => { + const groups: MenuGroup[] = [ + { + groupLabel: "Tools", + items: [ + { key: "search-tools", page: "search-tools", label: "Search Tools" }, + { key: "docs", page: "docs", label: "Docs", external_url: "https://example.com" }, + { key: "tools", page: "tools", label: "Tools", children: [{ key: "logs", page: "logs", label: "Logs" }] }, + ], + }, + ]; + + expect(flattenNavItems(groups).map(({ label }) => label)).toEqual(["Search Tools", "Logs"]); + }); +}); + +describe("visibleMenuGroups", () => { + const context = { + userRole: "Internal User", + isViewOnly: false, + isOrgAdmin: false, + isTeamAdmin: false, + }; + + it("applies the internal-user page allowlist", () => { + const visible = visibleMenuGroups(menuGroups, { + ...context, + enabledPagesInternalUsers: ["api-keys"], + }); + + expect(flattenNavItems(visible).map(({ route }) => route)).toEqual(["api-keys"]); + }); + + it("hides Projects when its UI setting is disabled", () => { + const projectAdminContext = { ...context, userRole: "Admin", isTeamAdmin: true }; + const visibleWithFlag = visibleMenuGroups(menuGroups, { ...projectAdminContext, enableProjectsUI: true }); + const visibleWithoutFlag = visibleMenuGroups(menuGroups, { ...projectAdminContext, enableProjectsUI: false }); + + expect(flattenNavItems(visibleWithFlag).some(({ route }) => route === "projects")).toBe(true); + expect(flattenNavItems(visibleWithoutFlag).some(({ route }) => route === "projects")).toBe(false); + }); + + it("keeps admin pages visible despite the internal-user page allowlist", () => { + const visible = visibleMenuGroups(menuGroups, { + ...context, + userRole: "Admin", + enabledPagesInternalUsers: ["api-keys"], + }); + + expect(flattenNavItems(visible).some(({ route }) => route === "users")).toBe(true); + expect(flattenNavItems(visible).some(({ route }) => route === "organizations")).toBe(true); + }); +}); + +describe("matchNavItems", () => { + const items: PaletteNavItem[] = [ + { key: "section", label: "Teams", section: "Logs", route: "teams" }, + { key: "substring", label: "Audit Logs", section: "Access Control", route: "audit-logs" }, + { key: "prefix", label: "Logging", section: "Observability", route: "logging" }, + ]; + + it("ranks label prefixes above label substrings and section matches", () => { + expect(matchNavItems(items, "log").map(({ key }) => key)).toEqual(["prefix", "substring", "section"]); + }); + + it("returns every item in its original order for an empty query", () => { + expect(matchNavItems(items, " ")).toEqual(items); + }); +}); + +describe("shortcutLabel", () => { + it("uses the command key on Mac and Ctrl elsewhere", () => { + expect(shortcutLabel("MacIntel", "Mozilla/5.0")).toBe("⌘K"); + expect(shortcutLabel("Linux x86_64", "Mozilla/5.0 (Windows NT 10.0)")).toBe("Ctrl K"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CommandPalette/utils.ts b/ui/litellm-dashboard/src/components/CommandPalette/utils.ts new file mode 100644 index 00000000000..ad02e155b0e --- /dev/null +++ b/ui/litellm-dashboard/src/components/CommandPalette/utils.ts @@ -0,0 +1,59 @@ +import type { ReactNode } from "react"; +import { labelText, routeOf, sectionText, type MenuGroup } from "@/components/leftnav"; + +export type PaletteScope = "keys" | "global"; + +export interface PaletteNavItem { + key: string; + label: string; + section: string; + route: string; + icon?: ReactNode; +} + +export const paletteScopeForRoute = (routeSegment: string): PaletteScope => + routeSegment === "" || routeSegment === "api-keys" ? "keys" : "global"; + +const matchRank = (item: PaletteNavItem, query: string): number => { + const label = item.label.toLowerCase(); + if (label.startsWith(query)) return 0; + if (label.includes(query)) return 1; + if (item.section.toLowerCase().includes(query)) return 2; + return -1; +}; + +export const flattenNavItems = (groups: readonly MenuGroup[]): PaletteNavItem[] => + groups.flatMap((group) => { + const flattenItems = (items: MenuGroup["items"]): PaletteNavItem[] => + items.flatMap((item) => { + if (item.external_url) return []; + if (item.children) return flattenItems(item.children); + return [ + { + key: item.key, + label: labelText(item), + section: sectionText(group.groupLabel), + route: routeOf(item), + icon: item.icon, + }, + ]; + }); + + return flattenItems(group.items); + }); + +export const matchNavItems = (items: readonly PaletteNavItem[], query: string): PaletteNavItem[] => { + const normalizedQuery = query.trim().toLowerCase(); + if (!normalizedQuery) return [...items]; + + return items + .map((item, index) => { + return { item, index, rank: matchRank(item, normalizedQuery) }; + }) + .filter(({ rank }) => rank >= 0) + .sort((left, right) => left.rank - right.rank || left.index - right.index) + .map(({ item }) => item); +}; + +export const shortcutLabel = (platform: string, userAgent: string): string => + [platform, userAgent].some((value) => value.toLowerCase().includes("mac")) ? "⌘K" : "Ctrl K"; diff --git a/ui/litellm-dashboard/src/components/DashboardHeader.integration.test.tsx b/ui/litellm-dashboard/src/components/DashboardHeader.integration.test.tsx index 756997b7ede..cb58670c20a 100644 --- a/ui/litellm-dashboard/src/components/DashboardHeader.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/DashboardHeader.integration.test.tsx @@ -3,6 +3,7 @@ import { render, screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { PluginModeProvider } from "@/contexts/PluginModeContext"; +import { CommandPaletteProvider } from "@/components/CommandPalette/CommandPaletteProvider"; import { DashboardHeader } from "./DashboardHeader"; vi.mock("next/navigation", () => ({ usePathname: () => "/ui/logs" })); @@ -57,7 +58,9 @@ async function openTools() { render( - + + + , ); diff --git a/ui/litellm-dashboard/src/components/DashboardHeader.test.tsx b/ui/litellm-dashboard/src/components/DashboardHeader.test.tsx index 63f7b7aa136..822cec9deb3 100644 --- a/ui/litellm-dashboard/src/components/DashboardHeader.test.tsx +++ b/ui/litellm-dashboard/src/components/DashboardHeader.test.tsx @@ -1,7 +1,8 @@ -import { afterEach, describe, expect, it, vi } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { act, fireEvent, render, screen, within } from "@testing-library/react"; import { DashboardHeader } from "./DashboardHeader"; import { NAV_PRODUCT_LINK_CLASS } from "@/components/Navbar/navProductLinkClass"; +import { CommandPaletteProvider } from "@/components/CommandPalette/CommandPaletteProvider"; const { mockUsePluginMode, mockUseUISettings, state } = vi.hoisted(() => { const state = { @@ -31,31 +32,43 @@ vi.mock("@/components/Navbar/NotificationsBell/NotificationsBell", () => ({ Noti vi.mock("@/components/Navbar/WorkerDropdown/WorkerDropdown", () => ({ default: () => null })); vi.mock("@/components/liteadmin/LiteAdmin", () => ({ default: () => })); +const renderDashboardHeader = () => + render( + + + , + ); + describe("DashboardHeader breadcrumb", () => { + beforeEach(() => { + localStorage.clear(); + }); + afterEach(() => { state.plugins = []; state.enableChatUI = false; state.pathname = "/ui/logs"; state.isDesktop = false; + localStorage.clear(); }); it("titles the breadcrumb from the current route, not from a sidebar page id", () => { state.pathname = "/ui/models-and-endpoints"; - render(); + renderDashboardHeader(); expect(screen.getByText("Models + Endpoints")).toBeInTheDocument(); }); it("titles the dashboard root as Virtual Keys", () => { state.pathname = "/ui/"; - render(); + renderDashboardHeader(); expect(screen.getByText("Virtual Keys")).toBeInTheDocument(); }); it("roots the breadcrumb in the AI Gateway selector (with a Chat option) and drops the static section crumb when the selector is available", async () => { state.enableChatUI = true; - render(); + renderDashboardHeader(); expect(screen.getByText("Logs")).toBeInTheDocument(); expect(screen.queryByText("Observability")).not.toBeInTheDocument(); @@ -68,7 +81,7 @@ describe("DashboardHeader breadcrumb", () => { }); it("keeps the AI Gateway selector at the root even when there is nothing to switch to (discovery)", () => { - render(); + renderDashboardHeader(); expect(screen.getByRole("button", { name: /AI Gateway/i })).toBeInTheDocument(); expect(screen.getByText("Logs")).toBeInTheDocument(); @@ -76,7 +89,7 @@ describe("DashboardHeader breadcrumb", () => { }); it("styles Docs with the shared product-link class instead of a muted toolbar button", () => { - render(); + renderDashboardHeader(); const docs = screen.getByRole("link", { name: "Docs" }); for (const cls of NAV_PRODUCT_LINK_CLASS.trim().split(/\s+/)) { @@ -86,7 +99,7 @@ describe("DashboardHeader breadcrumb", () => { }); it("renders the tools divider centered rather than stretched to the top of the row", () => { - const { container } = render(); + const { container } = renderDashboardHeader(); const separators = container.querySelectorAll('[data-slot="separator"][data-orientation="vertical"]'); expect(separators).toHaveLength(1); @@ -95,7 +108,7 @@ describe("DashboardHeader breadcrumb", () => { }); it("places LiteAdmin in the header tools ahead of Docs", () => { - render(); + renderDashboardHeader(); const liteAdmin = within(screen.getByRole("banner")).getByRole("button", { name: "LiteAdmin" }); expect(liteAdmin.compareDocumentPosition(screen.getByRole("link", { name: "Docs" }))).toBe( @@ -104,7 +117,7 @@ describe("DashboardHeader breadcrumb", () => { }); it("keeps the gateway selector and tools available from the compact header menu", async () => { - render(); + renderDashboardHeader(); fireEvent.click(screen.getByRole("button", { name: "More options" })); const tools = await screen.findByRole("dialog", { name: "Gateway tools" }); expect(within(tools).getByRole("button", { name: "AI Gateway" })).toBeInTheDocument(); @@ -113,16 +126,24 @@ describe("DashboardHeader breadcrumb", () => { }); it("closes mobile tools when switching to desktop and keeps them closed when returning", async () => { - const { rerender } = render(); + const { rerender } = renderDashboardHeader(); fireEvent.click(screen.getByRole("button", { name: "More options" })); expect(await screen.findByRole("dialog", { name: "Gateway tools" })).toBeInTheDocument(); state.isDesktop = true; - rerender(); + rerender( + + + , + ); expect(screen.queryByRole("dialog", { name: "Gateway tools" })).not.toBeInTheDocument(); state.isDesktop = false; - rerender(); + rerender( + + + , + ); expect(screen.queryByRole("dialog", { name: "Gateway tools" })).not.toBeInTheDocument(); }); }); diff --git a/ui/litellm-dashboard/src/components/DashboardHeader.tsx b/ui/litellm-dashboard/src/components/DashboardHeader.tsx index 43440b0c92d..7580f8b69c0 100644 --- a/ui/litellm-dashboard/src/components/DashboardHeader.tsx +++ b/ui/litellm-dashboard/src/components/DashboardHeader.tsx @@ -27,6 +27,7 @@ import { useMediaQuery } from "usehooks-ts"; import { Ellipsis } from "lucide-react"; import { Button } from "@/components/ui/button"; import { Popover, PopoverContent, PopoverTitle, PopoverTrigger } from "@/components/ui/popover"; +import { CommandPaletteTrigger } from "@/components/CommandPalette/CommandPaletteTrigger"; // Top bar for the dashboard shell. Sits only over the content column (the brand // lives in the sidebar header); mirrors the design's breadcrumb-left / tools-right layout. @@ -62,6 +63,8 @@ export function DashboardHeader({ navigationTrigger }: { navigationTrigger?: Rea + +
{showWorkerSwitch && ( <> @@ -78,6 +81,7 @@ export function DashboardHeader({ navigationTrigger }: { navigationTrigger?: Rea
+ }> diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx index 12efc5beb7c..220d46647f6 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.integration.test.tsx @@ -164,7 +164,6 @@ const createTestProps = (userRole = "proxy_admin", userId = "user-1", isTeamAdmi mountedValues: () => projectMountedValues(registry, form.getValues), handleOk: vi.fn().mockResolvedValue(true), setSelectedProvider: vi.fn(), - setProviderModelsFn: vi.fn(), getPlaceholder: vi.fn((provider: string) => `Enter ${provider} model name`), setShowAdvancedSettings: vi.fn(), selectedProvider: Providers.OpenAI, diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx index a416bf87602..b5b6c439dc2 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx @@ -52,7 +52,6 @@ interface AddModelFormProps { selectedProvider: string | null; setSelectedProvider: (provider: string | null) => void; providerModels: string[]; - setProviderModelsFn: (provider: string | null) => void; getPlaceholder: (provider: string) => string; showAdvancedSettings: boolean; setShowAdvancedSettings: (show: boolean) => void; @@ -76,7 +75,6 @@ const AddModelForm: React.FC = ({ selectedProvider, setSelectedProvider, providerModels, - setProviderModelsFn, getPlaceholder, showAdvancedSettings, setShowAdvancedSettings, @@ -166,7 +164,6 @@ const AddModelForm: React.FC = ({ const applyProviderSelection = (provider: string | null) => { setSelectedProvider(provider); - setProviderModelsFn(provider); form.setValue("model", []); form.setValue("model_name", undefined); }; diff --git a/ui/litellm-dashboard/src/components/add_model/litellm_model_name.tsx b/ui/litellm-dashboard/src/components/add_model/litellm_model_name.tsx index 8cf4b077f39..529c0e2a63e 100644 --- a/ui/litellm-dashboard/src/components/add_model/litellm_model_name.tsx +++ b/ui/litellm-dashboard/src/components/add_model/litellm_model_name.tsx @@ -131,12 +131,12 @@ const LiteLLMModelNameField: React.FC = ({ } }} /> - ) : providerModels.length > 0 ? ( + ) : providerModels.length > 0 && !(typeof control.value === "string" && control.value !== "") ? ( { control.onChange(value); handleModelChange(value); diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index 4cba12c9f7e..256fe578912 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -103,7 +103,7 @@ interface SidebarProps { allowVectorStoresForTeamAdmins?: boolean; } -interface MenuItem { +export interface MenuItem { key: string; page: string; route?: string; @@ -114,12 +114,92 @@ interface MenuItem { external_url?: string; } -interface MenuGroup { +export interface MenuGroup { groupLabel: string; items: MenuItem[]; roles?: string[]; } +export interface MenuVisibilityContext { + userRole: string; + isViewOnly: boolean; + isOrgAdmin: boolean; + isTeamAdmin: boolean; + enabledPagesInternalUsers?: string[] | null; + enableProjectsUI?: boolean; + disableAgentsForInternalUsers?: boolean; + allowAgentsForTeamAdmins?: boolean; + disableVectorStoresForInternalUsers?: boolean; + allowVectorStoresForTeamAdmins?: boolean; +} + +const adminPageVisibility = (item: MenuItem, context: MenuVisibilityContext, isAdmin: boolean): boolean | null => { + if (item.key !== "organizations" && item.key !== "users") return null; + const hasRoleAccess = !item.roles || item.roles.includes(context.userRole) || context.isOrgAdmin; + if (!hasRoleAccess) return false; + if (!isAdmin && context.enabledPagesInternalUsers != null) { + return context.enabledPagesInternalUsers.includes(item.page); + } + return true; +}; + +const projectPageIsVisible = (item: MenuItem, context: MenuVisibilityContext): boolean => { + if (item.key !== "projects") return true; + if (!context.enableProjectsUI) return false; + return canViewProjectsPage({ + userRole: context.userRole, + isOrgAdmin: context.isOrgAdmin, + isTeamAdmin: context.isTeamAdmin, + }); +}; + +const pageIsDisabledForInternalUsers = ( + item: MenuItem, + pageKey: "agents" | "vector-stores", + isAdmin: boolean, + context: MenuVisibilityContext, +): boolean => { + const isDisabled = + pageKey === "agents" ? context.disableAgentsForInternalUsers : context.disableVectorStoresForInternalUsers; + const allowedForTeamAdmins = + pageKey === "agents" ? context.allowAgentsForTeamAdmins : context.allowVectorStoresForTeamAdmins; + if (item.key !== pageKey || isAdmin || !isDisabled) return false; + return !(allowedForTeamAdmins && context.isTeamAdmin); +}; + +const pageIsEnabledForInternalUser = (item: MenuItem, context: MenuVisibilityContext): boolean => { + const enabledPages = context.enabledPagesInternalUsers; + if (enabledPages == null) return true; + if (item.children?.some((child) => enabledPages.includes(child.page))) return true; + return enabledPages.includes(item.page); +}; + +const menuItemIsVisible = (item: MenuItem, context: MenuVisibilityContext, isAdmin: boolean): boolean => { + if (item.children && item.children.length === 0) return false; + if (item.key === "llm-playground" && context.isViewOnly) return false; + const adminVisibility = adminPageVisibility(item, context, isAdmin); + if (adminVisibility !== null) return adminVisibility; + if (!projectPageIsVisible(item, context)) return false; + if (pageIsDisabledForInternalUsers(item, "agents", isAdmin, context)) return false; + if (pageIsDisabledForInternalUsers(item, "vector-stores", isAdmin, context)) return false; + if (item.roles && !item.roles.includes(context.userRole)) return false; + if (!isAdmin && context.enabledPagesInternalUsers != null) return pageIsEnabledForInternalUser(item, context); + return true; +}; + +export const visibleMenuGroups = (groups: readonly MenuGroup[], context: MenuVisibilityContext): MenuGroup[] => { + const isAdmin = isAdminRole(context.userRole); + const filterItems = (items: MenuItem[]): MenuItem[] => + items + .map((item) => ({ ...item, children: item.children ? filterItems(item.children) : undefined })) + .filter((item) => menuItemIsVisible(item, context, isAdmin)); + + return groups + .filter((group) => !group.roles || group.roles.includes(context.userRole)) + .map((group) => ({ groupLabel: group.groupLabel, items: filterItems(group.items) })) + .filter((group) => group.items.length > 0); +}; + // Menu groups organized by category - defined outside component for export. // Shape (key/page/label/roles/children) is consumed by page_utils.ts; only the // icons changed to lucide as part of the sidebar redesign. @@ -390,7 +470,7 @@ const menuGroups: MenuGroup[] = [ const HOME_ROUTE = "api-keys"; -const routeOf = (item: MenuItem): string => item.route ?? item.page; +export const routeOf = (item: MenuItem): string => item.route ?? item.page; const routeForPathname = (pathname: string): string => routeSegmentForPathname(pathname) || HOME_ROUTE; @@ -422,6 +502,8 @@ const SECTION_DISPLAY: Record = { SETTINGS: "Settings", }; +export const sectionText = (groupLabel: string): string => SECTION_DISPLAY[groupLabel] ?? groupLabel; + const prettify = (key: string): string => key .split(/[-_]/) @@ -435,7 +517,7 @@ export const getBreadcrumb = (pathname: string): { section: string | null; title const route = routeForPathname(pathname); for (const group of menuGroups) { for (const item of group.items) { - const section = SECTION_DISPLAY[group.groupLabel] ?? group.groupLabel; + const section = sectionText(group.groupLabel); if (routeOf(item) === route) return { section, title: labelText(item) }; const child = item.children?.find((c) => routeOf(c) === route); if (child) return { section, title: labelText(child) }; @@ -486,56 +568,32 @@ const Sidebar_: React.FC = ({ const isTeamAdmin = useMemo(() => isUserTeamAdminForAnyTeam(teams ?? null, userId ?? ""), [teams, userId]); - const filterItemsByRole = (items: MenuItem[]): MenuItem[] => { - const isAdmin = isAdminRole(userRole); - return items - .map((item) => ({ ...item, children: item.children ? filterItemsByRole(item.children) : undefined })) - .filter((item) => { - // A parent whose children were all filtered out renders as a leaf link - // to its own page id, which is not a real route. Drop it instead. - if (item.children && item.children.length === 0) return false; - if (item.key === "llm-playground" && isViewOnly) return false; - if (item.key === "organizations" || item.key === "users") { - const hasRoleAccess = !item.roles || item.roles.includes(userRole) || isOrgAdmin; - if (!hasRoleAccess) return false; - if (!isAdmin && enabledPagesInternalUsers != null) return enabledPagesInternalUsers.includes(item.page); - return true; - } - if ( - item.key === "projects" && - !(enableProjectsUI && canViewProjectsPage({ userRole, isOrgAdmin, isTeamAdmin })) - ) - return false; - if ( - !isAdmin && - item.key === "agents" && - disableAgentsForInternalUsers && - !(allowAgentsForTeamAdmins && isTeamAdmin) - ) - return false; - if ( - !isAdmin && - item.key === "vector-stores" && - disableVectorStoresForInternalUsers && - !(allowVectorStoresForTeamAdmins && isTeamAdmin) - ) - return false; - if (item.roles && !item.roles.includes(userRole)) return false; - if (!isAdmin && enabledPagesInternalUsers != null) { - if (item.children && item.children.length > 0) { - const hasVisibleChildren = item.children.some((child) => enabledPagesInternalUsers.includes(child.page)); - if (hasVisibleChildren) return true; - } - return enabledPagesInternalUsers.includes(item.page); - } - return true; - }); - }; - - const visibleGroups = menuGroups - .filter((group) => !group.roles || group.roles.includes(userRole)) - .map((group) => ({ groupLabel: group.groupLabel, items: filterItemsByRole(group.items) })) - .filter((group) => group.items.length > 0); + const visibleGroups = useMemo(() => { + const context: MenuVisibilityContext = { + userRole, + isViewOnly, + isOrgAdmin, + isTeamAdmin, + enabledPagesInternalUsers, + enableProjectsUI, + disableAgentsForInternalUsers, + allowAgentsForTeamAdmins, + disableVectorStoresForInternalUsers, + allowVectorStoresForTeamAdmins, + }; + return visibleMenuGroups(menuGroups, context); + }, [ + allowAgentsForTeamAdmins, + allowVectorStoresForTeamAdmins, + disableAgentsForInternalUsers, + disableVectorStoresForInternalUsers, + enableProjectsUI, + enabledPagesInternalUsers, + isOrgAdmin, + isTeamAdmin, + isViewOnly, + userRole, + ]); const toggleGroup = (key: string) => { if (collapsed) { diff --git a/ui/litellm-dashboard/src/components/lens/traces/detail/feedback/FeedbackPanel.tsx b/ui/litellm-dashboard/src/components/lens/traces/detail/feedback/FeedbackPanel.tsx index f06591acce8..74ea709a731 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/detail/feedback/FeedbackPanel.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/feedback/FeedbackPanel.tsx @@ -46,7 +46,6 @@ function Entry({ entry }: { entry: Feedback }) { ); } -/** End-user feedback on this run, shown first so a developer reads what the user said before the steps. */ export function FeedbackPanel({ summary, accessToken }: FeedbackPanelProps) { const api = useTracesApi(accessToken); const traceRef = summary.trace_ref ?? ""; diff --git a/ui/litellm-dashboard/src/components/lens/traces/detail/feedback/feedback.ts b/ui/litellm-dashboard/src/components/lens/traces/detail/feedback/feedback.ts index 6a572997657..01a26e788a7 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/detail/feedback/feedback.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/detail/feedback/feedback.ts @@ -6,7 +6,6 @@ export interface FeedbackView { readonly lowest: number; } -/** End-user feedback on one run, newest first, or null when nobody has rated it. */ export function feedbackView(feedback: readonly Feedback[]): FeedbackView | null { if (feedback.length === 0) return null; const scores = feedback.map((entry) => entry.score); diff --git a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.integration.test.tsx index f0ba76367b3..5a32aa4e610 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.integration.test.tsx @@ -81,7 +81,6 @@ const unrated = (traces: TraceKey[]): TraceFeedbackSummary[] => lowest: null, })); -/** Routes the shared POST mock: findings get `findings`, the feedback summary gets `feedback`. */ const stubPost = ( findings: (traces: TraceKey[]) => Promise, feedback: (traces: TraceKey[]) => Promise = async (traces) => unrated(traces), diff --git a/ui/litellm-dashboard/src/components/lens/traces/list/useTraceFeedback.ts b/ui/litellm-dashboard/src/components/lens/traces/list/useTraceFeedback.ts index b35542a0943..fa091c784f0 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/list/useTraceFeedback.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/list/useTraceFeedback.ts @@ -13,7 +13,6 @@ export type TraceFeedbackState = export const traceFeedbackKey = (accessToken: string) => ["traceFeedback", accessToken] as const; -/** A run some end user scored at or below the low-score threshold. */ export const isLowFeedback = (state: TraceFeedbackState | undefined): boolean => state?.status === "ready" && state.summary.count > 0 && (state.summary.lowest ?? Infinity) <= LOW_SCORE; diff --git a/ui/litellm-dashboard/src/components/networking.test.ts b/ui/litellm-dashboard/src/components/networking.test.ts index 62670467827..0fdf7a4799d 100644 --- a/ui/litellm-dashboard/src/components/networking.test.ts +++ b/ui/litellm-dashboard/src/components/networking.test.ts @@ -871,3 +871,20 @@ describe("schema-bound dashboard responses", () => { expect(result.users[0]).toEqual(user); }); }); + +describe("modelCostMap", () => { + afterEach(() => { + vi.unstubAllGlobals(); + }); + + it("requests the catalog-only map when catalogOnly is set and the full map otherwise", async () => { + const mockFetch = vi.fn().mockImplementation(async () => new Response(JSON.stringify({}))); + vi.stubGlobal("fetch", mockFetch); + + await Networking.modelCostMap(true); + await Networking.modelCostMap(); + + expect(mockFetch.mock.calls[0][0]).toMatch(/\/public\/litellm_model_cost_map\?catalog_only=true$/); + expect(mockFetch.mock.calls[1][0]).toMatch(/\/public\/litellm_model_cost_map$/); + }); +}); diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index bf35e22f228..a874f58f868 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -553,9 +553,10 @@ export const getOpenAPISchema = async () => { return jsonData; }; -export const modelCostMap = async () => { +export const modelCostMap = async (catalogOnly = false) => { try { - const url = proxyBaseUrl ? `${proxyBaseUrl}/public/litellm_model_cost_map` : `/public/litellm_model_cost_map`; + const path = catalogOnly ? "/public/litellm_model_cost_map?catalog_only=true" : "/public/litellm_model_cost_map"; + const url = proxyBaseUrl ? `${proxyBaseUrl}${path}` : path; const response = await fetch(url, { method: "GET", headers: { diff --git a/ui/litellm-dashboard/src/components/policies/types.ts b/ui/litellm-dashboard/src/components/policies/types.ts index 430864f93df..2830a22950c 100644 --- a/ui/litellm-dashboard/src/components/policies/types.ts +++ b/ui/litellm-dashboard/src/components/policies/types.ts @@ -102,7 +102,7 @@ export interface PolicyAttachmentListResponse { export interface PipelineStepResult { guardrail_name: string; - outcome: "pass" | "fail" | "error"; + outcome: "pass" | "fail" | "error" | "skip"; action_taken: string; modified_data: Record | null; error_detail: string | null; diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx index 7cfdaf3275d..cecd6ce01bf 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx @@ -214,6 +214,26 @@ describe("provider_info_helpers", () => { const { logo } = getProviderLogoAndName("tencent"); expect(logo).toContain("tencent"); }); + + it("should resolve the typesafe slug and TypeSafe enum key to the TypeSafe name and bundled logo", () => { + const fromSlug = getProviderLogoAndName("typesafe"); + expect(fromSlug.displayName).toBe(Providers.TypeSafe); + expect(fromSlug.logo).toContain("typesafe"); + + const fromEnumKey = getProviderLogoAndName("TypeSafe"); + expect(fromEnumKey.displayName).toBe(Providers.TypeSafe); + expect(fromEnumKey.logo).toBe(fromSlug.logo); + }); + + it("should resolve the strands_decider slug and StrandsDecider enum key to the Strands Decider name and bundled logo", () => { + const fromSlug = getProviderLogoAndName("strands_decider"); + expect(fromSlug.displayName).toBe("Strands Decider"); + expect(fromSlug.logo).toContain("strands"); + + const fromEnumKey = getProviderLogoAndName("StrandsDecider"); + expect(fromEnumKey.displayName).toBe(Providers.StrandsDecider); + expect(fromEnumKey.logo).toBe(fromSlug.logo); + }); }); describe("getPlaceholder", () => { @@ -318,6 +338,11 @@ describe("provider_info_helpers", () => { expect(getPlaceholder(Providers.Tencent)).toBe("tencent/deepseek-v4-pro"); }); + it("should return decision model placeholders for the TypeSafe and StrandsDecider dropdown keys", () => { + expect(getPlaceholder("TypeSafe")).toBe("typesafe/jev-latest"); + expect(getPlaceholder("StrandsDecider")).toBe("strands_decider/strands-decider-2B-hobson-v19"); + }); + it("should return default gpt-3.5-turbo placeholder for unknown provider", () => { expect(getPlaceholder("UnknownProvider" as any)).toBe("gpt-3.5-turbo"); }); @@ -421,6 +446,30 @@ describe("provider_info_helpers", () => { expect(getProviderModels("Sail" as Providers, modelMap)).toEqual(["sail/openai/gpt-oss-120b"]); }); + it("should list only typesafe decision models for the 'TypeSafe' provider key, not the OpenRouter-hosted one", () => { + const modelMap = { + "typesafe/jev-latest": { litellm_provider: "typesafe", mode: "evaluation" }, + "typesafe/jev-preview": { litellm_provider: "typesafe", mode: "evaluation" }, + "openrouter/typesafe/jev-1.13": { litellm_provider: "openrouter", mode: "evaluation" }, + "strands_decider/strands-decider-2B-hobson-v19": { litellm_provider: "strands_decider", mode: "evaluation" }, + }; + expect(getProviderModels("TypeSafe" as Providers, modelMap)).toEqual([ + "typesafe/jev-latest", + "typesafe/jev-preview", + ]); + }); + + it("should list only strands_decider models for the 'StrandsDecider' provider key", () => { + const modelMap = { + "strands_decider/strands-decider-2B-hobson-v19": { litellm_provider: "strands_decider", mode: "evaluation" }, + "typesafe/jev-latest": { litellm_provider: "typesafe", mode: "evaluation" }, + "openrouter/typesafe/jev-1.13": { litellm_provider: "openrouter", mode: "evaluation" }, + }; + expect(getProviderModels("StrandsDecider" as Providers, modelMap)).toEqual([ + "strands_decider/strands-decider-2B-hobson-v19", + ]); + }); + it("should include bedrock converse but exclude standalone bedrock_mantle when called with 'Bedrock' provider key", () => { const modelMap = { "bedrock-base": { litellm_provider: "bedrock" }, diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx index 5ea693bea10..177d4be3247 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx @@ -55,9 +55,11 @@ import sapLogo from "../../public/assets/logos/sap.png"; import scxAiLogo from "../../public/assets/logos/scx_ai.svg"; import snowflakeLogo from "../../public/assets/logos/snowflake.svg"; import sonioxLogo from "../../public/assets/logos/soniox.svg"; +import strandsLogo from "../../public/assets/logos/strands.svg"; import tencentLogo from "../../public/assets/logos/tencent.svg"; import togetheraiLogo from "../../public/assets/logos/togetherai.svg"; import topazLogo from "../../public/assets/logos/topaz.svg"; +import typesafeLogo from "../../public/assets/logos/typesafe.png"; import v0Logo from "../../public/assets/logos/v0.svg"; import vercelLogo from "../../public/assets/logos/vercel.svg"; import vllmLogo from "../../public/assets/logos/vllm.png"; @@ -167,18 +169,20 @@ export enum Providers { SCX_AI = "SCX.ai", Snowflake = "Snowflake", Soniox = "Soniox", + StrandsDecider = "Strands Decider", TEXT_COMPLETION_CODESTRAL = "Text-Completion-Codestral", Tencent = "Tencent", TogetherAI = "TogetherAI", TOPAZ = "Topaz", Triton = "Triton", + TypeSafe = "TypeSafe", V0 = "V0", VERCEL_AI_GATEWAY = "Vercel Ai Gateway", Vertex_AI = "Vertex AI (Anthropic, Gemini, etc.)", VERTEX_AI_BETA = "Vertex Ai Beta", VLLM = "Local vLLM", VolcEngine = "VolcEngine", - Voyage = "Voyage AI", + Voyage = "VoyageAI by MongoDB", WANDB = "Wandb", WATSONX = "Watsonx", WATSONX_TEXT = "Watsonx Text", @@ -287,11 +291,13 @@ export const provider_map: Record = { SCX_AI: "scx-ai", Snowflake: "snowflake", Soniox: "soniox", + StrandsDecider: "strands_decider", TEXT_COMPLETION_CODESTRAL: "text-completion-codestral", Tencent: "tencent", TogetherAI: "together_ai", TOPAZ: "topaz", Triton: "triton", + TypeSafe: "typesafe", V0: "v0", VERCEL_AI_GATEWAY: "vercel_ai_gateway", Vertex_AI: "vertex_ai", @@ -387,11 +393,13 @@ export const providerLogoMap: Partial> = { [Providers.SCX_AI]: scxAiLogo.src, [Providers.Snowflake]: snowflakeLogo.src, [Providers.Soniox]: sonioxLogo.src, + [Providers.StrandsDecider]: strandsLogo.src, [Providers.Tencent]: tencentLogo.src, [Providers.TEXT_COMPLETION_CODESTRAL]: mistralLogo.src, [Providers.TogetherAI]: togetheraiLogo.src, [Providers.TOPAZ]: topazLogo.src, [Providers.Triton]: nvidiaTritonLogo.src, + [Providers.TypeSafe]: typesafeLogo.src, [Providers.V0]: v0Logo.src, [Providers.VERCEL_AI_GATEWAY]: vercelLogo.src, [Providers.Vertex_AI]: googleLogo.src, @@ -457,7 +465,9 @@ const providerPlaceholderMap: Partial> = { [Providers.Sail]: "sail/openai/gpt-oss-120b", [Providers.SCX_AI]: "scx-ai/GLM-5.2", [Providers.Snowflake]: "snowflake/mistral-7b", + [Providers.StrandsDecider]: "strands_decider/strands-decider-2B-hobson-v19", [Providers.Tencent]: "tencent/deepseek-v4-pro", + [Providers.TypeSafe]: "typesafe/jev-latest", [Providers.Vertex_AI]: "gemini-pro", [Providers.VolcEngine]: "volcengine/", [Providers.Voyage]: "voyage/", diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.integration.test.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.integration.test.tsx index 456dc91b13f..3335ebf965a 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.integration.test.tsx @@ -75,6 +75,7 @@ vi.mock("@/app/(dashboard)/hooks/uiSettings/useUISettings", () => ({ vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({ useAllProxyModels: vi.fn(), + useModelAccessGroupNames: vi.fn(() => new Set()), })); vi.mock("@/app/(dashboard)/hooks/teams/useTeams", async (importOriginal) => ({ diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx index d721bdab2a5..b9df537884a 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx @@ -81,6 +81,7 @@ vi.mock("@/app/(dashboard)/hooks/uiSettings/useUISettings", () => ({ vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({ useAllProxyModels: vi.fn(), + useModelAccessGroupNames: vi.fn(() => new Set()), })); vi.mock("@/app/(dashboard)/hooks/teams/useTeams", async (importOriginal) => ({ @@ -231,7 +232,7 @@ vi.mock("../key_team_helpers/filter_helpers", () => ({ fetchAllOrganizations: vi.fn().mockResolvedValue([]), })); -import { useAllProxyModels } from "@/app/(dashboard)/hooks/models/useModels"; +import { useAllProxyModels, useModelAccessGroupNames } from "@/app/(dashboard)/hooks/models/useModels"; import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; import { useOrganization } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; import { teamKeys, teamsTableKeys, useTeam } from "@/app/(dashboard)/hooks/teams/useTeams"; @@ -242,6 +243,7 @@ import { useAccessGroups } from "@/app/(dashboard)/hooks/accessGroups/useAccessG import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; const mockUseAllProxyModels = vi.mocked(useAllProxyModels); +const mockUseModelAccessGroupNames = vi.mocked(useModelAccessGroupNames); const mockUseKeys = vi.mocked(useKeys); const mockUseTeam = vi.mocked(useTeam); const mockUseOrganization = vi.mocked(useOrganization); @@ -292,6 +294,7 @@ const createMockTeamData = (overrides = {}) => ({ }); const seedDefaultMocks = () => { + mockUseModelAccessGroupNames.mockReturnValue(new Set()); mockUseAllProxyModels.mockReturnValue({ data: { data: [] }, isLoading: false, @@ -361,6 +364,35 @@ describe("TeamInfoView", () => { }); describe("display and rendering", () => { + it("links direct model chips to their matching access-group or model filter", async () => { + vi.mocked(networking.teamInfoCall).mockResolvedValue( + createMockTeamData({ models: ["repro-access-group", "gpt-4.1"] }), + ); + mockUseModelAccessGroupNames.mockReturnValue(new Set(["repro-access-group"])); + + renderWithProviders(); + + expect(await screen.findByRole("link", { name: "repro-access-group" })).toHaveAttribute( + "href", + expect.stringMatching(/\?access_group=repro-access-group$/), + ); + expect(screen.getByRole("link", { name: "gpt-4.1" })).toHaveAttribute( + "href", + expect.stringMatching(/\?model_group=gpt-4\.1$/), + ); + + await userEvent.setup({ delay: null }).click(screen.getByRole("tab", { name: "Settings" })); + const settings = await screen.findByRole("tabpanel", { name: "Settings" }); + expect(within(settings).getByRole("link", { name: "repro-access-group" })).toHaveAttribute( + "href", + expect.stringMatching(/\?access_group=repro-access-group$/), + ); + expect(within(settings).getByRole("link", { name: "gpt-4.1" })).toHaveAttribute( + "href", + expect.stringMatching(/\?model_group=gpt-4\.1$/), + ); + }); + it("should render", async () => { vi.mocked(networking.teamInfoCall).mockResolvedValue(createMockTeamData()); diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index 4c62423a474..badfbee0cf1 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -26,7 +26,7 @@ import { ArrowLeftIcon } from "@heroicons/react/outline"; import { StatusBadge, type StatusTone } from "@/components/shared/table_cells/status_badge"; import { BadgeLink } from "@/components/shared/BadgeLink"; import { Badge } from "@/components/ui/badge"; -import { modelGroupHref } from "@/utils/entityLinks"; +import { modelGroupHref, modelOrAccessGroupHref } from "@/utils/entityLinks"; import { Card } from "@/components/ui/card"; import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; import { Input as UIInput } from "@/components/ui/input"; @@ -119,6 +119,7 @@ import { TEAM_INFO_TAB_LABELS, } from "./tabVisibilityUtils"; import TeamMembersComponent from "./TeamMemberTab"; +import { useModelAccessGroupNames } from "@/app/(dashboard)/hooks/models/useModels"; import { isValidThreshold, TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY, @@ -154,8 +155,14 @@ const TEAM_MODEL_BADGE_TONES: Record = { "access-group": "success", }; -const teamModelBadgeHref = (badge: TeamModelBadge): string | undefined => - badge.kind === "direct" || badge.kind === "access-group" ? modelGroupHref(badge.label) : undefined; +const teamModelBadgeHref = ( + badge: TeamModelBadge, + accessGroupNames: ReadonlySet | undefined, +): string | undefined => { + if (badge.kind === "direct") return modelOrAccessGroupHref(badge.label, accessGroupNames); + if (badge.kind === "access-group") return modelGroupHref(badge.label); + return undefined; +}; export type McpGrantResolution = | { readonly kind: "resolved"; readonly serverIds: ReadonlySet } @@ -622,6 +629,7 @@ const TeamInfoView: React.FC = ({ const routerSettingsRef = React.useRef(null); const [organization, setOrganization] = useState(null); const { userRole } = useAuthorized(); + const accessGroupNames = useModelAccessGroupNames(); const { data: allMcpServers = [], isError: mcpServersFailed, isLoading: mcpServersLoading } = useMCPServers(); const { data: allMcpToolsets = [], isError: mcpToolsetsFailed, isLoading: mcpToolsetsLoading } = useMCPToolsets(); const { data: allAccessGroups = [], isError: accessGroupsFailed, isLoading: accessGroupsLoading } = useAccessGroups(); @@ -1355,7 +1363,7 @@ const TeamInfoView: React.FC = ({ @@ -2191,7 +2199,7 @@ const TeamInfoView: React.FC = ({

Models

{info.models.map((model, index) => ( - + {model} ))} @@ -2202,7 +2210,7 @@ const TeamInfoView: React.FC = ({

Default Member Models

{info.default_team_member_models.map((model, index) => ( - + {model} ))} diff --git a/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx b/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx index 09042dea930..88f7fbb9f8c 100644 --- a/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx @@ -23,6 +23,10 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: mockUseAuthorized, })); +vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({ + useModelAccessGroupNames: vi.fn(() => new Set()), +})); + vi.mock("@/app/(dashboard)/hooks/organizations/useOrganizations", () => ({ useOrganizations: () => ({ data: [] }), })); diff --git a/ui/litellm-dashboard/src/components/templates/key_info_view.test.tsx b/ui/litellm-dashboard/src/components/templates/key_info_view.test.tsx index 48f2f47ceda..f4e1fe288ec 100644 --- a/ui/litellm-dashboard/src/components/templates/key_info_view.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_info_view.test.tsx @@ -1,7 +1,7 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import useTeams from "@/app/(dashboard)/hooks/useTeams"; import { renderWithProviders } from "../../../tests/test-utils"; -import { screen, waitFor } from "@testing-library/react"; +import { screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { KeyResponse, Team } from "../key_team_helpers/key_list"; @@ -9,6 +9,12 @@ import { keyDeleteCall, keyUpdateCall } from "../networking"; import { QueryClient } from "@tanstack/react-query"; import KeyInfoView, { needsLifetimeSpendBackfill } from "./key_info_view"; +vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({ + useModelAccessGroupNames: vi.fn(() => new Set()), +})); + +import { useModelAccessGroupNames } from "@/app/(dashboard)/hooks/models/useModels"; + const editViewMocks = vi.hoisted(() => ({ onSubmit: undefined as ((v: Record) => Promise) | undefined, })); @@ -639,6 +645,7 @@ describe("KeyInfoView", () => { beforeEach(() => { vi.mocked(useTeams).mockReturnValue({ teams: [mockTeam], setTeams: vi.fn() }); vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock); + vi.mocked(useModelAccessGroupNames).mockReturnValue(new Set()); }); it("links the key's team by alias, resolved from the teams list, to the team page", async () => { @@ -703,6 +710,37 @@ describe("KeyInfoView", () => { ); }); + it("links access-group model chips to the access-group filter", async () => { + vi.mocked(useModelAccessGroupNames).mockReturnValue(new Set(["repro-access-group"])); + const keyData = { ...MOCK_KEY_DATA, models: ["repro-access-group"] }; + renderWithProviders( + {}} keyId="test-key-id" onKeyDataUpdate={() => {}} teams={[]} />, + ); + + expect(await screen.findByRole("link", { name: "repro-access-group" })).toHaveAttribute( + "href", + expect.stringMatching(/\?access_group=repro-access-group$/), + ); + + await userEvent.setup({ delay: null }).click(screen.getByRole("tab", { name: "Settings" })); + const settings = await screen.findByRole("tabpanel", { name: "Settings" }); + expect(within(settings).getByRole("link", { name: "repro-access-group" })).toHaveAttribute( + "href", + expect.stringMatching(/\?access_group=repro-access-group$/), + ); + }); + + it("keeps access-group model chips unlinked while access-group names are loading", async () => { + vi.mocked(useModelAccessGroupNames).mockReturnValue(undefined); + const keyData = { ...MOCK_KEY_DATA, models: ["repro-access-group"] }; + renderWithProviders( + {}} keyId="test-key-id" onKeyDataUpdate={() => {}} teams={[]} />, + ); + + expect(await screen.findAllByText("repro-access-group")).not.toHaveLength(0); + expect(screen.queryByRole("link", { name: "repro-access-group" })).not.toBeInTheDocument(); + }); + it("keeps the all-proxy-models grant chip non-clickable", async () => { const keyData = { ...MOCK_KEY_DATA, models: ["all-proxy-models"] }; renderWithProviders( diff --git a/ui/litellm-dashboard/src/components/templates/key_info_view.tsx b/ui/litellm-dashboard/src/components/templates/key_info_view.tsx index cfb5e9fa1f8..04a792bb5ca 100644 --- a/ui/litellm-dashboard/src/components/templates/key_info_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_info_view.tsx @@ -1,4 +1,5 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { useModelAccessGroupNames } from "@/app/(dashboard)/hooks/models/useModels"; import { useProjects } from "@/app/(dashboard)/hooks/projects/useProjects"; import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; import { useApplyUserBudgetToTeamKeys } from "@/app/(dashboard)/hooks/uiSettings/useApplyUserBudgetToTeamKeys"; @@ -14,7 +15,7 @@ import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from " import { HoverCard, HoverCardContent, HoverCardTrigger } from "@/components/ui/hover-card"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { EntityLink } from "@/components/shared/EntityLink"; -import { modelGroupHref, teamDetailHref } from "@/utils/entityLinks"; +import { modelOrAccessGroupHref, teamDetailHref } from "@/utils/entityLinks"; import { BadgeLink } from "@/components/shared/BadgeLink"; import { KeyInfoHeader } from "./KeyInfoHeader"; import KeySavingsTab from "./KeySavingsTab"; @@ -95,6 +96,7 @@ export default function KeyInfoView({ backButtonText = "Back to Keys", }: KeyInfoViewProps) { const { accessToken, userId: userID, userRole, premiumUser } = useAuthorized(); + const accessGroupNames = useModelAccessGroupNames(); const activityDateRange = useActivityDateRange(); const queryClient = useQueryClient(); const canEditGuardrails = premiumUser || (userRole != null && rolesWithWriteAccess.includes(userRole)); @@ -752,7 +754,11 @@ export default function KeyInfoView({
{currentKeyData.models && currentKeyData.models.length > 0 ? ( currentKeyData.models.map((model, index) => ( - + {model} )) @@ -1104,7 +1110,11 @@ export default function KeyInfoView({
{currentKeyData.models && currentKeyData.models.length > 0 ? ( currentKeyData.models.map((model, index) => ( - + {model} )) diff --git a/ui/litellm-dashboard/src/components/ui/dialog.tsx b/ui/litellm-dashboard/src/components/ui/dialog.tsx index a4c17736595..e5af118a701 100644 --- a/ui/litellm-dashboard/src/components/ui/dialog.tsx +++ b/ui/litellm-dashboard/src/components/ui/dialog.tsx @@ -28,7 +28,10 @@ function DialogOverlay({ className, ...props }: DialogPrimitive.Backdrop.Props) - + ({ serverRootPath: "" })); -import { modelGroupHref, teamDetailHref, userDetailHref } from "./entityLinks"; +import { accessGroupHref, modelGroupHref, modelOrAccessGroupHref, teamDetailHref, userDetailHref } from "./entityLinks"; describe("userDetailHref", () => { it("targets the users page filtered to the encoded user id", () => { @@ -39,3 +39,30 @@ describe("modelGroupHref", () => { }, ); }); + +describe("accessGroupHref", () => { + it("targets the models page filtered to the encoded access group", () => { + expect(accessGroupHref("a b/c")).toMatch(/\/models-and-endpoints\?access_group=a%20b%2Fc$/); + }); +}); + +describe("modelOrAccessGroupHref", () => { + it("uses the access-group filter for a known access group", () => { + expect(modelOrAccessGroupHref("repro-access-group", new Set(["repro-access-group"]))).toMatch( + /\?access_group=repro-access-group$/, + ); + }); + + it("uses the model-group filter for a name outside the access-group set", () => { + expect(modelOrAccessGroupHref("gpt-4.1", new Set(["repro-access-group"]))).toMatch(/\?model_group=gpt-4\.1$/); + }); + + it("keeps grant sentinels without a link unless they are access groups", () => { + expect(modelOrAccessGroupHref("all-team-models", new Set())).toBeUndefined(); + }); + + it("keeps model names unlinked until access group names are available", () => { + expect(modelOrAccessGroupHref("repro-access-group", undefined)).toBeUndefined(); + expect(modelOrAccessGroupHref("gpt-4.1", undefined)).toBeUndefined(); + }); +}); diff --git a/ui/litellm-dashboard/src/utils/entityLinks.ts b/ui/litellm-dashboard/src/utils/entityLinks.ts index 33f2aa34976..69d47a8f497 100644 --- a/ui/litellm-dashboard/src/utils/entityLinks.ts +++ b/ui/litellm-dashboard/src/utils/entityLinks.ts @@ -29,3 +29,15 @@ export function modelGroupHref(modelGroup: string): string | undefined { if (MODEL_GRANT_SENTINELS.has(modelGroup)) return undefined; return `${uiHref("models-and-endpoints")}?model_group=${encodeURIComponent(modelGroup)}`; } + +export function accessGroupHref(accessGroup: string): string { + return `${uiHref("models-and-endpoints")}?access_group=${encodeURIComponent(accessGroup)}`; +} + +export function modelOrAccessGroupHref( + name: string, + accessGroupNames: ReadonlySet | undefined, +): string | undefined { + if (accessGroupNames === undefined) return undefined; + return accessGroupNames.has(name) ? accessGroupHref(name) : modelGroupHref(name); +} diff --git a/uv.lock b/uv.lock index 231e8cb0d06..a481ecb2a2c 100644 --- a/uv.lock +++ b/uv.lock @@ -4996,12 +4996,12 @@ provides-extras = ["test"] [[package]] name = "litellm-enterprise" -version = "0.1.74" +version = "0.1.75" source = { editable = "enterprise" } [[package]] name = "litellm-proxy-extras" -version = "0.4.106" +version = "0.4.107" source = { editable = "litellm-proxy-extras" } dependencies = [ { name = "psycopg" }, diff --git a/whitelisted_bedrock_models.txt b/whitelisted_bedrock_models.txt index 4bca4af2435..521df60a8ab 100644 --- a/whitelisted_bedrock_models.txt +++ b/whitelisted_bedrock_models.txt @@ -7,6 +7,8 @@ twelvelabs.pegasus-1-2-v1:0 us.twelvelabs.pegasus-1-2-v1:0 eu.twelvelabs.pegasus-1-2-v1:0 global.twelvelabs.pegasus-1-2-v1:0 +us.twelvelabs.pegasus-1-5-v1:0 +global.twelvelabs.pegasus-1-5-v1:0 amazon.titan-text-express-v1 amazon.titan-text-lite-v1 amazon.titan-text-premier-v1:0