diff --git a/.circleci/config.yml b/.circleci/config.yml index 1798abe9de5..2653135da41 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -38,7 +38,7 @@ commands: parameters: category: type: enum - enum: ["backend", "client", "provider-harness"] + enum: ["backend", "client", "provider-harness", "redis-compat"] default: "backend" steps: - run: @@ -229,6 +229,19 @@ commands: - wait_for_service: url: tcp://localhost:6379 timeout: "60" + install_codecov_cli: + steps: + - run: + name: Install Codecov CLI (pinned v11.3.1) + when: always + command: | + curl -sSLf -o /tmp/codecov https://cli.codecov.io/v11.3.1/linux/codecov + curl -sSLf -o /tmp/codecov.SHA256SUM https://cli.codecov.io/v11.3.1/linux/codecov.SHA256SUM + [ "$(cat /tmp/codecov.SHA256SUM)" = "ca1d64196d2d34771084afe76ea657d581bf628e31d993ff8e52ea09cc88a56d codecov" ] + (cd /tmp && sha256sum -c codecov.SHA256SUM) + chmod +x /tmp/codecov + mkdir -p "$HOME/.local/bin" + mv /tmp/codecov "$HOME/.local/bin/codecov" start_openai_record_replay_proxy: description: "Start the record/replay proxy (tests/_openai_record_replay_proxy.py) on host port 8090 and wait until healthy. Models whose api_base points here replay recorded provider responses, so the E2E run neither pays for nor depends on the live provider. The default upstream is OpenAI; a non-OpenAI model must point its api_base at /__recorder_upstream// so the recorder forwards there instead of defaulting to OpenAI. Run after uv deps are synced." steps: @@ -3279,30 +3292,164 @@ jobs: - store_artifacts: path: test-results - unit: + postgres_suite: + parameters: + test_path: + type: string + seed: + type: boolean + coverage_flag: + type: string + default: "" + timeout_minutes: + type: integer machine: image: ubuntu-2204:2024.04.1 resource_class: large working_directory: ~/project - parallelism: 4 + environment: + DATABASE_URL: postgresql://postgres:postgres@localhost:5432/litellm_test steps: + - checkout + - skip_if_unrelated_changes + - setup_litellm_test_deps + - start_postgres: + db_name: litellm_test + image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5 + - run: + name: Generate Prisma client + command: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma + - when: + condition: << parameters.seed >> + steps: + - run: + name: Seed database schema + command: uv run --no-sync prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss + - run: + name: Run << parameters.test_path >> + command: | + mkdir -p test-results + coverage_args=() + if [ -n "<< parameters.coverage_flag >>" ]; then + coverage_args=(--cov=./litellm --cov-report=xml:coverage.xml) + fi + timeout --signal=TERM << parameters.timeout_minutes >>m uv run --no-sync pytest \ + << parameters.test_path >> -vv --tb=short --durations=10 -o junit_family=xunit1 \ + --junitxml=test-results/junit.xml "${coverage_args[@]}" + - when: + condition: << parameters.coverage_flag >> + steps: + - install_codecov_cli + - run: + name: Upload coverage + when: always + command: | + [ -f coverage.xml ] || { echo "no coverage.xml produced"; exit 1; } + codecov upload-process --disable-search --fail-on-error -f coverage.xml \ + -F << parameters.coverage_flag >> -C "$CIRCLE_SHA1" \ + -n "<< parameters.coverage_flag >>-${CIRCLE_BUILD_NUM}" --git-service github + - store_test_results: + path: test-results + - store_artifacts: + path: test-results + + mcp_integration: + machine: + image: ubuntu-2204:2024.04.1 + resource_class: large + working_directory: ~/project + environment: + COVERAGE_CORE: sysmon + LITELLM_LOCAL_MODEL_COST_MAP: "True" + steps: + - checkout + - skip_if_unrelated_changes - setup_litellm_test_deps - run: name: Generate Prisma client command: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma - run: - name: Run unit tests + name: Install MCP SDK1 peer command: | - mkdir -p test-results/unit - shard="$(find tests/unit -name 'test_*.py' | sort | circleci tests split --split-by=timings --timings-type=filename)" - if [ -z "${shard}" ]; then echo "shard ${CIRCLE_NODE_INDEX} received no tests/unit files; nothing to run"; exit 0; fi - mapfile -t files < <(printf '%s\n' "${shard}") - set +e - LITELLM_LOCAL_MODEL_COST_MAP=True uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short -o junit_family=xunit1 --junitxml=test-results/unit/junit.xml - status=$? - set -e - if [ "$status" -eq 5 ]; then echo "pytest collected no tests from tests/unit; passing"; exit 0; fi - exit "$status" + uv venv --python 3.12 .venv-mcp-peer + uv pip install --python .venv-mcp-peer 'mcp==1.28.1' 'langchain-mcp-adapters==0.2.1' + echo "export MCP_TEST_PEER_PYTHON=$PWD/.venv-mcp-peer/bin/python" >> "$BASH_ENV" + - run: + name: Run MCP integration tests + command: | + mkdir -p test-results + env -u OPENAI_API_KEY -u ANTHROPIC_API_KEY \ + timeout --signal=TERM 20m uv run --no-sync pytest \ + tests/mcp_tests tests/unit/experimental_mcp_client tests/unit/proxy/_experimental/mcp_server \ + tests/unit/responses/mcp --tb=short -vv --maxfail=10 -n 2 --dist=loadscope --reruns 0 \ + --reruns-delay 1 --timeout=120 --rerun-except "from pytest-timeout" --durations=20 \ + --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml \ + --cov-config=pyproject.toml -o junit_family=xunit1 --junitxml=test-results/junit.xml + - install_codecov_cli + - run: + name: Upload MCP integration coverage + when: always + command: | + [ -f coverage.xml ] || { echo "no coverage.xml produced"; exit 1; } + codecov upload-process --disable-search --fail-on-error -f coverage.xml \ + -F mcp-integration -C "$CIRCLE_SHA1" \ + -n "mcp-integration-${CIRCLE_BUILD_NUM}" --git-service github + - store_test_results: + path: test-results + - store_artifacts: + path: test-results + + redis_compat: + parameters: + redis_py: + type: string + machine: + image: ubuntu-2204:2024.04.1 + resource_class: large + working_directory: ~/project + steps: + - checkout + - skip_if_unrelated_changes: + category: redis-compat + - setup_litellm_test_deps + - run: + name: Pin redis-py version + command: | + uv pip install "redis==<< parameters.redis_py >>" + uv run --no-sync python -c "import redis; assert redis.__version__ == '<< parameters.redis_py >>', redis.__version__; print('redis-py', redis.__version__)" + - run: + name: Install redis-server 7.2.16 + command: | + mkdir -p "$HOME/.local/bin" + cid="$(docker create redis:7.2.16@sha256:0637954999d01b7c9ce9167db2da50656e2590d3b884f1c600c5f63bb6e6773c)" + docker cp "${cid}":/usr/local/bin/redis-server "$HOME/.local/bin/redis-server" + docker rm "$cid" + redis-server --version | grep -q 'v=7.2.16' + - run: + name: Run Redis compatibility tests + command: | + mkdir -p test-results + env -u CASSETTE_REDIS_URL -u AZURE_CLIENT_ID -u AZURE_CLIENT_SECRET -u AZURE_TENANT_ID \ + timeout --signal=TERM 15m uv run --no-sync pytest \ + tests/unit/test_redis.py tests/unit/caching/test_redis_connection_pool.py \ + tests/unit/caching/test_redis_cluster_cache.py tests/unit/caching/test_evicted_client_closer.py \ + tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \ + tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \ + --tb=short -vv --reruns 2 --reruns-delay 1 --durations=20 --cov=./litellm \ + --cov-report=xml:coverage.xml -o junit_family=xunit1 --junitxml=test-results/junit.xml + - when: + condition: + equal: ["5.3.1", << parameters.redis_py >>] + steps: + - install_codecov_cli + - run: + name: Upload Redis compatibility coverage + when: always + command: | + [ -f coverage.xml ] || { echo "no coverage.xml produced"; exit 1; } + codecov upload-process --disable-search --fail-on-error -f coverage.xml \ + -F redis-compat -C "$CIRCLE_SHA1" \ + -n "redis-compat-${CIRCLE_BUILD_NUM}" --git-service github - store_test_results: path: test-results - store_artifacts: @@ -3376,6 +3523,39 @@ workflows: parameters: suite: [management, database] mode: [replica] + - postgres_suite: + name: proxy-behavior + test_path: tests/proxy_behavior + seed: true + coverage_flag: lens-postgres + timeout_minutes: 25 + - postgres_suite: + name: proxy-security + test_path: tests/proxy_security_tests + seed: true + timeout_minutes: 15 + - postgres_suite: + name: schema-migration + test_path: tests/proxy_migration_tests + seed: false + timeout_minutes: 20 + - postgres_suite: + name: roi-database + test_path: tests/integration/database/test_roi_observed.py + seed: false + coverage_flag: roi-postgres + timeout_minutes: 10 + - mcp_integration: + name: mcp-integration + - redis_compat: + name: redis-compat-<< matrix.redis_py >> + matrix: + parameters: + redis_py: + - "5.3.1" + - "6.4.0" + - "7.4.1" + - "8.0.1" build_and_test: unless: or: @@ -3385,7 +3565,6 @@ workflows: jobs: - using_litellm_on_windows - windows_release_wheel - - unit - provider_replay_harness - base_sdk_install - local_testing_part1 diff --git a/.circleci/scripts/classify_changes.sh b/.circleci/scripts/classify_changes.sh index 387197b65d7..273716025e5 100755 --- a/.circleci/scripts/classify_changes.sh +++ b/.circleci/scripts/classify_changes.sh @@ -1,7 +1,7 @@ #!/usr/bin/env bash set -uo pipefail -category="${1:?usage: classify_changes.sh }" +category="${1:?usage: classify_changes.sh }" has_client=false has_backend=false @@ -10,9 +10,14 @@ has_provider_harness=false has_cost_map=false has_mcp_dependencies=false has_windows_release=false +has_redis_compat=false outside_cost_map_set=false while IFS= read -r file || [ -n "$file" ]; do [ -n "$file" ] || continue + case "$file" in + litellm/_redis.py | litellm/_redis_credential_provider.py | litellm/caching/redis_cache.py | litellm/caching/evicted_client_closer.py | tests/unit/test_redis.py | tests/local_testing/test_caching.py | tests/unit/caching/test_redis_connection_pool.py | tests/unit/caching/test_redis_cluster_cache.py | tests/unit/caching/test_evicted_client_closer.py | .circleci/config.yml | .circleci/scripts/classify_changes.sh | .circleci/scripts/path_filter.sh | pyproject.toml | uv.lock) + has_redis_compat=true ;; + esac case "$file" in *.md | *.mdx) : ;; pyproject.toml | */pyproject.toml | uv.lock | uv.toml | .python-version | rust-toolchain.toml | litellm-rust/* | litellm/__init__.py | litellm/proxy/proxy_server.py | litellm/*mcp* | tests/*mcp* | litellm/integrations/arize/* | tests/base_sdk_tests/* | scripts/check_mcp_sdk_install.py | .github/workflows/test-mcp-dependency-resolution.yml | .github/actions/detect-changes/* | .github/actions/setup-uv-with-retries/* | .github/actions/cache-cargo-build/* | .github/scripts/detect_changes.sh | .github/scripts/uv_sync_with_retries.sh | .circleci/scripts/classify_changes.sh | tests/unit/test_circleci_path_filter.py | tests/unit/test_detect_changes.py) @@ -54,6 +59,9 @@ case "$category" in windows-release) [ "$has_windows_release" = true ] && echo run || echo skip ;; + redis-compat) + [ "$has_redis_compat" = true ] && echo run || echo skip + ;; backend) [ "$has_backend" = true ] && echo run || echo skip ;; diff --git a/.circleci/scripts/path_filter.sh b/.circleci/scripts/path_filter.sh index 3050674f562..ccc68a0750b 100755 --- a/.circleci/scripts/path_filter.sh +++ b/.circleci/scripts/path_filter.sh @@ -1,7 +1,7 @@ #!/usr/bin/env bash set -uo pipefail -category="${1:?usage: path_filter.sh }" +category="${1:?usage: path_filter.sh }" here="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" run_full() { diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh deleted file mode 100755 index 02df32d5eab..00000000000 --- a/.circleci/scripts/unit_selection.sh +++ /dev/null @@ -1,187 +0,0 @@ -#!/usr/bin/env bash -set -euo pipefail - -flag="${1:?usage: unit_selection.sh }" - -legacy_flags=( - caching-local - core-utils - enterprise-package - enterprise-routing - integrations - llm-other-providers - llm-vertex-ai - mcp-integration - misc - proxy-db-auth-checks - proxy-db-budgets - proxy-db-custom-logging - proxy-db-db-and-spend - proxy-db-endpoints-and-responses - proxy-db-guardrails-hooks - proxy-db-jwt-and-keys - proxy-db-key-generation - proxy-db-logging-misc - proxy-db-proxy-runtime - proxy-db-proxy-server-core - proxy-db-proxy-utils - proxy-extras - proxy-infra - responses-caching-types -) - -legacy_paths() { - case "$1" in - caching-local) echo tests/unit/caching ;; - core-utils) echo tests/unit/litellm_core_utils ;; - enterprise-package) - echo tests/unit/enterprise/integrations - echo tests/unit/enterprise/proxy/auth - echo tests/unit/enterprise/proxy/guardrails - echo tests/unit/enterprise/proxy/hooks - echo tests/unit/enterprise/proxy/management_endpoints - echo tests/unit/enterprise/proxy/test_audit_logging_endpoints.py - echo tests/unit/enterprise/proxy/test_liteadmin.py - echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;; - enterprise-routing) - echo tests/unit/google_genai - echo tests/unit/router_strategy - echo tests/unit/router_utils - echo tests/unit/proxy/common_utils/test_cache_aware_routing.py - echo tests/unit/enterprise/enterprise_callbacks/send_emails - echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py - echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py - echo tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py - echo tests/unit/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py - echo tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py - echo tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py - echo tests/unit/enterprise/proxy/test_enterprise_routes.py - echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py - echo tests/unit/enterprise/proxy/test_managed_files_access_check.py - echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;; - integrations) echo tests/unit/integrations ;; - llm-other-providers) find tests/unit/llms -name 'test_*.py' -not -path 'tests/unit/llms/vertex_ai/*' ;; - llm-vertex-ai) echo tests/unit/llms/vertex_ai ;; - mcp-integration) - echo tests/unit/experimental_mcp_client - echo tests/unit/proxy/_experimental/mcp_server - echo tests/unit/responses/mcp - echo tests/mcp_tests/test_proxy_mcp_e2e.py ;; - misc) - find tests/unit -maxdepth 1 -name 'test_*.py' - echo tests/unit/test_router - echo tests/unit/a2a_protocol - echo tests/unit/batches - echo tests/unit/chat_completions - echo tests/unit/completion_extras - echo tests/unit/containers - echo tests/unit/embeddings - echo tests/unit/endpoints - echo tests/unit/files - echo tests/unit/harness - echo tests/unit/images - echo tests/unit/interactions - echo tests/unit/messages - echo tests/unit/rag - echo tests/unit/rerank_api - echo tests/unit/rust_bridge - echo tests/unit/secret_managers - echo tests/unit/vector_stores - echo tests/unit/videos ;; - proxy-db-auth-checks) - echo tests/unit/proxy/auth/test_auth_checks.py - echo tests/unit/proxy/auth/test_user_api_key_auth.py - echo tests/unit/proxy/test_credential_slot_registry.py - echo tests/unit/proxy/test_deprecated_key_grace_period.py ;; - proxy-db-budgets) - echo tests/unit/proxy/auth/test_default_end_user_budget_simple.py - echo tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py - echo tests/unit/proxy/test_zero_cost_model_budget_bypass.py ;; - proxy-db-custom-logging) - echo tests/unit/proxy/test_custom_callback_input.py - echo tests/unit/proxy/test_custom_logger_s3_gcs.py ;; - proxy-db-db-and-spend) - echo tests/unit/proxy/common_utils/test_proxy_encrypt_decrypt.py - echo tests/unit/proxy/db/db_transaction_queue/test_e2e_pod_lock_manager.py - echo tests/unit/proxy/db/test_update_daily_tag_spend.py - echo tests/unit/proxy/test_db_schema_changes.py - echo tests/unit/proxy/test_prisma_client_backoff_retry.py - echo tests/unit/proxy/test_update_spend.py - echo tests/unit/skills/test_skills_db.py ;; - proxy-db-endpoints-and-responses) - echo tests/unit/proxy/lens - echo tests/unit/proxy/auth/test_models_fallback_endpoint.py - echo tests/unit/proxy/common_utils/test_check_batch_cost.py - echo tests/unit/proxy/common_utils/test_check_responses_cost.py - echo tests/unit/proxy/common_utils/test_realtime_cache.py - echo tests/unit/proxy/google_endpoints/test_gemini_agents_endpoints.py - echo tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py - echo tests/unit/proxy/google_endpoints/test_google_gemini_proxy_request.py - echo tests/unit/proxy/public_endpoints/test_blog_posts_endpoint.py - echo tests/unit/proxy/response_polling - echo tests/unit/proxy/test_custom_tokenizer_bug.py - echo tests/unit/proxy/test_get_favicon.py - echo tests/unit/proxy/test_get_image.py - echo tests/unit/proxy/test_prompt_test_endpoint.py - echo tests/unit/proxy/test_reducto_ocr_route.py - echo tests/unit/proxy/test_response_polling_pre_call_checks.py - echo tests/unit/proxy/test_ui_path_detection.py ;; - proxy-db-guardrails-hooks) - echo tests/unit/proxy/hooks/test_banned_keyword_list.py - echo tests/unit/proxy/test_proxy_setting_guardrails.py - echo tests/unit/proxy/test_unit_test_proxy_hooks.py ;; - proxy-db-jwt-and-keys) - echo tests/unit/proxy/auth/test_jwt.py - echo tests/unit/proxy/management_endpoints/test_jwt_key_mapping.py - echo tests/unit/proxy/test_proxy_custom_auth.py ;; - proxy-db-key-generation) echo tests/unit/proxy/management_endpoints/test_key_generate_prisma.py ;; - proxy-db-logging-misc) - echo tests/unit/proxy/management_helpers/test_audit_logs_proxy.py - echo tests/unit/proxy/spend_tracking/test_search_api_logging.py - echo tests/unit/proxy/test_proxy_reject_logging.py ;; - proxy-db-proxy-runtime) - echo tests/unit/proxy/auth/test_multipart_bypass_repro.py - echo tests/unit/proxy/auth/test_proxy_routes.py - echo tests/unit/proxy/middleware/test_request_size_limit_middleware.py - echo tests/unit/proxy/test_proxy_config_unit_test.py - echo tests/unit/proxy/test_proxy_token_counter.py - echo tests/unit/proxy/test_server_root_path.py ;; - proxy-db-proxy-server-core) - echo tests/unit/proxy/test__lazy_features.py - echo tests/unit/proxy/test_aproxy_startup.py - echo tests/unit/proxy/test_proxy_server.py ;; - proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;; - proxy-extras) echo tests/unit/litellm_proxy_extras ;; - proxy-infra) - echo tests/unit/gateway - echo tests/unit/proxy/management - echo tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py - echo tests/unit/proxy/roi_calculator ;; - responses-caching-types) - find tests/unit/responses -name 'test_*.py' -not -path 'tests/unit/responses/mcp/*' - echo tests/unit/types ;; - *) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;; - esac -} - -expand() { - while read -r path; do - if [ -d "$path" ]; then - find "$path" -name 'test_*.py' - elif [ -f "$path" ]; then - echo "$path" - else - echo "unit_selection.sh: $path does not exist" >&2 - exit 1 - fi - done -} - -if [ "$flag" = unit ]; then - comm -23 \ - <(find tests/unit -name 'test_*.py' | sort) \ - <(for legacy in "${legacy_flags[@]}"; do legacy_paths "$legacy"; done | expand | sort) - exit 0 -fi - -legacy_paths "$flag" | expand | sort diff --git a/.circleci/tests.yml b/.circleci/tests.yml deleted file mode 100644 index a9cd21bad5e..00000000000 --- a/.circleci/tests.yml +++ /dev/null @@ -1,416 +0,0 @@ -version: 2.1 - -commands: - wait_for_service: - parameters: - url: - type: string - timeout: - type: string - default: "60" - steps: - - run: - name: "Wait for << parameters.url >>" - command: | - TIMEOUT=<< parameters.timeout >> - URL="<< parameters.url >>" - ELAPSED=0 - echo "Waiting up to ${TIMEOUT}s for ${URL} ..." - if echo "$URL" | grep -q '^tcp://'; then - HOST=$(echo "$URL" | sed 's|tcp://||' | cut -d: -f1) - PORT=$(echo "$URL" | sed 's|tcp://||' | cut -d: -f2) - while ! bash -c "echo > /dev/tcp/$HOST/$PORT" 2>/dev/null; do - sleep 2; ELAPSED=$((ELAPSED+2)) - if [ "$ELAPSED" -ge "$TIMEOUT" ]; then echo "Timed out"; exit 1; fi - done - else - while ! curl -sf --max-time 5 "$URL" > /dev/null 2>&1; do - sleep 2; ELAPSED=$((ELAPSED+2)) - if [ "$ELAPSED" -ge "$TIMEOUT" ]; then echo "Timed out"; exit 1; fi - done - fi - echo "Service ready after ${ELAPSED}s" - install_uv: - steps: - - run: - name: Install uv (pinned 0.10.9) - command: | - curl -LsSf -o /tmp/uv-install.sh https://astral.sh/uv/0.10.9/install.sh - echo "7fc46e39cb97290b57169c0c813a17970585ac519139f19006453c99b5f2f45f /tmp/uv-install.sh" | sha256sum -c - - env UV_NO_MODIFY_PATH=1 sh /tmp/uv-install.sh - rm -f /tmp/uv-install.sh - echo 'export PATH="$HOME/.local/bin:$PATH"' >> "$BASH_ENV" - export PATH="$HOME/.local/bin:$PATH" - install_rust: - steps: - - run: - name: Install Rust (rustup 1.28.2, toolchain 1.98.0) - command: | - case "$(uname -m)" in - x86_64) - RUSTUP_TRIPLE=x86_64-unknown-linux-gnu - RUSTUP_SHA256=20a06e644b0d9bd2fbdbfd52d42540bdde820ea7df86e92e533c073da0cdd43c - ;; - aarch64) - RUSTUP_TRIPLE=aarch64-unknown-linux-gnu - RUSTUP_SHA256=e3853c5a252fca15252d07cb23a1bdd9377a8c6f3efa01531109281ae47f841c - ;; - *) - echo "install_rust: unsupported architecture $(uname -m)" >&2 - exit 1 - ;; - esac - curl -sSLf -o /tmp/rustup-init \ - "https://static.rust-lang.org/rustup/archive/1.28.2/${RUSTUP_TRIPLE}/rustup-init" - echo "${RUSTUP_SHA256} /tmp/rustup-init" | sha256sum -c - - chmod +x /tmp/rustup-init - /tmp/rustup-init -y --no-modify-path --profile minimal --default-toolchain 1.98.0 - rm -f /tmp/rustup-init - echo 'export PATH="$HOME/.cargo/bin:$PATH"' >> "$BASH_ENV" - export PATH="$HOME/.cargo/bin:$PATH" - rustc --version - cargo --version - install_codecov_cli: - steps: - - run: - name: Install Codecov CLI (pinned v11.3.1) - when: always - command: | - curl -sSLf -o /tmp/codecov https://cli.codecov.io/v11.3.1/linux/codecov - curl -sSLf -o /tmp/codecov.SHA256SUM https://cli.codecov.io/v11.3.1/linux/codecov.SHA256SUM - [ "$(cat /tmp/codecov.SHA256SUM)" = "ca1d64196d2d34771084afe76ea657d581bf628e31d993ff8e52ea09cc88a56d codecov" ] - (cd /tmp && sha256sum -c codecov.SHA256SUM) - chmod +x /tmp/codecov - mkdir -p "$HOME/.local/bin" - mv /tmp/codecov "$HOME/.local/bin/codecov" - setup_litellm_enterprise_pip: - steps: - - run: - name: "Install local version of litellm-enterprise" - command: | - uv run --no-sync python -c "import litellm_enterprise; print('litellm-enterprise OK:', litellm_enterprise.__file__)" - setup_test_deps: - steps: - - install_uv - - install_rust - - restore_cache: - keys: - - v3-integration-uv-cache-{{ checksum "uv.lock" }} - - run: - name: Install Dependencies - command: | - uv sync --frozen --all-groups --all-extras --python 3.12 - - setup_litellm_enterprise_pip - - save_cache: - paths: - - ~/.cache/uv - key: v3-integration-uv-cache-{{ checksum "uv.lock" }} - - run: - name: Generate Prisma client - command: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma - skip_unless_relevant: - parameters: - category: - type: string - default: backend - base_ref: - type: string - default: "" - pull_request_url: - type: string - default: "" - steps: - - run: - name: "Skip job when no << parameters.category >>-relevant files changed" - command: | - export CIRCLE_PULL_REQUEST="${CIRCLE_PULL_REQUEST:-<< parameters.pull_request_url >>}" - export PATH_FILTER_BASE_BRANCH="<< parameters.base_ref >>" - [ -n "$PATH_FILTER_BASE_BRANCH" ] || unset PATH_FILTER_BASE_BRANCH - bash .circleci/scripts/path_filter.sh << parameters.category >> - start_postgres: - parameters: - db_name: - type: string - default: circle_test - image: - type: string - default: postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26 - steps: - - run: - name: Start PostgreSQL - command: | - docker run -d \ - --name postgres-db \ - -e POSTGRES_USER=postgres \ - -e POSTGRES_PASSWORD=postgres \ - -e POSTGRES_DB=<< parameters.db_name >> \ - -p 5432:5432 \ - << parameters.image >> - - wait_for_service: - url: tcp://localhost:5432 - timeout: "60" - start_redis: - steps: - - run: - name: Start Redis - command: | - docker run -d \ - --name redis-cache \ - -p 6379:6379 \ - redis:7-alpine@sha256:7aec734b2bb298a1d769fd8729f13b8514a41bf90fcdd1f38ec52267fbaa8ee6 - - wait_for_service: - url: tcp://localhost:6379 - timeout: "60" - -jobs: - unit: - parameters: - flag: - type: string - default: unit - shards: - type: integer - default: 6 - workers: - type: integer - default: 4 - dist: - type: string - default: loadscope - base_ref: - type: string - default: "" - pull_request_url: - type: string - default: "" - legacy_mcp_peer: - type: boolean - default: false - reruns: - type: integer - default: 0 - machine: - image: ubuntu-2204:2024.04.1 - resource_class: large - working_directory: ~/project - parallelism: << parameters.shards >> - environment: - COVERAGE_CORE: sysmon - LITELLM_LOCAL_MODEL_COST_MAP: "True" - steps: - - checkout - - skip_unless_relevant: - base_ref: << parameters.base_ref >> - pull_request_url: << parameters.pull_request_url >> - - setup_test_deps - - when: - condition: << parameters.legacy_mcp_peer >> - steps: - - run: - name: Install MCP SDK1 peer - command: | - uv venv --python 3.12 .venv-mcp-peer - uv pip install --python .venv-mcp-peer 'mcp==1.28.1' 'langchain-mcp-adapters==0.2.1' - echo "export MCP_TEST_PEER_PYTHON=$PWD/.venv-mcp-peer/bin/python" >> "$BASH_ENV" - - run: - name: "Run << parameters.flag >> shard" - no_output_timeout: 20m - command: | - mkdir -p test-results/<< parameters.flag >> - selection="$(bash .circleci/scripts/unit_selection.sh << parameters.flag >>)" || { echo "unit_selection.sh failed for << parameters.flag >>"; exit 1; } - [ -n "${selection}" ] || { echo "unit_selection.sh produced no files for << parameters.flag >>"; exit 1; } - shard="$(printf '%s\n' "${selection}" | circleci tests split --split-by=timings --timings-type=filename)" || { echo "circleci tests split failed for << parameters.flag >>"; exit 1; } - [ -n "${shard}" ] || { echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.flag >> files; nothing to run"; exit 0; } - mapfile -t files < <(printf '%s\n' "${shard}") - xdist_args=() - if [ "<< parameters.workers >>" -gt 0 ]; then xdist_args=(-n << parameters.workers >> --dist=<< parameters.dist >>); fi - rerun_args=(-p no:rerunfailures) - if [ "<< parameters.reruns >>" -gt 0 ]; then rerun_args=(--reruns << parameters.reruns >> --reruns-delay 1 --rerun-except "from pytest-timeout"); fi - test_env=(PATH="$PATH" HOME="$HOME" CI=true COVERAGE_CORE="$COVERAGE_CORE" LITELLM_LOCAL_MODEL_COST_MAP="$LITELLM_LOCAL_MODEL_COST_MAP") - if [ -n "${MCP_TEST_PEER_PYTHON:-}" ]; then test_env+=(MCP_TEST_PEER_PYTHON="$MCP_TEST_PEER_PYTHON"); fi - set +e - env -i "${test_env[@]}" \ - uv run --no-sync pytest "${files[@]}" "${rerun_args[@]}" -p no:pytest-retry --timeout=90 "${xdist_args[@]}" --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml - status=$? - set -e - if [ "$status" -eq 5 ]; then echo "pytest collected no tests from the shard; passing"; exit 0; fi - exit "$status" - - install_codecov_cli - - run: - name: Upload coverage - when: always - command: | - [ -f coverage.xml ] || { echo "no coverage.xml produced; skipping upload"; exit 0; } - codecov upload-process --disable-search -f coverage.xml -F << parameters.flag >> -C "$CIRCLE_SHA1" -n "<< parameters.flag >>-${CIRCLE_NODE_INDEX}-${CIRCLE_BUILD_NUM}" --git-service github - - store_test_results: - path: test-results - - store_artifacts: - path: test-results - - store_artifacts: - path: coverage.xml - documentation: - machine: - image: ubuntu-2204:2024.04.1 - resource_class: large - working_directory: ~/project - steps: - - checkout - - setup_test_deps - - run: - name: Checkout litellm-docs - command: rm -rf docs/my-website && git clone --depth 1 https://github.com/BerriAI/litellm-docs.git docs/my-website - - run: - name: Run documentation validation - command: | - uv run --no-sync python ./tests/documentation_tests/test_env_keys.py - uv run --no-sync python ./tests/documentation_tests/test_router_settings.py - uv run --no-sync python ./tests/documentation_tests/test_api_docs.py - uv run --no-sync python ./tests/documentation_tests/test_circular_imports.py - integration: - parameters: - suite: - type: string - base_ref: - type: string - default: "" - pull_request_url: - type: string - default: "" - machine: - image: ubuntu-2204:2024.04.1 - resource_class: large - working_directory: ~/project - steps: - - checkout - - skip_unless_relevant: - base_ref: << parameters.base_ref >> - pull_request_url: << parameters.pull_request_url >> - - setup_test_deps - - start_postgres: - image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5 - - start_redis - - run: - name: Run owned integration contracts - command: env -i PATH="$PATH" HOME="$HOME" CIRCLE_SHA1="$CIRCLE_SHA1" CIRCLE_WORKFLOW_ID="$CIRCLE_WORKFLOW_ID" bash .circleci/scripts/run_integration.sh << parameters.suite >> - no_output_timeout: 15m - - run: - name: Stop owned database and Redis - when: always - command: | - mkdir -p test-results/integration-<< parameters.suite >> - docker logs postgres-db > test-results/integration-<< parameters.suite >>/postgres.log 2>&1 || true - docker logs redis-cache > test-results/integration-<< parameters.suite >>/redis.log 2>&1 || true - docker rm -f postgres-db redis-cache - test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)" - - store_test_results: - path: test-results - - store_artifacts: - path: test-results - -workflows: - tests: - when: (pipeline.event.name == "push" and pipeline.git.branch == "main") or pipeline.event.name == "api" or (pipeline.event.name == "pull_request" and (pipeline.event.github.pull_request.base.ref == "main" or pipeline.event.github.pull_request.base.ref starts-with "litellm_")) - jobs: - - unit: - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - - unit: - name: unit-<< matrix.flag >> - shards: 1 - workers: 2 - reruns: 2 - matrix: - parameters: - flag: [caching-local, proxy-extras, enterprise-routing] - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - - unit: - name: unit-mcp-integration - flag: mcp-integration - shards: 1 - workers: 2 - legacy_mcp_peer: true - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - - unit: - name: unit-<< matrix.flag >> - shards: 1 - reruns: 2 - matrix: - parameters: - flag: - - enterprise-package - - proxy-infra - - responses-caching-types - - proxy-db-auth-checks - - proxy-db-jwt-and-keys - - proxy-db-proxy-server-core - - proxy-db-proxy-runtime - - proxy-db-custom-logging - - proxy-db-logging-misc - - proxy-db-db-and-spend - - proxy-db-guardrails-hooks - - proxy-db-budgets - - proxy-db-endpoints-and-responses - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - - unit: - name: unit-llm-vertex-ai - flag: llm-vertex-ai - shards: 2 - workers: 1 - reruns: 2 - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - - unit: - name: unit-llm-other-providers - flag: llm-other-providers - shards: 3 - reruns: 2 - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - - unit: - name: unit-core-utils - flag: core-utils - shards: 2 - reruns: 1 - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - - unit: - name: unit-integrations - flag: integrations - shards: 2 - reruns: 3 - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - - unit: - name: unit-misc - flag: misc - shards: 2 - reruns: 2 - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - - unit: - name: unit-proxy-db-proxy-utils - flag: proxy-db-proxy-utils - shards: 1 - reruns: 2 - dist: worksteal - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - - unit: - name: unit-proxy-db-key-generation - flag: proxy-db-key-generation - shards: 1 - workers: 0 - reruns: 2 - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - - documentation - - integration: - name: integration-<< matrix.suite >> - matrix: - parameters: - suite: [sdk] - base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> - pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> diff --git a/.github/ci-coverage-allowlist.yml b/.github/ci-coverage-allowlist.yml index eea25e8e285..8ca45ba3f75 100644 --- a/.github/ci-coverage-allowlist.yml +++ b/.github/ci-coverage-allowlist.yml @@ -21,7 +21,7 @@ test_paths: Live-provider caching cases in tests/local_testing that remain outside CI. Jobs that glob that directory either deselect them (local_testing_part1 and part2 carry `-k "... and not caching and not cache"`) or keep only another keyword (langfuse, router, assistants). - Separately, test-redis-compat.yml selects two IAM cluster authentication tests in + Separately, the CircleCI redis-compat jobs select two IAM cluster authentication tests in test_caching.py by node ID. It does not run that file's other tests. The gap was eight files and 118 tests when measured 2026-08-20; the five keyless files now run in the caching-local shard, leaving live cases in these three. Measured 2026-08-21 with no provider diff --git a/.github/e2e-stack/down.sh b/.github/e2e-stack/down.sh deleted file mode 100755 index 740d626beea..00000000000 --- a/.github/e2e-stack/down.sh +++ /dev/null @@ -1,17 +0,0 @@ -#!/usr/bin/env bash -set -uo pipefail - -STACK_DIR="${E2E_STACK_DIR:-${RUNNER_TEMP:-/tmp}/litellm-e2e-stack}" - -for pid_file in "${STACK_DIR}"/pids/*.pid; do - [[ -f "${pid_file}" ]] || continue - pkill -TERM -P "$(cat "${pid_file}")" 2>/dev/null - kill -TERM "$(cat "${pid_file}")" 2>/dev/null - rm -f "${pid_file}" -done - -for container in e2e-nginx e2e-keycloak e2e-valkey e2e-jaeger e2e-postgres; do - docker rm -f "${container}" >/dev/null 2>&1 -done - -exit 0 diff --git a/.github/e2e-stack/redact_output.py b/.github/e2e-stack/redact_output.py deleted file mode 100644 index 233a8be3e2d..00000000000 --- a/.github/e2e-stack/redact_output.py +++ /dev/null @@ -1,83 +0,0 @@ -import argparse -import os -import sys -from functools import reduce -from pathlib import Path -from typing import Final -from xml.sax.saxutils import escape - -from pydantic import JsonValue, TypeAdapter, ValidationError -from secrets_to_env import MIN_MASKED_LENGTH - -REDACTED: Final = "***" -json_adapter: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) - - -def string_leaves(node: JsonValue) -> tuple[str, ...]: - match node: - case str(): - return (node,) - case list(): - return tuple(leaf for child in node for leaf in string_leaves(child)) - case dict(): - return tuple(leaf for child in node.values() for leaf in string_leaves(child)) - return () - - -def field_lines(value: str) -> tuple[str, ...]: - try: - return tuple(line for leaf in string_leaves(json_adapter.validate_json(value)) for line in leaf.splitlines()) - except ValidationError: - return () - - -def masked_values(values_files: tuple[Path, ...]) -> tuple[str, ...]: - values: Final = frozenset( - line.split("=", 1)[1].strip().strip("'") - for path in values_files - for line in path.read_text().splitlines() - if "=" in line - ) - texts: Final = frozenset(text for value in values for text in (value, *field_lines(value))) - renderings: Final = frozenset( - rendering - for text in texts - if len(text) >= MIN_MASKED_LENGTH - for rendering in (text, escape(text), escape(text, {'"': """})) - ) - return tuple(sorted(renderings, key=lambda rendering: (-len(rendering), rendering))) - - -def redact(text: str, values: tuple[str, ...]) -> str: - return reduce(lambda redacted, value: redacted.replace(value, REDACTED), values, text) - - -def write_redacted(source: Path, out_dir: Path, values: tuple[str, ...]) -> None: - target: Final = out_dir / source.name - with os.fdopen(os.open(target, os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW, 0o600), "w") as handle: - _ = handle.write(redact(source.read_text(errors="replace"), values)) - - -def main() -> int: - parser: Final = argparse.ArgumentParser() - _ = parser.add_argument("--values", action="append", type=Path, required=True) - _ = parser.add_argument("--out", type=Path, required=True) - _ = parser.add_argument("files", nargs="*", type=Path) - args: Final = parser.parse_args() - values_files: Final = tuple(args.values) - out_dir: Final[Path] = args.out - sources: Final = tuple(args.files) - try: - values: Final = masked_values(values_files) - out_dir.mkdir(mode=0o700, exist_ok=True) - for source in sources: - write_redacted(source, out_dir, values) - except OSError as error: - _ = sys.stderr.write(f"could not redact {error.filename}\n") - return 1 - _ = sys.stdout.write(f"redacted {len(sources)} file(s) into {out_dir}\n") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/.github/e2e-stack/secrets_to_env.py b/.github/e2e-stack/secrets_to_env.py deleted file mode 100644 index 691c10203bf..00000000000 --- a/.github/e2e-stack/secrets_to_env.py +++ /dev/null @@ -1,55 +0,0 @@ -import os -import re -import sys -from pathlib import Path -from typing import Final - -from pydantic import TypeAdapter, ValidationError - -secrets_adapter: Final[TypeAdapter[dict[str, str]]] = TypeAdapter(dict[str, str]) -ENV_NAME: Final = re.compile(r"[A-Za-z_][A-Za-z0-9_]*") -MIN_MASKED_LENGTH: Final = 8 -ACTIONS_RUNNER_FLAG: Final = "GITHUB_ACTIONS" - - -def main() -> int: - env_path: Final = Path(sys.argv[1]) - try: - secrets: Final = { - key: value.rstrip("\r\n") for key, value in secrets_adapter.validate_json(sys.stdin.read()).items() - } - except (ValidationError, UnicodeError): - _ = sys.stderr.write("expected a JSON object containing string environment values\n") - return 1 - unusable: Final = tuple( - key - for key, value in secrets.items() - if ENV_NAME.fullmatch(key) is None or any(char in value for char in "'\n\r\0") - ) - if unusable: - _ = sys.stderr.write( - f"these names or values cannot be represented in both bash and dotenv: {' '.join(sorted(unusable))}\n" - ) - return 1 - if os.environ.get(ACTIONS_RUNNER_FLAG) == "true": - _ = sys.stdout.write( - "".join( - f"::add-mask::{value.replace('%', '%25')}\n" - for value in secrets.values() - if len(value) >= MIN_MASKED_LENGTH - ) - ) - sys.stdout.flush() - lines: Final = tuple(f"{key}='{value}'" for key, value in secrets.items() if value) - try: - with os.fdopen(os.open(env_path, os.O_WRONLY | os.O_APPEND | os.O_CREAT | os.O_NOFOLLOW, 0o600), "w") as handle: - os.fchmod(handle.fileno(), 0o600) - _ = handle.write("\n".join(lines) + "\n") - except OSError: - _ = sys.stderr.write("could not write the environment file\n") - return 1 - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/.github/e2e-stack/select_tests.py b/.github/e2e-stack/select_tests.py deleted file mode 100644 index 792da5ae09c..00000000000 --- a/.github/e2e-stack/select_tests.py +++ /dev/null @@ -1,52 +0,0 @@ -import re -import sys -from typing import Final - -SELECTABLE: Final = re.compile(r"^tests/e2e/([A-Za-z0-9_.-]+/)*test_[A-Za-z0-9_.-]+\.py$") -UNSUPPORTED: Final = re.compile( - r"^tests/e2e/(ui|claude_code|load|migrations)/" - r"|^tests/e2e/mcp/test_mcp_oauth_happy_path_e2e\.py$" - r"|^tests/e2e/llm_translation/realtime/test_realtime_pipecat_audio_e2e\.py$" - r"|^tests/e2e/batches/test_managed_files_enforcement_e2e\.py$" - r"|^tests/e2e/guardrails/test_presidio_masking_e2e\.py$" - r"|^tests/e2e/logging/test_otel_v2_langfuse_generation_output_e2e\.py$" - r"|^tests/e2e/logging/test_langsmith_batch_serialization_e2e\.py$" - r"|^tests/e2e/logging/test_s3_log_e2e\.py$" - r"|^tests/e2e/secret_manager/" -) -HARNESS: Final = re.compile( - r"^tests/e2e/[A-Za-z0-9_.-]+\.(py|ini)$" - r"|^tests/e2e/idp_realm\.json$" - r"|^tests/e2e/management/(management_client|jwt_actors|conftest)\.py$" - r"|^tests/e2e/coverage_registry/management_cases\.py$" - r"|^tests/e2e/gateway/" - r"|^\.github/e2e-stack/" - r"|^\.github/workflows/test-e2e-changed\.yml$" -) -UNEXPANDED: Final = re.compile(r"[*?\[]") - - -def is_selectable(path: str) -> bool: - return SELECTABLE.match(path) is not None and UNSUPPORTED.match(path) is None - - -def select(changed: tuple[str, ...], canary: tuple[str, ...]) -> tuple[str, ...]: - direct: Final = frozenset(path for path in changed if is_selectable(path)) - harness_changed: Final = any(HARNESS.match(path) for path in changed) - canary_tests: Final = frozenset(path for path in canary if harness_changed and is_selectable(path)) - return tuple(sorted(direct | canary_tests)) - - -def main() -> int: - canary: Final = tuple(sys.argv[1:]) - unexpanded: Final = tuple(path for path in canary if UNEXPANDED.search(path)) - if unexpanded: - _ = sys.stderr.write(f"the canary paths reached the selector unexpanded: {' '.join(unexpanded)}\n") - return 1 - changed: Final = tuple(line.strip() for line in sys.stdin if line.strip()) - _ = sys.stdout.write(" ".join(select(changed, canary)) + "\n") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/.github/e2e-stack/up.sh b/.github/e2e-stack/up.sh deleted file mode 100755 index 931b5eb2169..00000000000 --- a/.github/e2e-stack/up.sh +++ /dev/null @@ -1,231 +0,0 @@ -#!/usr/bin/env bash -set -euo pipefail -umask 077 - -REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)" -STACK_DIR="${E2E_STACK_DIR:-${RUNNER_TEMP:-/tmp}/litellm-e2e-stack}" -CERTS_DIR="${STACK_DIR}/certs" -LOGS_DIR="${STACK_DIR}/logs" -PIDS_DIR="${STACK_DIR}/pids" - -POSTGRES_IMAGE="${E2E_POSTGRES_IMAGE:-postgres:16.6}" -VALKEY_IMAGE="${E2E_VALKEY_IMAGE:-valkey/valkey:8.1.4@sha256:81db6d39e1bba3b3ff32bd3a1b19a6d69690f94a3954ec131277b9a26b95b3aa}" -JAEGER_IMAGE="${E2E_JAEGER_IMAGE:-jaegertracing/jaeger:2.10.0}" -NGINX_IMAGE="${E2E_NGINX_IMAGE:-nginx:1.29.1-alpine@sha256:42a516af16b852e33b7682d5ef8acbd5d13fe08fecadc7ed98605ba5e3b26ab8}" - -LB_PORT="${E2E_LB_PORT:-4000}" -GATEWAY_PORT_1="${E2E_GATEWAY_PORT_1:-4010}" -GATEWAY_PORT_2="${E2E_GATEWAY_PORT_2:-4011}" -BACKEND_PORT="${E2E_BACKEND_PORT:-4001}" -REDIS_PORT="${E2E_REDIS_PORT:-6379}" -DATABASE_HOST="${E2E_DATABASE_HOST:-127.0.0.1}" -DATABASE_PORT="${E2E_DATABASE_PORT:-5432}" -DATABASE_USER="${E2E_DATABASE_USER:-litellm}" -DATABASE_PASSWORD="${E2E_DATABASE_PASSWORD:-dbpassword9090}" -DATABASE_NAME="${E2E_DATABASE_NAME:-litellm}" -JAEGER_OTLP_PORT="${E2E_JAEGER_OTLP_PORT:-4318}" -JAEGER_OTLP_TLS_PORT="${E2E_JAEGER_OTLP_TLS_PORT:-4319}" -JAEGER_QUERY_PORT="${E2E_JAEGER_QUERY_PORT:-16686}" -KEYCLOAK_PORT="${E2E_KEYCLOAK_PORT:-8081}" - -MASTER_KEY="${LITELLM_MASTER_KEY:-sk-e2e-$(openssl rand -hex 16)}" - -mkdir -p "${CERTS_DIR}" "${LOGS_DIR}" "${PIDS_DIR}" -chmod 700 "${STACK_DIR}" "${LOGS_DIR}" "${PIDS_DIR}" -chmod 755 "${CERTS_DIR}" - -log() { printf 'e2e-stack: %s\n' "$*"; } - -port_open() { (exec 3<>"/dev/tcp/127.0.0.1/$1") 2>/dev/null; } - -wait_for() { - local label="$1" check="$2" deadline=$((SECONDS + ${3:-120})) - until eval "${check}"; do - if ((SECONDS >= deadline)); then - log "timed out waiting for ${label}" - exit 1 - fi - sleep 2 - done - log "${label} is up" -} - -if [[ -f "${REPO_ROOT}/tests/e2e/.env" ]]; then - set -a - source "${REPO_ROOT}/tests/e2e/.env" - set +a -fi - -if [[ -z "${DD_API_KEY:-}" ]]; then - log "DD_API_KEY is empty; the gateway config enables the datadog callback, so put a Datadog API key in tests/e2e/.env" - exit 1 -fi -export DD_SITE="${DD_SITE:-datadoghq.com}" - -if ! port_open "${DATABASE_PORT}"; then - docker run -d --name e2e-postgres -p "${DATABASE_PORT}:5432" \ - -e "POSTGRES_USER=${DATABASE_USER}" -e "POSTGRES_PASSWORD=${DATABASE_PASSWORD}" -e "POSTGRES_DB=${DATABASE_NAME}" \ - "${POSTGRES_IMAGE}" >/dev/null -fi -wait_for "postgres" "port_open ${DATABASE_PORT}" - -if ! port_open "${JAEGER_QUERY_PORT}"; then - docker run -d --name e2e-jaeger -p "${JAEGER_OTLP_PORT}:4318" -p "${JAEGER_QUERY_PORT}:16686" \ - "${JAEGER_IMAGE}" >/dev/null -fi -wait_for "jaeger" "curl -fs http://127.0.0.1:${JAEGER_QUERY_PORT}/api/services >/dev/null" - -openssl genrsa -out "${CERTS_DIR}/ca.key" 2048 2>/dev/null -openssl req -x509 -new -nodes -key "${CERTS_DIR}/ca.key" -sha256 -days 7 \ - -subj "/CN=litellm-e2e-ca" \ - -addext "basicConstraints=critical,CA:TRUE" -addext "keyUsage=critical,keyCertSign,cRLSign" \ - -out "${CERTS_DIR}/ca.crt" 2>/dev/null -openssl genrsa -out "${CERTS_DIR}/server.key" 2048 2>/dev/null -openssl req -new -key "${CERTS_DIR}/server.key" -subj "/CN=localhost" -out "${CERTS_DIR}/server.csr" 2>/dev/null -openssl x509 -req -in "${CERTS_DIR}/server.csr" -CA "${CERTS_DIR}/ca.crt" -CAkey "${CERTS_DIR}/ca.key" \ - -CAcreateserial -days 7 -sha256 \ - -extfile <(printf 'basicConstraints=CA:FALSE\nkeyUsage=critical,digitalSignature,keyEncipherment\nextendedKeyUsage=serverAuth\nsubjectAltName=DNS:localhost,IP:127.0.0.1\n') \ - -out "${CERTS_DIR}/server.crt" 2>/dev/null -chmod 644 "${CERTS_DIR}"/*.key "${CERTS_DIR}"/*.crt - -CERTIFI_BUNDLE="$(cd "${REPO_ROOT}" && uv run --no-sync python -c 'import certifi; print(certifi.where())')" -cat "${CERTIFI_BUNDLE}" "${CERTS_DIR}/ca.crt" > "${CERTS_DIR}/ca-bundle.pem" - -docker rm -f e2e-valkey >/dev/null 2>&1 || true -docker run -d --name e2e-valkey -p "${REDIS_PORT}:${REDIS_PORT}" -v "${CERTS_DIR}:/certs:ro" \ - "${VALKEY_IMAGE}" valkey-server \ - --cluster-enabled yes --port 0 --tls-port "${REDIS_PORT}" \ - --tls-cert-file /certs/server.crt --tls-key-file /certs/server.key --tls-ca-cert-file /certs/ca.crt \ - --tls-auth-clients no --cluster-announce-ip 127.0.0.1 >/dev/null -VALKEY_CLI="docker exec e2e-valkey valkey-cli --tls --cacert /certs/ca.crt -h 127.0.0.1 -p ${REDIS_PORT}" -wait_for "valkey" "${VALKEY_CLI} ping 2>/dev/null | grep -q PONG" -${VALKEY_CLI} cluster addslotsrange 0 16383 >/dev/null -wait_for "valkey cluster" "${VALKEY_CLI} cluster info 2>/dev/null | grep -q cluster_state:ok" - -CONFIG_SOURCE="${REPO_ROOT}/tests/e2e/gateway/stage_mirror_ci_config.yml" -CONFIG_PATH="${CONFIG_SOURCE}" -if [[ "${REDIS_PORT}" != "6379" ]]; then - CONFIG_PATH="${STACK_DIR}/litellm-config.yml" - sed "s/port: 6379/port: ${REDIS_PORT}/" "${CONFIG_SOURCE}" > "${CONFIG_PATH}" -fi - -SERVER_ENV=( - "LITELLM_MASTER_KEY=${MASTER_KEY}" - "DATABASE_HOST=${DATABASE_HOST}" - "DATABASE_PORT=${DATABASE_PORT}" - "DATABASE_USER=${DATABASE_USER}" - "DATABASE_PASSWORD=${DATABASE_PASSWORD}" - "DATABASE_NAME=${DATABASE_NAME}" - "DISABLE_SCHEMA_UPDATE=true" - "REDIS_HOST=127.0.0.1" - "REDIS_PORT=${REDIS_PORT}" - "REDIS_CLUSTER_NODES=[{\"host\":\"127.0.0.1\",\"port\":${REDIS_PORT}}]" - "CONFIG_FILE_PATH=${CONFIG_PATH}" - "STORE_MODEL_IN_DB=True" - "OTEL_EXPORTER_OTLP_PROTOCOL=http/protobuf" - "OTEL_EXPORTER_OTLP_ENDPOINT=https://127.0.0.1:${JAEGER_OTLP_TLS_PORT}" - "SSL_CERT_FILE=${CERTS_DIR}/ca-bundle.pem" - "PYTHONPATH=${REPO_ROOT}" - "JWT_PUBLIC_KEY_URL=http://127.0.0.1:${KEYCLOAK_PORT}/realms/litellm-e2e/protocol/openid-connect/certs" - "JWT_ISSUER=http://127.0.0.1:${KEYCLOAK_PORT}/realms/litellm-e2e" - "JWT_AUDIENCE=litellm-e2e" -) -if [[ -n "${VERTEXAI_CREDENTIALS:-}" ]]; then - printf '%s' "${VERTEXAI_CREDENTIALS}" > "${STACK_DIR}/vertex-adc.json" - SERVER_ENV+=("GOOGLE_APPLICATION_CREDENTIALS=${STACK_DIR}/vertex-adc.json") -fi - -cd "${REPO_ROOT}" - -env "${SERVER_ENV[@]}" "E2E_KEYCLOAK_PORT=${KEYCLOAK_PORT}" bash .github/e2e-stack/start-idp.sh - -log "running migrations" -env "${SERVER_ENV[@]}" uv run --no-sync python migrations/run.py >"${LOGS_DIR}/migrations.log" 2>&1 - -start_server() { - local name="$1"; shift - env -u AWS_ROLE_NAME "${SERVER_ENV[@]}" "$@" >"${LOGS_DIR}/${name}.log" 2>&1 & - echo $! > "${PIDS_DIR}/${name}.pid" -} - -if [[ "$(uname)" == "Linux" ]]; then - NGINX_UPSTREAM_HOST=127.0.0.1 - NGINX_DOCKER_ARGS=(--network host) -else - NGINX_UPSTREAM_HOST=host.docker.internal - NGINX_DOCKER_ARGS=(-p "${LB_PORT}:${LB_PORT}" -p "${JAEGER_OTLP_TLS_PORT}:${JAEGER_OTLP_TLS_PORT}") -fi - -cat > "${STACK_DIR}/nginx.conf" </dev/null 2>&1 || true -docker run -d --name e2e-nginx "${NGINX_DOCKER_ARGS[@]}" \ - -v "${STACK_DIR}/nginx.conf:/etc/nginx/nginx.conf:ro" \ - -v "${CERTS_DIR}:/certs:ro" "${NGINX_IMAGE}" >/dev/null - -wait_for "Jaeger OTLP TLS listener" \ - "curl -sS --cacert ${CERTS_DIR}/ca.crt https://127.0.0.1:${JAEGER_OTLP_TLS_PORT}/ -o /dev/null -w '%{http_code}' | grep -qE '^[2345]'" - -start_server backend uv run --no-sync uvicorn backend.main:app --host 0.0.0.0 --port "${BACKEND_PORT}" -start_server gateway-1 uv run --no-sync uvicorn gateway.main:app --workers 1 --host 0.0.0.0 --port "${GATEWAY_PORT_1}" -start_server gateway-2 uv run --no-sync uvicorn gateway.main:app --workers 1 --host 0.0.0.0 --port "${GATEWAY_PORT_2}" - -wait_for "backend" "curl -fs http://127.0.0.1:${BACKEND_PORT}/health/liveliness >/dev/null" 300 -wait_for "gateway-1" "curl -fs http://127.0.0.1:${GATEWAY_PORT_1}/health/liveliness >/dev/null" 300 -wait_for "gateway-2" "curl -fs http://127.0.0.1:${GATEWAY_PORT_2}/health/liveliness >/dev/null" 300 -wait_for "load balancer" "curl -fs http://127.0.0.1:${LB_PORT}/health/liveliness >/dev/null" 60 - -cat > "${STACK_DIR}/stack.env" < bool: + return any(_token_covers(token, relative_path) for token in self.included) and not any( + _token_covers(token, relative_path) for token in self.ignored + ) + + @dataclass(frozen=True, slots=True) class Finding: subject: str @@ -111,44 +122,32 @@ def _uncommented(value: str) -> str: return COMMENT_RE.sub("", value) -def _invoked_test_tokens(scalars: Iterable[Scalar]) -> frozenset[str]: - return frozenset( - match.group(0).rstrip("/") +def _selection_for_scalar(scalar: Scalar) -> Selection: + text: Final = _uncommented(scalar.value) + ignored: Final = frozenset().union( + *( + frozenset(token.rstrip("/") for token in TEST_TOKEN_RE.findall(ignored_argument)) + for ignored_argument in IGNORE_ARG_RE.findall(text) + ) + ) + included: Final = frozenset( + match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(IGNORE_ARG_RE.sub("", text)) + ) + return Selection(included=included, ignored=ignored) + + +def _invoked_selections(scalars: Iterable[Scalar]) -> tuple[Selection, ...]: + selected_scalars: Final = tuple( + scalar for scalar in scalars if scalar.key in TEST_PATH_KEYS or TEST_RUNNER_RE.search(scalar.value) - for match in TEST_TOKEN_RE.finditer(_uncommented(scalar.value)) ) + selections: Final = tuple(_selection_for_scalar(scalar) for scalar in selected_scalars) + return tuple(selection for selection in selections if selection.included) -SELECTION_ARM_RE = re.compile(r"(?ms)^\s*([A-Za-z0-9_|*-]+)\)\s*(.*?);;") - - -def _unit_selection_arms(repo_root: pathlib.Path = REPO_ROOT) -> Mapping[str, frozenset[str]]: - script: Final = repo_root / ".circleci/scripts/unit_selection.sh" - if not script.is_file(): - return MappingProxyType({}) - text: Final = _uncommented(script.read_text()) - return MappingProxyType( - { - label: frozenset(match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(body)) - for label, body in SELECTION_ARM_RE.findall(text) - } - ) - - -def _unit_selection_tokens(repo_root: pathlib.Path = REPO_ROOT) -> frozenset[str]: - return frozenset(token for tokens in _unit_selection_arms(repo_root).values() for token in tokens) - - -def _wired_unit_flags(scalars: Iterable[Scalar]) -> frozenset[str]: - return frozenset(scalar.value for scalar in scalars if scalar.key == "unit-flag" and "${{" not in scalar.value) - - -def _shard_tokens(scalars: Iterable[Scalar], arms: Mapping[str, frozenset[str]]) -> frozenset[str]: - wired: Final = _wired_unit_flags(scalars) - return _invoked_test_tokens(scalars) | frozenset( - token for label, tokens in arms.items() if label in wired for token in tokens - ) +def _invoked_test_tokens(scalars: Iterable[Scalar]) -> frozenset[str]: + return frozenset().union(*(selection.included for selection in _invoked_selections(scalars))) def _built_dockerfile_tokens(scalars: Iterable[Scalar]) -> frozenset[str]: @@ -211,11 +210,12 @@ def _dockerfiles() -> tuple[str, ...]: ) -def _uncovered_tests(allowlist: Allowlist, tokens: frozenset[str]) -> tuple[Finding, ...]: +def _uncovered_tests(allowlist: Allowlist, selections: tuple[Selection, ...]) -> tuple[Finding, ...]: uncovered = tuple( relative_path for relative_path in _test_files() - if not any(_token_covers(token, relative_path) for token in tokens) and not allowlist.covers_test(relative_path) + if not any(selection.covers(relative_path) for selection in selections) + and not allowlist.covers_test(relative_path) ) directories = tuple(dict.fromkeys(path.rsplit("/", 1)[0] for path in uncovered)) return tuple( @@ -505,7 +505,7 @@ def _check_slices() -> int: def _check_shards() -> int: - findings = _unassigned_shard_children(_shard_tokens(_all_scalars(), _unit_selection_arms())) + findings = _unassigned_shard_children(_invoked_test_tokens(_all_scalars())) if findings: _report( "test directories and files that no shard claims", @@ -568,6 +568,9 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens browser_paths: Final = frozenset(node.split("::", 1)[0] for node in browser_nodes) circle_path: Final = repo_root / ".circleci/config.yml" circle: Final = yaml.safe_load(circle_path.read_text()) if circle_path.exists() else {} + circle_test_path_tokens: Final = _invoked_test_tokens( + scalar for scalar in _scalars(circle, "config.yml") if scalar.key == "test_path" + ) steps: Final = circle.get("jobs", {}).get("integration_contracts", {}).get("steps", ()) invoked: Final = any( ".circleci/scripts/run_integration.sh" in scalar.value @@ -609,9 +612,9 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens if any(_token_covers(token, path) for token in gha_tokens) ) + tuple( - Finding(path, "GitHub-owned integration contract has no invoking workflow") + Finding(path, "GitHub-owned integration contract has no invoking job") for path in sorted(github_files) - if not any(_token_covers(token, path) for token in gha_tokens) + if not any(_token_covers(token, path) for token in gha_tokens | circle_test_path_tokens) ) + tuple( Finding(path, "GitHub-owned integration file is missing") @@ -675,7 +678,11 @@ def main() -> int: integration_paths, ownership_findings = _integration_ownership() test_findings = ( - _uncovered_tests(allowlist, _invoked_test_tokens(scalars) | _unit_selection_tokens() | integration_paths) + _uncovered_tests( + allowlist, + _invoked_selections(scalars) + + (Selection(included=integration_paths, ignored=frozenset()),), + ) + ownership_findings ) dockerfile_findings = _uncovered_dockerfiles(allowlist, _built_dockerfile_tokens(scalars)) diff --git a/.github/scripts/verify_linux_native_wheel.py b/.github/scripts/verify_linux_native_wheel.py index 465918f5a81..4c4cb399dce 100644 --- a/.github/scripts/verify_linux_native_wheel.py +++ b/.github/scripts/verify_linux_native_wheel.py @@ -214,7 +214,7 @@ def main( native_module: Final = load_native_module(native_path) native_module_loads: Final = native_module is not None panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test") - native_size_limit: Final = 45_000_000 + native_size_limit: Final = 48_000_000 native_size_within_limit: Final = native_member.file_size <= native_size_limit validations: Final = ( (f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG), diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index 6d67bef44cb..7db01d721cf 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -13,16 +13,6 @@ on: have its path existence-checked like any other token. required: true type: string - unit-flag: - description: >- - Codecov flag of the `.circleci/tests.yml` job that now owns part of - this shard. The shard also runs the files - `.circleci/scripts/unit_selection.sh` lists for the flag, on every - event, because the CircleCI pipeline is manual-only while the tests - migrate. - required: false - type: string - default: "" workers: description: "Number of pytest-xdist workers" required: false @@ -79,12 +69,6 @@ on: description: "Unique name for the coverage artifact (must be unique per run)" required: true type: string - legacy-mcp-peer: - description: "Install the isolated SDK1 peer for MCP compatibility tests" - required: false - type: boolean - default: false - permissions: contents: read @@ -142,17 +126,10 @@ jobs: - name: Install dependencies if: steps.changes.outputs.decision != 'skip' timeout-minutes: 8 - env: - LEGACY_MCP_PEER: ${{ inputs.legacy-mcp-peer }} run: | diff -u model_prices_and_context_window.json litellm/model_prices_and_context_window_backup.json - .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml + .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 uv run --no-sync python -c 'import os, sys; print(sys.version); assert f"{sys.version_info.major}.{sys.version_info.minor}" == os.environ["UV_PYTHON"]' - if [ "$LEGACY_MCP_PEER" = "true" ]; then - uv venv --python "${UV_PYTHON}" .venv-mcp-peer - uv pip install --python .venv-mcp-peer 'mcp==1.28.1' 'langchain-mcp-adapters==0.2.1' - echo "MCP_TEST_PEER_PYTHON=$GITHUB_WORKSPACE/.venv-mcp-peer/bin/python" >> "$GITHUB_ENV" - fi - name: Cache Prisma binaries if: steps.changes.outputs.decision != 'skip' @@ -171,7 +148,6 @@ jobs: timeout-minutes: ${{ inputs.timeout-minutes }} env: TEST_PATH: ${{ inputs.test-path }} - UNIT_FLAG: ${{ inputs.unit-flag }} MAX_FAILURES: ${{ inputs.max-failures }} WORKERS: ${{ inputs.workers }} RERUNS: ${{ inputs.reruns }} @@ -181,9 +157,6 @@ jobs: run: | echo "has-coverage=false" >> "$GITHUB_OUTPUT" selection="${TEST_PATH}" - if [ -n "${UNIT_FLAG}" ]; then - selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${UNIT_FLAG}" | tr '\n' ' ')" - fi if [ -z "${selection// /}" ]; then echo "shard selection is empty; nothing to run" exit 0 diff --git a/.github/workflows/test-e2e-changed.yml b/.github/workflows/test-e2e-changed.yml deleted file mode 100644 index 228e23f60d7..00000000000 --- a/.github/workflows/test-e2e-changed.yml +++ /dev/null @@ -1,265 +0,0 @@ -name: e2e-changed-tests - -on: - pull_request: - -concurrency: - group: e2e-changed-${{ github.event.pull_request.number }} - cancel-in-progress: true - -permissions: {} - -jobs: - detect: - name: Detect changed e2e tests - runs-on: ubuntu-latest - timeout-minutes: 5 - permissions: - contents: read - pull-requests: read - outputs: - tests: ${{ steps.changed.outputs.tests }} - any: ${{ steps.changed.outputs.any }} - steps: - - name: Checkout the selector and the canary suite - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - sparse-checkout: | - .github/e2e-stack - tests/e2e/access_control - tests/e2e/management/test_jwt_management_e2e.py - tests/e2e/other/test_jwt_auth_e2e.py - persist-credentials: false - ref: ${{ github.sha }} - - - name: List the e2e test files this PR added or modified - id: changed - env: - GH_TOKEN: ${{ github.token }} - REPO: ${{ github.repository }} - PR_NUMBER: ${{ github.event.pull_request.number }} - HEAD_SHA: ${{ github.event.pull_request.head.sha }} - run: | - gh api "repos/${REPO}/pulls/${PR_NUMBER}" \ - --jq 'select(.head.sha == env.HEAD_SHA and .changed_files < 3000) | .head.sha' \ - | grep -Fxq "${HEAD_SHA}" - files="$(gh api "repos/${REPO}/pulls/${PR_NUMBER}/files" --paginate \ - --jq '.[] | select(.status != "removed") | .filename')" - gh api "repos/${REPO}/pulls/${PR_NUMBER}" --jq '.head.sha' | grep -Fxq "${HEAD_SHA}" - tests="$(printf '%s\n' "${files}" \ - | python3 .github/e2e-stack/select_tests.py tests/e2e/access_control/test_*.py \ - tests/e2e/management/test_jwt_management_e2e.py tests/e2e/other/test_jwt_auth_e2e.py)" - echo "tests=${tests}" >> "${GITHUB_OUTPUT}" - if [ -n "${tests}" ]; then - echo "any=true" >> "${GITHUB_OUTPUT}" - echo "selected e2e tests: ${tests}" - else - echo "any=false" >> "${GITHUB_OUTPUT}" - echo "no changed e2e test files supported by this stack; nothing to run" - fi - - run: - name: Run changed e2e tests against the stage-mirror stack - needs: detect - if: needs.detect.outputs.any == 'true' && github.event.pull_request.head.repo.full_name == github.repository - runs-on: ubuntu-latest - timeout-minutes: 90 - environment: e2e-changed - permissions: - contents: read - id-token: write - services: - postgres: - image: postgres:16.6 - env: - POSTGRES_USER: litellm - POSTGRES_PASSWORD: dbpassword9090 - POSTGRES_DB: litellm - ports: - - 5432:5432 - options: >- - --health-cmd "pg_isready -U litellm" - --health-interval 5s - --health-timeout 5s - --health-retries 10 - jaeger: - image: jaegertracing/jaeger:2.10.0 - ports: - - 4318:4318 - - 16686:16686 - steps: - - name: Validate configuration - env: - ROLE: ${{ vars.E2E_AWS_ROLE_TO_ASSUME }} - run: test -n "${ROLE}" || { echo "::error::Set repo variable E2E_AWS_ROLE_TO_ASSUME to an OIDC role with read access to the e2e secrets"; exit 1; } - - - name: Checkout - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - persist-credentials: false - ref: ${{ github.sha }} - - - name: Set up Python - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 - with: - python-version: "3.13" - - - 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: Install dependencies - run: | - .github/scripts/uv_sync_with_retries.sh --frozen \ - --extra proxy --extra proxy-runtime --extra extra_proxy \ - --extra semantic-router --extra bedrock-realtime \ - --group ci --group proxy-dev --group e2e-dev - uv pip install "pipecat-ai[openai]==1.4.0" - - - name: Cache Prisma binaries - uses: ./.github/actions/cache-prisma-binaries - - - name: Generate Prisma client - run: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma - - - name: Install Playwright chromium - run: uv run --no-sync playwright install --with-deps chromium - - - name: Configure AWS credentials - id: aws - uses: aws-actions/configure-aws-credentials@e7f100cf4c008499ea8adda475de1042d6975c7b # v6.2.0 - with: - role-to-assume: ${{ vars.E2E_AWS_ROLE_TO_ASSUME }} - aws-region: us-east-1 - role-session-name: litellm-e2e-changed-${{ github.run_id }} - role-duration-seconds: 900 - output-env-credentials: false - output-credentials: true - - - name: Fetch provider credentials from AWS Secrets Manager - env: - AWS_ACCESS_KEY_ID: ${{ steps.aws.outputs.aws-access-key-id }} - AWS_SECRET_ACCESS_KEY: ${{ steps.aws.outputs.aws-secret-access-key }} - AWS_SESSION_TOKEN: ${{ steps.aws.outputs.aws-session-token }} - AWS_DEFAULT_REGION: us-east-1 - run: | - umask 077 - aws secretsmanager get-secret-value --secret-id litellm-e2e-changed-provider-keys \ - --query SecretString --output text \ - | uv run --no-sync python .github/e2e-stack/secrets_to_env.py tests/e2e/.env - aws secretsmanager get-secret-value --secret-id litellm-e2e-changed-license \ - --query SecretString --output text \ - | jq -R -s '{"LITELLM_LICENSE": .}' \ - | uv run --no-sync python .github/e2e-stack/secrets_to_env.py tests/e2e/.env - - - name: Boot the stage-mirror stack - id: boot - run: | - umask 077 - if ! bash .github/e2e-stack/up.sh > "${RUNNER_TEMP}/e2e-boot.log" 2>&1; then - echo "::error::stage-mirror stack failed to boot; raw logs are not published" - exit 1 - fi - - - name: Export stack environment - run: | - master_key="$(grep '^LITELLM_MASTER_KEY=' "${RUNNER_TEMP}/litellm-e2e-stack/stack.env" | cut -d= -f2-)" - echo "::add-mask::${master_key}" - cat "${RUNNER_TEMP}/litellm-e2e-stack/stack.env" >> "${GITHUB_ENV}" - - - name: Run the selected tests three times - env: - TESTS: ${{ needs.detect.outputs.tests }} - E2E_FIXTURE_MODE: live - E2E_PROVIDER_EDGE_HOST_REACHABLE: '1' - E2E_OWNED_GATEWAY: '1' - COLUMNS: '400' - run: | - umask 077 - read -r -a test_files <<< "${TESTS}" - for pass in 1 2 3; do - report="${RUNNER_TEMP}/e2e-pass-${pass}.xml" - log="${RUNNER_TEMP}/e2e-pass-${pass}.log" - echo "::group::pass ${pass} of 3" - set +e - uv run --no-sync pytest "${test_files[@]}" --rootdir=. -v --reruns 0 -p no:cacheprovider \ - -o junit_family=xunit1 --junitxml="${report}" > "${log}" 2>&1 - status=$? - uv run --no-sync python .github/e2e-stack/assert_tests_ran.py "${report}" "${test_files[@]}" - verified=$? - set -e - grep -E '^(FAILED|ERROR) ' "${log}" || true - grep -E '^=+ .* in [0-9.]+s( \([0-9:]+\))? =+$' "${log}" | tail -n 1 - echo "::endgroup::" - if [ "${status}" = "5" ]; then - echo "::error::the selected files collected no runnable tests, so nothing was verified" - exit 1 - fi - if [ "${status}" != "0" ]; then - echo "::error::pass ${pass} of 3 failed with exit code ${status}" - exit "${status}" - fi - if [ "${verified}" != "0" ]; then - echo "::error::pass ${pass} of 3 did not verify every selected file" - exit 1 - fi - echo "pass ${pass} of 3 passed" - done - - - name: Redact the pytest output - if: always() && steps.boot.outcome == 'success' - run: | - umask 077 - shopt -s nullglob - uv run --no-sync python .github/e2e-stack/redact_output.py \ - --values tests/e2e/.env --values "${RUNNER_TEMP}/litellm-e2e-stack/stack.env" \ - --out "${RUNNER_TEMP}/e2e-redacted" "${RUNNER_TEMP}"/e2e-pass-*.log "${RUNNER_TEMP}"/e2e-pass-*.xml - - - name: Keep the redacted pytest output - if: always() && steps.boot.outcome == 'success' - uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1 - with: - name: e2e-changed-pytest-output-${{ github.run_attempt }} - path: ${{ runner.temp }}/e2e-redacted - retention-days: 14 - if-no-files-found: ignore - - - name: Stop the stack - if: always() && steps.boot.outcome != 'skipped' - run: bash .github/e2e-stack/down.sh - - - name: Remove credentials and raw output - if: always() - run: | - rm -f tests/e2e/.env "${RUNNER_TEMP}/e2e-boot.log" "${RUNNER_TEMP}"/e2e-pass-*.log "${RUNNER_TEMP}"/e2e-pass-*.xml - rm -rf "${RUNNER_TEMP}/litellm-e2e-stack" "${RUNNER_TEMP}/e2e-redacted" - - gate: - name: e2e-changed-tests - needs: [detect, run] - if: always() - runs-on: ubuntu-latest - timeout-minutes: 5 - steps: - - name: Require three successful passes when tests changed - env: - DETECT_RESULT: ${{ needs.detect.result }} - ANY_TESTS: ${{ needs.detect.outputs.any }} - RUN_RESULT: ${{ needs.run.result }} - run: | - if [ "${DETECT_RESULT}" != "success" ]; then - echo "::error::changed-test detection did not succeed" - exit 1 - fi - if [ "${ANY_TESTS}" = "false" ]; then - echo "no changed e2e test files supported by this stack; nothing to run" - exit 0 - fi - if [ "${ANY_TESTS}" != "true" ] || [ "${RUN_RESULT}" != "success" ]; then - echo "::error::selected e2e tests require an approved, successful run; fork PRs must run from a reviewed same-repository branch" - exit 1 - fi diff --git a/.github/workflows/test-mcp-dependency-resolution.yml b/.github/workflows/test-mcp-dependency-resolution.yml index 8f70375a181..1772bdeabc6 100644 --- a/.github/workflows/test-mcp-dependency-resolution.yml +++ b/.github/workflows/test-mcp-dependency-resolution.yml @@ -5,6 +5,27 @@ on: branches: - main - "litellm_**" + paths: + - "**/pyproject.toml" + - "uv.lock" + - "uv.toml" + - ".python-version" + - "rust-toolchain.toml" + - "litellm-rust/**" + - "litellm/__init__.py" + - "litellm/proxy/proxy_server.py" + - "litellm/**/*mcp*" + - "litellm/**/*mcp*/**" + - "litellm/integrations/arize/**" + - "scripts/check_mcp_sdk_install.py" + - "tests/base_sdk_tests/**" + - ".github/workflows/test-mcp-dependency-resolution.yml" + - ".github/actions/detect-changes/**" + - ".github/actions/setup-uv-with-retries/**" + - ".github/actions/cache-cargo-build/**" + - ".github/scripts/detect_changes.sh" + - ".github/scripts/uv_sync_with_retries.sh" + - ".circleci/scripts/classify_changes.sh" permissions: contents: read diff --git a/.github/workflows/test-postgres.yml b/.github/workflows/test-postgres.yml deleted file mode 100644 index a1c639acbeb..00000000000 --- a/.github/workflows/test-postgres.yml +++ /dev/null @@ -1,185 +0,0 @@ -name: "Postgres Tests" - -on: - pull_request: - branches: - - main - - "litellm_**" - push: - branches: - - main - workflow_dispatch: - -permissions: - contents: read - -concurrency: - group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }} - cancel-in-progress: ${{ github.event_name == 'pull_request' }} - -jobs: - postgres: - name: ${{ matrix.shard }} - runs-on: ubuntu-latest - timeout-minutes: ${{ matrix.job-timeout-minutes }} - permissions: - contents: read - id-token: write - - services: - postgres: - image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5 - env: - POSTGRES_USER: postgres - POSTGRES_PASSWORD: postgres - POSTGRES_DB: litellm_test - ports: - - 5432:5432 - options: >- - --health-cmd pg_isready - --health-interval 10s - --health-timeout 5s - --health-retries 10 - - strategy: - fail-fast: false - matrix: - include: - - shard: roi-database - test-path: "tests/integration/database/test_roi_observed.py" - seed: none - workers: 0 - timeout-minutes: 10 - job-timeout-minutes: 35 - - - shard: proxy-behavior - test-path: "tests/proxy_behavior" - seed: db-push - workers: 0 - timeout-minutes: 25 - job-timeout-minutes: 50 - - - shard: proxy-security - test-path: "tests/proxy_security_tests" - seed: db-push - workers: 0 - timeout-minutes: 15 - job-timeout-minutes: 40 - - - shard: schema-migration - test-path: "tests/proxy_migration_tests" - seed: none - workers: 0 - timeout-minutes: 20 - job-timeout-minutes: 45 - - env: - DATABASE_URL: "postgresql://postgres:postgres@localhost:5432/litellm_test" - - steps: - - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - timeout-minutes: 3 - with: - persist-credentials: false - - - name: Detect relevant changes - id: changes - timeout-minutes: 2 - uses: ./.github/actions/detect-changes - - - name: Set up Python - if: steps.changes.outputs.decision != 'skip' - timeout-minutes: 3 - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 - with: - python-version: "3.12" - - - name: Set up uv - if: steps.changes.outputs.decision != 'skip' - timeout-minutes: 3 - uses: ./.github/actions/setup-uv-with-retries - with: - version: "0.10.9" - - - name: Cache uv dependencies - if: steps.changes.outputs.decision != 'skip' && github.ref == 'refs/heads/main' - timeout-minutes: 5 - uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0 - with: - path: | - ~/.cache/uv - .venv - key: ${{ runner.os }}-uv-postgres-${{ hashFiles('uv.lock') }} - restore-keys: | - ${{ runner.os }}-uv-postgres- - - - name: Cache uv dependencies - if: steps.changes.outputs.decision != 'skip' && github.ref != 'refs/heads/main' - timeout-minutes: 5 - uses: actions/cache/restore@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0 - with: - path: | - ~/.cache/uv - .venv - key: ${{ runner.os }}-uv-postgres-${{ hashFiles('uv.lock') }} - restore-keys: | - ${{ runner.os }}-uv-postgres- - - - name: Install dependencies - if: steps.changes.outputs.decision != 'skip' - timeout-minutes: 12 - run: | - .github/scripts/uv_sync_with_retries.sh --frozen --all-groups --all-extras - - - name: Cache Prisma binaries - if: steps.changes.outputs.decision != 'skip' - timeout-minutes: 3 - uses: ./.github/actions/cache-prisma-binaries - - - name: Generate Prisma client - if: steps.changes.outputs.decision != 'skip' - timeout-minutes: 5 - run: | - uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma - - - name: Seed database schema - if: steps.changes.outputs.decision != 'skip' && matrix.seed != 'none' - timeout-minutes: 10 - run: | - uv run --no-sync prisma db push --schema litellm/proxy/schema.prisma --accept-data-loss - - - name: Run tests - if: steps.changes.outputs.decision != 'skip' - timeout-minutes: ${{ matrix.timeout-minutes }} - env: - TEST_PATH: ${{ matrix.test-path }} - WORKERS: ${{ matrix.workers }} - PYTEST_ADDOPTS: ${{ matrix.shard == 'proxy-behavior' && '--cov=./litellm --cov-report=xml:coverage-lens-postgres.xml' || matrix.shard == 'roi-database' && '--cov=./litellm --cov-report=xml:coverage-roi-postgres.xml' || '' }} - run: | - if [ "${WORKERS}" = "0" ]; then - uv run --no-sync pytest ${TEST_PATH:?} -vv --tb=short --durations=10 - else - uv run --no-sync pytest ${TEST_PATH:?} -vv --tb=short --durations=10 -n "${WORKERS}" - fi - - - name: Upload Lens database coverage - if: steps.changes.outputs.decision != 'skip' && matrix.shard == 'proxy-behavior' && !cancelled() - uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1 - with: - use_oidc: true - version: v11.3.1 - root_dir: ${{ github.workspace }} - files: coverage-lens-postgres.xml - flags: lens-postgres - fail_ci_if_error: true - - - name: Upload ROI database coverage - if: steps.changes.outputs.decision != 'skip' && matrix.shard == 'roi-database' && !cancelled() - uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1 - with: - use_oidc: true - version: v11.3.1 - root_dir: ${{ github.workspace }} - files: coverage-roi-postgres.xml - flags: roi-postgres - fail_ci_if_error: true diff --git a/.github/workflows/test-redis-compat.yml b/.github/workflows/test-redis-compat.yml deleted file mode 100644 index d6cfacccace..00000000000 --- a/.github/workflows/test-redis-compat.yml +++ /dev/null @@ -1,106 +0,0 @@ -name: "Unit Tests: Redis Client Version Compatibility" - -on: - pull_request: - branches: - - main - - "litellm_**" - paths: - - "litellm/_redis.py" - - "litellm/_redis_credential_provider.py" - - "litellm/caching/redis_cache.py" - - "litellm/caching/evicted_client_closer.py" - - "tests/unit/test_redis.py" - - "tests/local_testing/test_caching.py" - - "tests/unit/caching/test_redis_connection_pool.py" - - "tests/unit/caching/test_redis_cluster_cache.py" - - "tests/unit/caching/test_evicted_client_closer.py" - - ".github/workflows/test-redis-compat.yml" - - "pyproject.toml" - - "uv.lock" - -permissions: - contents: read - -concurrency: - group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} - cancel-in-progress: true - -jobs: - redis-compat: - name: "redis-py ${{ matrix.redis-version }}" - runs-on: ubuntu-latest - timeout-minutes: 15 - permissions: - contents: read - id-token: write - - strategy: - fail-fast: false - matrix: - # 5.3.1 is the version pinned in uv.lock (redisvl caps it below 6); the - # newer legs prove the inspect.signature introspection in litellm/_redis.py - # keeps extracting kwargs on the redis-py releases people actually run now. - # Only the exact release 6.0.0 is skipped: rq (pulled by the proxy extra) - # specifies `redis != 6`, which excludes 6.0.0 alone, so 6.4.0 stands in - # for the 6.x line. - redis-version: ["5.3.1", "6.4.0", "7.4.1", "8.0.1"] - - steps: - - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - persist-credentials: false - - - name: Set up Python - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 - with: - python-version: "3.12" - - - name: Set up uv - uses: ./.github/actions/setup-uv-with-retries - with: - version: "0.10.9" - - - name: Install dependencies - run: | - .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra extra_proxy --extra semantic-router - - - name: Pin redis-py to the matrix version - env: - REDIS_VERSION: ${{ matrix.redis-version }} - run: | - uv pip install "redis==${REDIS_VERSION:?}" - uv run --no-sync python -c "import redis; assert redis.__version__ == '${REDIS_VERSION:?}', redis.__version__; print('redis-py', redis.__version__)" - - - name: Build Redis for cluster authentication tests - run: | - curl --fail --location --retry 3 https://download.redis.io/releases/redis-7.2.16.tar.gz -o "$RUNNER_TEMP/redis-7.2.16.tar.gz" - echo "960a8ec15e34ff40e57ff16837b26b33bd81f2da6d24497bb63de532a323a18e $RUNNER_TEMP/redis-7.2.16.tar.gz" | sha256sum --check - tar -xzf "$RUNNER_TEMP/redis-7.2.16.tar.gz" -C "$RUNNER_TEMP" - make -C "$RUNNER_TEMP/redis-7.2.16" -j2 MALLOC=libc OPTIMIZATION=-O1 redis-server - echo "$RUNNER_TEMP/redis-7.2.16/src" >> "$GITHUB_PATH" - - - name: Run redis unit tests - run: | - redis-server --version - uv run --no-sync pytest \ - tests/unit/test_redis.py \ - tests/unit/caching/test_redis_connection_pool.py \ - tests/unit/caching/test_redis_cluster_cache.py \ - tests/unit/caching/test_evicted_client_closer.py \ - tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \ - tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \ - --tb=short -vv \ - --reruns 2 \ - --reruns-delay 1 \ - --durations=20 \ - --cov=./litellm --cov-report=xml:coverage-redis.xml - - - name: Upload Redis coverage - if: matrix.redis-version == '5.3.1' - uses: codecov/codecov-action@0fb7174895f61a3b6b78fc075e0cd60383518dac # v5.5.5 - with: - use_oidc: true - files: coverage-redis.xml - flags: redis-compat - fail_ci_if_error: false diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index da4477b6947..c4397c2067b 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -20,13 +20,6 @@ concurrency: # rather than alphabetical letter ranges. Adding a new test file means adding it # to whichever group it belongs to, not reshuffling slices. # -# `.circleci/tests.yml` runs each group's files on same-repo events under the -# `proxy-db-` Codecov flag; `.circleci/scripts/unit_selection.sh` holds -# the file lists. That pipeline is manual-only while the tests migrate, so -# `unit-flag` makes the shard run that list on every event. `test-path` keeps -# the files that still reach real providers and never left -# tests/proxy_unit_tests. -# # 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 + @@ -78,93 +71,136 @@ jobs: include: # Must run serially โ€” event-loop conflict with the logging worker. - test-group: key-generation - test-path: "" - unit-flag: proxy-db-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: "" - unit-flag: proxy-db-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: "" - unit-flag: proxy-db-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: "" - unit-flag: proxy-db-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" - unit-flag: proxy-db-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: "" - unit-flag: proxy-db-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" - unit-flag: proxy-db-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: "" - unit-flag: proxy-db-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: "" - unit-flag: proxy-db-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: "" - unit-flag: proxy-db-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: "" - unit-flag: proxy-db-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" - unit-flag: proxy-db-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 }} - unit-flag: ${{ matrix.unit-flag }} workers: ${{ matrix.workers }} reruns: 2 timeout-minutes: ${{ matrix.timeout }} diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index b769c8f6a3d..5f9ea430069 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -35,10 +35,6 @@ concurrency: # 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. -# -# `unit-flag` names the `.circleci/tests.yml` job that now runs part of the -# shard under the same Codecov flag. That pipeline is manual-only while the -# tests migrate, so the shard also runs those files on every event. jobs: unit: name: ${{ matrix.shard }} @@ -50,19 +46,11 @@ jobs: fail-fast: false matrix: include: - - shard: mcp-integration - artifact-name: mcp-integration - test-path: "tests/mcp_tests" - unit-flag: mcp-integration - workers: 2 - reruns: 0 - timeout-minutes: 20 - job-timeout-minutes: 60 - - shard: core-utils artifact-name: core-utils - test-path: tests/unit/decisions - unit-flag: core-utils + test-path: >- + tests/unit/decisions + tests/unit/litellm_core_utils workers: 2 reruns: 1 timeout-minutes: 20 @@ -70,8 +58,22 @@ jobs: - shard: enterprise-routing artifact-name: enterprise-routing - test-path: "" - unit-flag: enterprise-routing + test-path: >- + tests/unit/google_genai + tests/unit/router_strategy + tests/unit/router_utils + tests/unit/proxy/common_utils/test_cache_aware_routing.py + tests/unit/enterprise/enterprise_callbacks/send_emails + tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py + tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py + tests/unit/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py + tests/unit/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py + tests/unit/enterprise/proxy/test_batch_update_db_managed_output_file_id.py + tests/unit/enterprise/proxy/test_deleted_file_returns_403_not_404.py + tests/unit/enterprise/proxy/test_enterprise_routes.py + tests/unit/enterprise/proxy/test_file_deletion_blocking.py + tests/unit/enterprise/proxy/test_managed_files_access_check.py + tests/unit/enterprise/proxy/test_managed_files_hook.py workers: 2 reruns: 2 timeout-minutes: 20 @@ -82,7 +84,7 @@ jobs: test-path: >- tests/test_litellm/integrations tests/test_litellm/tracing - unit-flag: integrations + tests/unit/integrations workers: 2 reruns: 3 timeout-minutes: 20 @@ -90,8 +92,8 @@ jobs: - shard: Vertex AI artifact-name: llm-vertex-ai - test-path: "" - unit-flag: llm-vertex-ai + test-path: >- + tests/unit/llms/vertex_ai workers: 1 reruns: 2 timeout-minutes: 20 @@ -99,8 +101,10 @@ jobs: - shard: All Other Providers artifact-name: llm-other-providers - test-path: "" - unit-flag: llm-other-providers + test-path: >- + tests/unit/llms + --ignore=tests/unit/llms/vertex_ai + --ignore=tests/unit/llms/base_llm/batches/base_batches_config_test.py workers: 2 reruns: 2 timeout-minutes: 20 @@ -110,7 +114,27 @@ jobs: artifact-name: misc test-path: >- tests/test_litellm/test_*.py - unit-flag: misc + tests/unit/test_*.py + tests/unit/test_router + tests/unit/a2a_protocol + tests/unit/batches + tests/unit/chat_completions + tests/unit/completion_extras + tests/unit/containers + tests/unit/embeddings + tests/unit/endpoints + tests/unit/files + tests/unit/harness + tests/unit/images + tests/unit/interactions + tests/unit/messages + tests/unit/rag + tests/unit/rerank_api + tests/unit/rust_bridge + tests/unit/secret_managers + tests/unit/vector_stores + tests/unit/videos + --ignore=tests/unit/rust_bridge/native_route_wheel_test.py workers: 2 reruns: 2 timeout-minutes: 20 @@ -218,7 +242,10 @@ jobs: tests/unit/proxy/enterprise_billing tests/unit/proxy/types_utils tests/unit/proxy/logging_endpoints - unit-flag: proxy-infra + tests/unit/gateway + tests/unit/proxy/management + tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py + tests/unit/proxy/roi_calculator workers: 4 reruns: 2 timeout-minutes: 20 @@ -260,8 +287,8 @@ jobs: - shard: caching-local artifact-name: caching-local - test-path: "" - unit-flag: caching-local + test-path: >- + tests/unit/caching workers: 2 reruns: 2 timeout-minutes: 20 @@ -269,8 +296,8 @@ jobs: - shard: proxy-extras artifact-name: proxy-extras - test-path: "" - unit-flag: proxy-extras + test-path: >- + tests/unit/litellm_proxy_extras workers: 2 reruns: 2 timeout-minutes: 20 @@ -278,8 +305,15 @@ jobs: - shard: enterprise-package artifact-name: enterprise-package - test-path: "" - unit-flag: enterprise-package + test-path: >- + tests/unit/enterprise/integrations + tests/unit/enterprise/proxy/auth + tests/unit/enterprise/proxy/guardrails + tests/unit/enterprise/proxy/hooks + tests/unit/enterprise/proxy/management_endpoints + tests/unit/enterprise/proxy/test_audit_logging_endpoints.py + tests/unit/enterprise/proxy/test_liteadmin.py + tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py workers: 4 reruns: 2 timeout-minutes: 20 @@ -287,19 +321,41 @@ jobs: - shard: responses-caching-types artifact-name: responses-caching-types - test-path: "" - unit-flag: responses-caching-types + test-path: >- + tests/unit/responses + tests/unit/types + --ignore=tests/unit/responses/mcp workers: 2 reruns: 2 timeout-minutes: 20 job-timeout-minutes: 60 + + - shard: unit + artifact-name: unit + test-path: >- + tests/unit/anthropic_interface + tests/unit/compression + tests/unit/enterprise/enterprise_callbacks/test_callback_controls.py + tests/unit/enterprise/enterprise_callbacks/test_llm_guard.py + tests/unit/enterprise/enterprise_callbacks/test_secret_detection.py + tests/unit/integration_support + tests/unit/models + tests/unit/ocr + tests/unit/passthrough + tests/unit/realtime_api + tests/unit/repositories + tests/unit/sandbox + tests/unit/skills/test_skills_main.py + tests/unit/tracing + workers: 2 + reruns: 0 + timeout-minutes: 20 + job-timeout-minutes: 60 uses: ./.github/workflows/_test-unit-base.yml with: test-path: ${{ matrix.test-path }} - unit-flag: ${{ matrix.unit-flag || '' }} workers: ${{ matrix.workers }} reruns: ${{ matrix.reruns }} timeout-minutes: ${{ matrix.timeout-minutes }} job-timeout-minutes: ${{ matrix.job-timeout-minutes }} artifact-name: ${{ matrix.artifact-name }} - legacy-mcp-peer: ${{ matrix.shard == 'mcp-integration' }} diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py index a29f0a1b43a..c648f7faea0 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py @@ -46,6 +46,8 @@ from litellm.proxy._types import ( UserAPIKeyAuth, WebhookEvent, ) +from litellm.repositories.table_repositories import InvitationLinkRepository +from litellm.repositories.user_repository import UserRepository from litellm.secret_managers.main import get_secret_bool from litellm.types.integrations.slack_alerting import LITELLM_LOGO_URL @@ -844,7 +846,7 @@ class BaseEmailLogger(CustomLogger): ) return None - user_row = await prisma_client.db.litellm_usertable.find_unique( + user_row = await UserRepository(prisma_client).table.find_unique( where={"user_id": user_id} ) @@ -929,7 +931,7 @@ class BaseEmailLogger(CustomLogger): try: # Try to get existing invitation existing_invitations = ( - await prisma_client.db.litellm_invitationlink.find_many( + await InvitationLinkRepository(prisma_client).table.find_many( where={"user_id": user_id}, order={"created_at": "desc"}, ) 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 35fc3510414..985a9b6de38 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -15,6 +15,10 @@ from litellm.constants import ( MANAGED_OBJECT_STALENESS_CUTOFF_DAYS, MAX_OBJECTS_PER_POLL_CYCLE, ) +from litellm.repositories.table_repositories import ManagedObjectRepository +from litellm.repositories.team_repository import TeamRepository +from litellm.repositories.user_repository import UserRepository +from litellm.repositories.verification_token_repository import VerificationTokenRepository if TYPE_CHECKING: from prisma import models as prisma_models @@ -58,25 +62,19 @@ class _ManagedObjectRow(Protocol): def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]": - table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable - return table + return ManagedObjectRepository(prisma_client).table def _user_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_UserTable]": - table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = prisma_client.db.litellm_usertable - return table + return UserRepository(prisma_client).table def _token_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_VerificationToken]": - table: Final[TableActions[prisma_models.LiteLLM_VerificationToken]] = ( - prisma_client.db.litellm_verificationtoken - ) - return table + return VerificationTokenRepository(prisma_client).table def _team_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_TeamTable]": - table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = prisma_client.db.litellm_teamtable - return table + return TeamRepository(prisma_client).table class CheckBatchCost: 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 cdeea0d3d4b..738b9ebe1e0 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -16,6 +16,7 @@ from litellm.constants import ( MAX_OBJECTS_PER_POLL_CYCLE, STALE_OBJECT_CLEANUP_BATCH_SIZE, ) +from litellm.repositories.table_repositories import ManagedObjectRepository from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN @@ -43,8 +44,7 @@ class _ManagedObjectRow(Protocol): def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]": - table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable - return table + return ManagedObjectRepository(prisma_client).table class CheckResponsesCost: diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index e0a94612646..011bf95defd 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -1647,7 +1647,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): owner_filter: Final = build_owner_filter(user_api_key_dict) if owner_filter is None: - return FileListPage(**build_list_page([])) + return FileListPage.model_validate(build_list_page([])) if after: cursor_row = await _managed_file_table(self.prisma_client).find_first( @@ -1686,7 +1686,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): cursor_id = chunk[-1].unified_file_id chunk_size = max(chunk_size, FILE_LIST_CONTINUATION_CHUNK_SIZE) - return FileListPage(**build_list_page(matches[:page_size], has_more=len(matches) > page_size)) + return FileListPage.model_validate(build_list_page(matches[:page_size], has_more=len(matches) > page_size)) def _is_batch_polling_enabled(self) -> bool: """ diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index f7b667c8ab2..934a1f0bb6b 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -4085,26 +4085,6 @@ dependencies = [ "strum", ] -[[package]] -name = "litellm-migrate" -version = "0.1.0" -dependencies = [ - "litellm-migrate-macros", - "rstest", -] - -[[package]] -name = "litellm-migrate-macros" -version = "0.1.0" -dependencies = [ - "proc-macro2", - "quote", - "rstest", - "syn 2.0.119", - "tempfile", - "thiserror 2.0.19", -] - [[package]] name = "litellm-model-catalog" version = "0.1.0" @@ -4170,6 +4150,7 @@ dependencies = [ "serde_json", "serde_with", "sha2 0.10.9", + "sqlx", "strum", "thiserror 2.0.19", "tokio", @@ -4368,6 +4349,9 @@ dependencies = [ "rstest", "serde", "serde_json", + "serde_with", + "sqlx", + "testcontainers-modules", "thiserror 2.0.19", "tokio", "url", @@ -4498,7 +4482,6 @@ dependencies = [ "hmac 0.12.1", "jsonschema", "litellm-http", - "litellm-migrate", "litellm-storage-clickhouse", "litellm-traces", "litellm-traces-cache", @@ -4509,6 +4492,7 @@ dependencies = [ "serde", "serde_json", "sha2 0.10.9", + "sqlx", "strum", "testcontainers-modules", "thiserror 2.0.19", diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index b0766f11e87..6ef90c59d2f 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -16,8 +16,6 @@ litellm-traces = { path = "crates/traces" } litellm-traces-cache = { path = "crates/traces-cache" } litellm-traces-clickhouse = { path = "crates/traces-clickhouse" } litellm-storage-clickhouse = { path = "crates/storage-clickhouse" } -litellm-migrate = { path = "crates/migrate" } -litellm-migrate-macros = { path = "crates/migrate-macros" } litellm-core = { path = "crates/core" } litellm-gateway-mcp = { path = "crates/gateway-mcp" } litellm-gateway = { path = "crates/gateway" } diff --git a/litellm-rust/crates/migrate-macros/Cargo.toml b/litellm-rust/crates/migrate-macros/Cargo.toml deleted file mode 100644 index 5cd68415ca2..00000000000 --- a/litellm-rust/crates/migrate-macros/Cargo.toml +++ /dev/null @@ -1,19 +0,0 @@ -[package] -name = "litellm-migrate-macros" -version = "0.1.0" -edition.workspace = true -license.workspace = true -repository.workspace = true - -[lib] -proc-macro = true - -[dependencies] -proc-macro2.workspace = true -quote.workspace = true -syn = { workspace = true, features = ["parsing", "printing", "proc-macro"] } -thiserror.workspace = true - -[dev-dependencies] -rstest.workspace = true -tempfile.workspace = true diff --git a/litellm-rust/crates/migrate-macros/src/error.rs b/litellm-rust/crates/migrate-macros/src/error.rs deleted file mode 100644 index 9833009517b..00000000000 --- a/litellm-rust/crates/migrate-macros/src/error.rs +++ /dev/null @@ -1,21 +0,0 @@ -use std::io; - -#[derive(Debug, thiserror::Error)] -pub enum Error { - #[error("could not read migrations directory `{path}`")] - ReadDirectory { - path: String, - #[source] - source: io::Error, - }, - #[error( - "migration name `{name}` must be `_.sql` with a `[a-z0-9_]` description" - )] - InvalidName { name: String }, - #[error("migration version `{version}` is declared more than once")] - DuplicateVersion { version: u64 }, - #[error("migrations directory `{path}` contains no migrations")] - Empty { path: String }, - #[error("migration path `{path}` is not valid UTF-8")] - NonUtf8Path { path: String }, -} diff --git a/litellm-rust/crates/migrate-macros/src/lib.rs b/litellm-rust/crates/migrate-macros/src/lib.rs deleted file mode 100644 index 501f59e6fc2..00000000000 --- a/litellm-rust/crates/migrate-macros/src/lib.rs +++ /dev/null @@ -1,199 +0,0 @@ -mod error; - -use std::path::{Path, PathBuf}; - -use error::Error; -use proc_macro::TokenStream; -use quote::quote; -use syn::LitStr; - -struct Entry { - version: u64, - description: String, - path: PathBuf, -} - -fn resolve(dir: &Path) -> Result, Error> { - let mut entries = Vec::new(); - let files = std::fs::read_dir(dir).map_err(|source| Error::ReadDirectory { - path: dir.display().to_string(), - source, - })?; - for file in files { - let file = file.map_err(|source| Error::ReadDirectory { - path: dir.display().to_string(), - source, - })?; - let path = file.path(); - let name = path - .file_name() - .and_then(|name| name.to_str()) - .ok_or_else(|| Error::NonUtf8Path { - path: path.display().to_string(), - })? - .to_owned(); - let invalid = || Error::InvalidName { name: name.clone() }; - let stem = name - .strip_suffix(".sql") - .filter(|_| file.file_type().is_ok_and(|kind| kind.is_file())) - .and_then(|stem| stem.split_once('_')) - .filter(|(version, description)| { - !version.is_empty() - && version.bytes().all(|b| b.is_ascii_digit()) - && !description.is_empty() - && description - .bytes() - .all(|b| b.is_ascii_lowercase() || b.is_ascii_digit() || b == b'_') - }) - .ok_or_else(invalid)?; - let version = stem.0.parse::().map_err(|_| invalid())?; - entries.push(Entry { - version, - description: stem.1.to_owned(), - path, - }); - } - if entries.is_empty() { - return Err(Error::Empty { - path: dir.display().to_string(), - }); - } - entries.sort_by_key(|entry| entry.version); - for pair in entries.windows(2) { - if pair[0].version == pair[1].version { - return Err(Error::DuplicateVersion { - version: pair[0].version, - }); - } - } - Ok(entries) -} - -fn resolve_input(lit: &LitStr) -> Result, Error> { - let root = std::env::var("CARGO_MANIFEST_DIR") - .map(PathBuf::from) - .unwrap_or_default(); - let dir = root.join(lit.value()); - let dir = dir.canonicalize().map_err(|source| Error::ReadDirectory { - path: dir.display().to_string(), - source, - })?; - if dir.to_str().is_none() { - return Err(Error::NonUtf8Path { - path: dir.display().to_string(), - }); - } - resolve(&dir) -} - -#[proc_macro] -pub fn migrate(input: TokenStream) -> TokenStream { - let lit = syn::parse_macro_input!(input as LitStr); - match resolve_input(&lit) { - Ok(entries) => { - let migrations = entries.iter().map(|entry| { - let version = entry.version; - let description = &entry.description; - let path = entry - .path - .to_str() - .expect("canonical migration path is UTF-8"); - quote! { - ::litellm_migrate::Migration { - version: #version, - description: #description, - sql: ::core::include_str!(#path), - } - } - }); - quote! { &[#(#migrations),*] }.into() - } - Err(err) => syn::Error::new(lit.span(), err).to_compile_error().into(), - } -} - -#[cfg(test)] -mod tests { - use std::fs; - - use rstest::rstest; - use tempfile::TempDir; - - use super::{Error, resolve}; - - fn migrations_dir(files: &[&str]) -> TempDir { - let dir = TempDir::new().expect("tempdir"); - for file in files { - fs::write(dir.path().join(file), "SELECT 1").expect("write fixture"); - } - dir - } - - #[rstest] - fn orders_versions_numerically() { - let dir = migrations_dir(&["10_tenth.sql", "2_second.sql", "1_first.sql"]); - let entries = resolve(dir.path()).expect("resolves"); - let versions: Vec = entries.iter().map(|entry| entry.version).collect(); - let descriptions: Vec<&str> = entries - .iter() - .map(|entry| entry.description.as_str()) - .collect(); - assert_eq!(versions, [1, 2, 10]); - assert_eq!(descriptions, ["first", "second", "tenth"]); - } - - #[rstest] - #[case::dash_in_version(&["0001-dash.sql"])] - #[case::not_sql(&["notes.txt"])] - #[case::empty_description(&["0001_.sql"])] - #[case::non_digit_version(&["x_name.sql"])] - #[case::uppercase_description(&["0001_Upper.sql"])] - #[case::no_underscore(&["0001.sql"])] - #[case::plus_sign_version(&["+10_add.sql"])] - fn rejects_invalid_names(#[case] files: &[&str]) { - let dir = migrations_dir(files); - assert!(matches!( - resolve(dir.path()), - Err(Error::InvalidName { .. }) - )); - } - - #[rstest] - fn rejects_subdirectories() { - let dir = migrations_dir(&["0001_a.sql"]); - fs::create_dir(dir.path().join("0002_b.sql")).expect("subdir"); - assert!(matches!( - resolve(dir.path()), - Err(Error::InvalidName { .. }) - )); - } - - #[cfg(unix)] - #[rstest] - fn rejects_symlinks() { - let dir = migrations_dir(&["0001_a.sql"]); - let target = TempDir::new().expect("tempdir"); - let target_file = target.path().join("real.sql"); - fs::write(&target_file, "SELECT 2").expect("write fixture"); - std::os::unix::fs::symlink(&target_file, dir.path().join("0002_b.sql")).expect("symlink"); - assert!(matches!( - resolve(dir.path()), - Err(Error::InvalidName { .. }) - )); - } - - #[rstest] - fn rejects_duplicate_versions() { - let dir = migrations_dir(&["0001_a.sql", "1_b.sql"]); - assert!(matches!( - resolve(dir.path()), - Err(Error::DuplicateVersion { version: 1 }) - )); - } - - #[rstest] - fn rejects_empty_directory() { - let dir = migrations_dir(&[]); - assert!(matches!(resolve(dir.path()), Err(Error::Empty { .. }))); - } -} diff --git a/litellm-rust/crates/migrate/Cargo.toml b/litellm-rust/crates/migrate/Cargo.toml deleted file mode 100644 index bb1ecaa3128..00000000000 --- a/litellm-rust/crates/migrate/Cargo.toml +++ /dev/null @@ -1,12 +0,0 @@ -[package] -name = "litellm-migrate" -version = "0.1.0" -edition.workspace = true -license.workspace = true -repository.workspace = true - -[dependencies] -litellm-migrate-macros.workspace = true - -[dev-dependencies] -rstest.workspace = true diff --git a/litellm-rust/crates/migrate/README.md b/litellm-rust/crates/migrate/README.md deleted file mode 100644 index 4817029451c..00000000000 --- a/litellm-rust/crates/migrate/README.md +++ /dev/null @@ -1,5 +0,0 @@ -# Migrations - -`litellm-migrate` exports the `Migration` struct and the `migrate!` macro that embeds a directory of `_.sql` files at compile time, sorted by numeric version - -The crate does not apply or track migrations; callers decide how and when the embedded SQL runs diff --git a/litellm-rust/crates/migrate/src/lib.rs b/litellm-rust/crates/migrate/src/lib.rs deleted file mode 100644 index f4e065e1b53..00000000000 --- a/litellm-rust/crates/migrate/src/lib.rs +++ /dev/null @@ -1,8 +0,0 @@ -pub use litellm_migrate_macros::migrate; - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct Migration { - pub version: u64, - pub description: &'static str, - pub sql: &'static str, -} diff --git a/litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql b/litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql deleted file mode 100644 index 31807719e9c..00000000000 --- a/litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql +++ /dev/null @@ -1 +0,0 @@ -SELECT 10; diff --git a/litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql b/litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql deleted file mode 100644 index e0ac49d1ecf..00000000000 --- a/litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql +++ /dev/null @@ -1 +0,0 @@ -SELECT 1; diff --git a/litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql b/litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql deleted file mode 100644 index e7f8100648d..00000000000 --- a/litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql +++ /dev/null @@ -1 +0,0 @@ -SELECT 2; diff --git a/litellm-rust/crates/migrate/tests/migrate.rs b/litellm-rust/crates/migrate/tests/migrate.rs deleted file mode 100644 index 61c80351cf4..00000000000 --- a/litellm-rust/crates/migrate/tests/migrate.rs +++ /dev/null @@ -1,21 +0,0 @@ -use litellm_migrate::Migration; -use rstest::rstest; - -const MIGRATIONS: &[Migration] = litellm_migrate::migrate!("tests/fixtures/migrations"); - -#[rstest] -#[case::first(0, 1, "first", include_str!("fixtures/migrations/1_first.sql"))] -#[case::second(1, 2, "second", include_str!("fixtures/migrations/2_second.sql"))] -#[case::tenth(2, 10, "tenth", include_str!("fixtures/migrations/10_tenth.sql"))] -fn embeds_every_file_sorted_by_numeric_version( - #[case] index: usize, - #[case] version: u64, - #[case] description: &str, - #[case] sql: &str, -) { - assert_eq!(MIGRATIONS.len(), 3); - let migration = &MIGRATIONS[index]; - assert_eq!(migration.version, version); - assert_eq!(migration.description, description); - assert_eq!(migration.sql, sql); -} diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 226cecd5b55..7458e64d374 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -74,6 +74,7 @@ serde_with.workspace = true criterion.workspace = true futures-util.workspace = true rstest.workspace = true +sqlx = { workspace = true, features = ["migrate"] } sha2.workspace = true tokio-tungstenite.workspace = true wiremock.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index f442d6c31de..65db1439d13 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -52,8 +52,7 @@ fn map_error_ref(error: &Error) -> PyErr { | Error::InvalidParameters | Error::InvalidScope => PyValueError::new_err(error.to_string()), Error::Task - | Error::SchemaFailed(_) - | Error::SchemaTransport + | Error::Migration(_) | Error::MissingSecret | Error::Busy | Error::ProvisionFailed(_) @@ -452,7 +451,14 @@ mod tests { )] #[case::insert_budget(Error::InsertTooLarge, "OverflowError")] #[case::scope(Error::InvalidScope, "ValueError")] - #[case::schema(Error::SchemaFailed(503), "RuntimeError")] + #[case::schema( + Error::Storage(litellm_storage_clickhouse::Error::SchemaFailed(503)), + "RuntimeError" + )] + #[case::migration( + Error::Migration(sqlx::migrate::MigrateError::VersionMismatch(1)), + "RuntimeError" + )] #[case::reader(Error::MissingSecret, "RuntimeError")] #[case::storage( Error::Storage(litellm_storage_clickhouse::Error::InvalidUrl), diff --git a/litellm-rust/crates/storage-clickhouse/AGENTS.md b/litellm-rust/crates/storage-clickhouse/AGENTS.md index 959ffdffb88..3f87e5d0759 100644 --- a/litellm-rust/crates/storage-clickhouse/AGENTS.md +++ b/litellm-rust/crates/storage-clickhouse/AGENTS.md @@ -3,3 +3,7 @@ `litellm-storage-clickhouse` exports `Storage`, a writer and bounded reader derived from one ClickHouse URL and database. It also exports bounded HTTP read and insert execution The crate has no trace tables, OTLP types, or named trace queries. `litellm-traces-clickhouse` supplies those rules and uses this storage for both trace rows and spend rows + +It also applies embedded SQLx migrations through the `_sqlx_migrations` ledger + +Migration files are append-only, and changed applied files are rejected by their checksums. Startup migrations must be replay-safe schema changes because the runner records success after execution without dirty states or locks. Backfills belong in coordinated jobs outside proxy startup. The replay policy lives in `ClickHouseMigrate::apply`, `dirty_version`, and `lock`; a Keeper-backed or deploy-time runner changes only those methods diff --git a/litellm-rust/crates/storage-clickhouse/Cargo.toml b/litellm-rust/crates/storage-clickhouse/Cargo.toml index acce941f2c9..4e72e823dc2 100644 --- a/litellm-rust/crates/storage-clickhouse/Cargo.toml +++ b/litellm-rust/crates/storage-clickhouse/Cargo.toml @@ -11,11 +11,14 @@ flate2.workspace = true litellm-http.workspace = true serde.workspace = true serde_json.workspace = true +serde_with.workspace = true +sqlx = { workspace = true, features = ["migrate"] } thiserror.workspace = true url.workspace = true [dev-dependencies] litellm-http = { workspace = true, features = ["test-support"] } rstest.workspace = true +testcontainers-modules = { version = "0.15.0", features = ["clickhouse"] } tokio.workspace = true wiremock.workspace = true diff --git a/litellm-rust/crates/storage-clickhouse/src/error.rs b/litellm-rust/crates/storage-clickhouse/src/error.rs index acec9b91675..fa0c5eabd73 100644 --- a/litellm-rust/crates/storage-clickhouse/src/error.rs +++ b/litellm-rust/crates/storage-clickhouse/src/error.rs @@ -1,4 +1,4 @@ -#[derive(Debug, thiserror::Error)] +#[derive(Clone, Debug, thiserror::Error)] pub enum Error { #[error("invalid ClickHouse insert row")] InvalidRow, diff --git a/litellm-rust/crates/storage-clickhouse/src/lib.rs b/litellm-rust/crates/storage-clickhouse/src/lib.rs index 7ab2aa9bc0a..e2ef62d1bc3 100644 --- a/litellm-rust/crates/storage-clickhouse/src/lib.rs +++ b/litellm-rust/crates/storage-clickhouse/src/lib.rs @@ -1,9 +1,11 @@ mod error; mod insert; +mod migrate; mod read; pub use error::Error; pub use insert::{insert_compressed_rows, insert_encoded_rows}; +pub use migrate::{ClickHouseMigrate, execute_statement, storage_error}; pub use read::{Parameter, Query, READ_LIMITS, ReadLimits, execute_read, fetch, fetch_json}; use url::Url; diff --git a/litellm-rust/crates/storage-clickhouse/src/migrate.rs b/litellm-rust/crates/storage-clickhouse/src/migrate.rs new file mode 100644 index 00000000000..c249020faaf --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/src/migrate.rs @@ -0,0 +1,308 @@ +use std::{ + future::Future, + pin::Pin, + time::{Duration, Instant}, +}; + +use litellm_http::Client; +use serde::Deserialize; +use serde_with::{DisplayFromStr, PickFirst, serde_as}; +use sqlx::{ + Error as SqlxError, + migrate::{AppliedMigration, Migrate, MigrateError, Migration}, +}; + +use crate::{Connection, Error, READ_LIMITS, valid_identifier}; + +pub async fn execute_statement( + client: &Client, + connection: &Connection, + sql: &str, + timeout: Duration, +) -> Result<(), Error> { + let body = execute_sql(client, connection, sql, timeout).await?; + if !body.trim().is_empty() { + return Err(Error::InvalidResponse); + } + Ok(()) +} + +async fn execute_sql( + client: &Client, + connection: &Connection, + sql: &str, + timeout: Duration, +) -> Result { + let mut url = connection.url().clone(); + let pairs: Vec<_> = url + .query_pairs() + .filter(|(key, _)| { + !matches!( + key.as_ref(), + "query" | "wait_end_of_query" | "send_progress_in_http_headers" | "async_insert" + ) + }) + .map(|(key, value)| (key.into_owned(), value.into_owned())) + .collect(); + url.query_pairs_mut() + .clear() + .extend_pairs(pairs) + .append_pair("wait_end_of_query", "1") + .append_pair("send_progress_in_http_headers", "0") + .append_pair("async_insert", "0"); + let mut response = client + .post(url) + .timeout(timeout) + .body(sql.to_owned()) + .send() + .await + .map_err(|_| Error::Transport)?; + if !response.status().is_success() { + return Err(Error::SchemaFailed(response.status().as_u16())); + } + let mut body = Vec::new(); + while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? { + if body.len() + chunk.len() > READ_LIMITS.response_bytes { + return Err(Error::ResponseTooLarge); + } + body.extend_from_slice(&chunk); + } + String::from_utf8(body).map_err(|_| Error::InvalidResponse) +} + +/// The startup runner records success after execution, never marks migrations dirty, and skips locks +/// A Keeper-backed or deploy-time runner changes only `apply`, `dirty_version`, and `lock` +pub struct ClickHouseMigrate<'a, R> { + client: &'a Client, + connection: &'a Connection, + database: &'a str, + render: R, + timeout: Duration, +} + +impl<'a, R> ClickHouseMigrate<'a, R> +where + R: Fn(&str) -> String + Send + Sync, +{ + pub fn new( + client: &'a Client, + connection: &'a Connection, + database: &'a str, + render: R, + timeout: Duration, + ) -> Result { + if !valid_identifier(database) { + return Err(Error::InvalidSchema); + } + Ok(Self { + client, + connection, + database, + render, + timeout, + }) + } +} + +#[serde_as] +#[derive(Deserialize)] +struct Applied { + #[serde_as(as = "PickFirst<(_, DisplayFromStr)>")] + version: i64, + checksum: String, +} + +fn migrate_error(error: Error) -> MigrateError { + MigrateError::Execute(SqlxError::AnyDriverError(Box::new(error))) +} + +fn migrate_execution_error(error: Error, version: i64) -> MigrateError { + MigrateError::ExecuteMigration(SqlxError::AnyDriverError(Box::new(error)), version) +} + +fn encode_hex(bytes: &[u8]) -> String { + bytes + .iter() + .map(|byte| format!("{byte:02x}")) + .collect::>() + .join("") +} + +fn decode_hex(value: &str) -> Result, Error> { + let (pairs, remainder) = value.as_bytes().as_chunks::<2>(); + if !remainder.is_empty() { + return Err(Error::InvalidResponse); + } + pairs + .iter() + .map(|pair| { + let high = decode_hex_digit(pair[0]).ok_or(Error::InvalidResponse)?; + let low = decode_hex_digit(pair[1]).ok_or(Error::InvalidResponse)?; + Ok((high << 4) | low) + }) + .collect() +} + +fn decode_hex_digit(value: u8) -> Option { + match value { + b'0'..=b'9' => Some(value - b'0'), + b'a'..=b'f' => Some(value - b'a' + 10), + b'A'..=b'F' => Some(value - b'A' + 10), + _ => None, + } +} + +fn escape_sql_string(value: &str) -> String { + value.replace('\\', "\\\\").replace('\'', "\\'") +} + +type MigrateFuture<'e, T> = Pin + Send + 'e>>; + +impl Migrate for ClickHouseMigrate<'_, R> +where + R: Fn(&str) -> String + Send + Sync, +{ + fn create_schema_if_not_exists<'e>( + &'e mut self, + schema_name: &'e str, + ) -> MigrateFuture<'e, Result<(), MigrateError>> { + Box::pin(async move { + if !valid_identifier(schema_name) { + return Err(migrate_error(Error::InvalidSchema)); + } + let statement = format!("CREATE DATABASE IF NOT EXISTS `{schema_name}`"); + execute_statement(self.client, self.connection, &statement, self.timeout) + .await + .map_err(migrate_error) + }) + } + + fn ensure_migrations_table<'e>( + &'e mut self, + table_name: &'e str, + ) -> MigrateFuture<'e, Result<(), MigrateError>> { + Box::pin(async move { + let database = format!("`{}`", self.database); + execute_statement( + self.client, + self.connection, + &format!("CREATE DATABASE IF NOT EXISTS {database}"), + self.timeout, + ) + .await + .map_err(migrate_error)?; + execute_statement( + self.client, + self.connection, + &format!( + "CREATE TABLE IF NOT EXISTS {database}.{table_name} \ + (version Int64, description String, installed_on DateTime64(3) DEFAULT now64(3), \ + success Bool, checksum String, execution_time Int64) ENGINE = MergeTree ORDER BY version" + ), + self.timeout, + ) + .await + .map_err(migrate_error) + }) + } + + fn dirty_version<'e>( + &'e mut self, + _table_name: &'e str, + ) -> MigrateFuture<'e, Result, MigrateError>> { + Box::pin(async { Ok(None) }) + } + + fn list_applied_migrations<'e>( + &'e mut self, + table_name: &'e str, + ) -> MigrateFuture<'e, Result, MigrateError>> { + Box::pin(async move { + let database = format!("`{}`", self.database); + let statement = format!( + "SELECT DISTINCT version, checksum FROM {database}.{table_name} \ + WHERE success ORDER BY version FORMAT JSONEachRow" + ); + let response = execute_sql(self.client, self.connection, &statement, self.timeout) + .await + .map_err(migrate_error)?; + response + .lines() + .filter(|line| !line.trim().is_empty()) + .map(|line| { + let row = serde_json::from_str::(line) + .map_err(|_| migrate_error(Error::InvalidResponse))?; + let checksum = decode_hex(&row.checksum).map_err(migrate_error)?; + Ok(AppliedMigration { + version: row.version, + checksum: checksum.into(), + }) + }) + .collect() + }) + } + + fn lock(&mut self) -> MigrateFuture<'_, Result<(), MigrateError>> { + Box::pin(async { Ok(()) }) + } + + fn unlock(&mut self) -> MigrateFuture<'_, Result<(), MigrateError>> { + Box::pin(async { Ok(()) }) + } + + fn apply<'e>( + &'e mut self, + table_name: &'e str, + migration: &'e Migration, + ) -> MigrateFuture<'e, Result> { + Box::pin(async move { + let started_at = Instant::now(); + let statement = (self.render)(migration.sql.as_str()); + execute_statement(self.client, self.connection, &statement, self.timeout) + .await + .map_err(|error| migrate_execution_error(error, migration.version))?; + let elapsed = started_at.elapsed(); + let execution_time = elapsed.as_nanos().min(i64::MAX as u128) as i64; + let description = escape_sql_string(&migration.description); + let checksum = encode_hex(&migration.checksum); + let database = format!("`{}`", self.database); + execute_statement( + self.client, + self.connection, + &format!( + "INSERT INTO {database}.{table_name} \ + (version, description, success, checksum, execution_time) \ + VALUES ({}, '{}', true, '{}', {execution_time})", + migration.version, description, checksum + ), + self.timeout, + ) + .await + .map_err(migrate_error)?; + Ok(elapsed) + }) + } + + fn revert<'e>( + &'e mut self, + _table_name: &'e str, + _migration: &'e Migration, + ) -> MigrateFuture<'e, Result> { + Box::pin(async { + Err(MigrateError::Execute(SqlxError::AnyDriverError(Box::new( + std::io::Error::other("ClickHouse migrations are forward-only"), + )))) + }) + } +} + +pub fn storage_error(error: &MigrateError) -> Option<&Error> { + let error = match error { + MigrateError::Execute(error) | MigrateError::ExecuteMigration(error, _) => error, + _ => return None, + }; + match error { + SqlxError::AnyDriverError(error) => error.downcast_ref(), + _ => None, + } +} diff --git a/litellm-rust/crates/storage-clickhouse/tests/migrations.rs b/litellm-rust/crates/storage-clickhouse/tests/migrations.rs new file mode 100644 index 00000000000..29a85a8c4ec --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/tests/migrations.rs @@ -0,0 +1,476 @@ +use std::time::Duration; + +use litellm_http::Client; +use litellm_storage_clickhouse::{ + ClickHouseMigrate, Connection, Error, READ_LIMITS, execute_statement, storage_error, +}; +use rstest::{fixture, rstest}; +use sqlx::{ + SqlStr, + migrate::{Migrate, MigrateError, Migration, MigrationType, Migrator}, +}; +use testcontainers_modules::{ + clickhouse::ClickHouse, + testcontainers::{ContainerAsync, ImageExt, runners::AsyncRunner}, +}; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_string, method, query_param}, +}; + +const CLICKHOUSE_TAG: &str = + "26.9.6.6@sha256:eb4870e7ca7ed70c259eebfcfbee6cf797017f6b5436c2926bbbfe3d4d28486e"; +const DATABASE: &str = "storage_migrate_test"; +const REQUEST_TIMEOUT: Duration = Duration::from_secs(10); +const SELECT_APPLIED: &str = "SELECT DISTINCT version, checksum FROM `trace_test`._sqlx_migrations \ + WHERE success ORDER BY version FORMAT JSONEachRow"; + +type TestResult = Result>; + +struct ClickHouseDatabase { + _container: ContainerAsync, + url: String, + client: Client, +} + +#[fixture] +async fn database() -> TestResult { + let container = ClickHouse::default() + .with_tag(CLICKHOUSE_TAG) + .with_env_var("CLICKHOUSE_SKIP_USER_SETUP", "1") + .start() + .await?; + let url = format!( + "http://{}:{}", + container.get_host().await?, + container.get_host_port_ipv4(8123).await? + ); + Ok(ClickHouseDatabase { + _container: container, + url, + client: Client::no_redirect_for_test(), + }) +} + +#[fixture] +async fn mock_server() -> MockServer { + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200)) + .with_priority(10) + .mount(&server) + .await; + server +} + +fn migration(version: i64, sql: &'static str) -> Migration { + Migration::new( + version, + format!("migration_{version}").into(), + MigrationType::Simple, + SqlStr::from_static(sql), + false, + ) +} + +fn migrator(migrations: Vec) -> Migrator { + Migrator { + ignore_missing: true, + locking: false, + ..Migrator::with_migrations(migrations) + } +} + +fn render_database(sql: &str) -> String { + sql.replace("{database}", &format!("`{DATABASE}`")) +} + +async fn run_migrations( + database: &ClickHouseDatabase, + migrator: &Migrator, + schema: &str, + render: R, +) -> Result<(), MigrateError> +where + R: Fn(&str) -> String + Send + Sync, +{ + let connection = Connection::writer(&database.url).expect("valid ClickHouse URL"); + let mut adapter = ClickHouseMigrate::new( + &database.client, + &connection, + schema, + render, + REQUEST_TIMEOUT, + ) + .expect("valid schema"); + migrator.run_direct(None, &mut adapter, false).await +} + +async fn execute_write(database: &ClickHouseDatabase, sql: &str) -> TestResult { + database + .client + .post(&database.url) + .body(sql.to_owned()) + .send() + .await? + .error_for_status()?; + Ok(()) +} + +async fn read_json(database: &ClickHouseDatabase, sql: &str) -> TestResult { + let response = database + .client + .post(&database.url) + .body(sql.to_owned()) + .send() + .await? + .error_for_status()?; + Ok(serde_json::from_str(&response.text().await?)?) +} + +async fn ledger_versions(database: &ClickHouseDatabase) -> TestResult> { + let response = read_json( + database, + &format!( + "SELECT version FROM `{DATABASE}`._sqlx_migrations \ + GROUP BY version ORDER BY version FORMAT JSON" + ), + ) + .await?; + Ok(response["data"] + .as_array() + .expect("ClickHouse returns versions") + .iter() + .map(|row| row["version"].as_i64().expect("version is Int64")) + .collect()) +} + +fn encode_hex(bytes: &[u8]) -> String { + bytes + .iter() + .map(|byte| format!("{byte:02x}")) + .collect::>() + .join("") +} + +#[rstest] +#[tokio::test] +async fn only_pending_migrations_execute_on_the_second_run( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let migrator = migrator(vec![ + migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.migration_one (id UInt8) ENGINE = MergeTree ORDER BY id", + ), + migration( + 2, + "CREATE TABLE IF NOT EXISTS {database}.migration_two (id UInt8) ENGINE = MergeTree ORDER BY id", + ), + ]); + run_migrations(&database, &migrator, DATABASE, render_database).await?; + run_migrations(&database, &migrator, DATABASE, render_database).await?; + execute_statement( + &database.client, + &Connection::writer(&database.url)?, + "SYSTEM FLUSH LOGS", + REQUEST_TIMEOUT, + ) + .await?; + + let queries = read_json( + &database, + "SELECT count() AS executions FROM system.query_log \ + WHERE type = 'QueryFinish' AND query LIKE \ + 'CREATE TABLE IF NOT EXISTS `storage_migrate_test`.migration_%' FORMAT JSON", + ) + .await?; + assert_eq!(queries["data"][0]["executions"].as_u64(), Some(2)); + assert_eq!(ledger_versions(&database).await?, vec![1, 2]); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn edited_migration_checksum_returns_version_mismatch( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let original = migrator(vec![migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.original (id UInt8) ENGINE = MergeTree ORDER BY id", + )]); + let changed = migrator(vec![migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.changed (id UInt8) ENGINE = MergeTree ORDER BY id", + )]); + run_migrations(&database, &original, DATABASE, render_database).await?; + + assert!(matches!( + run_migrations(&database, &changed, DATABASE, render_database).await, + Err(MigrateError::VersionMismatch(1)) + )); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn failed_migration_is_not_recorded_and_retains_storage_error( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let migrator = migrator(vec![migration(1, "THIS IS NOT VALID CLICKHOUSE SQL")]); + let error = run_migrations(&database, &migrator, DATABASE, render_database) + .await + .expect_err("invalid SQL must fail"); + assert!(matches!(&error, MigrateError::ExecuteMigration(_, 1))); + assert!(matches!( + storage_error(&error), + Some(Error::SchemaFailed(_)) + )); + + let rows = read_json( + &database, + &format!( + "SELECT count() AS rows FROM `{DATABASE}`._sqlx_migrations \ + WHERE version = 1 FORMAT JSON" + ), + ) + .await?; + assert_eq!(rows["data"][0]["rows"].as_u64(), Some(0)); + Ok(()) +} + +#[rstest] +fn invalid_database_identifier_is_rejected() { + let client = Client::no_redirect_for_test(); + let connection = Connection::writer("http://127.0.0.1:1").expect("valid URL"); + assert!(matches!( + ClickHouseMigrate::new( + &client, + &connection, + "storage_test; DROP DATABASE default", + str::to_owned, + REQUEST_TIMEOUT, + ), + Err(Error::InvalidSchema) + )); +} + +#[rstest] +#[tokio::test] +async fn unknown_source_version_is_tolerated( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + run_migrations(&database, &migrator(vec![]), DATABASE, render_database).await?; + execute_write( + &database, + &format!( + "INSERT INTO `{DATABASE}`._sqlx_migrations \ + (version, description, success, checksum, execution_time) \ + VALUES (99, 'unknown', true, '{}', 0)", + "00".repeat(48) + ), + ) + .await?; + let migrator = migrator(vec![migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.known (id UInt8) ENGINE = MergeTree ORDER BY id", + )]); + + run_migrations(&database, &migrator, DATABASE, render_database).await?; + + assert_eq!(ledger_versions(&database).await?, vec![1, 99]); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn duplicate_ledger_rows_are_tolerated( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let applied = migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.duplicate_test (id UInt8) ENGINE = MergeTree ORDER BY id", + ); + let checksum = encode_hex(&applied.checksum); + let migrator = migrator(vec![applied]); + run_migrations(&database, &migrator, DATABASE, render_database).await?; + execute_write( + &database, + &format!( + "INSERT INTO `{DATABASE}`._sqlx_migrations \ + (version, description, success, checksum, execution_time) \ + VALUES (1, 'migration_1', true, '{checksum}', 0)" + ), + ) + .await?; + + run_migrations(&database, &migrator, DATABASE, render_database).await?; + + let rows = read_json( + &database, + &format!( + "SELECT count() AS rows, uniqExact(version) AS versions \ + FROM `{DATABASE}`._sqlx_migrations FORMAT JSON" + ), + ) + .await?; + assert_eq!(rows["data"][0]["rows"].as_u64(), Some(2)); + assert_eq!(rows["data"][0]["versions"].as_u64(), Some(1)); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn revert_reports_forward_only_error() { + let client = Client::no_redirect_for_test(); + let connection = Connection::writer("http://127.0.0.1:1").expect("valid URL"); + let migration = migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.revert_test (id UInt8) ENGINE = MergeTree ORDER BY id", + ); + let mut adapter = ClickHouseMigrate::new( + &client, + &connection, + DATABASE, + render_database, + REQUEST_TIMEOUT, + ) + .expect("valid schema"); + + let error = adapter + .revert("_sqlx_migrations", &migration) + .await + .expect_err("ClickHouse migrations cannot be reverted"); + assert!( + error + .to_string() + .contains("ClickHouse migrations are forward-only") + ); +} + +#[rstest] +#[case::numeric("1")] +#[case::quoted("\"1\"")] +#[tokio::test] +async fn applied_int64_versions_accept_numeric_and_quoted_json( + #[future(awt)] mock_server: MockServer, + #[case] version: &str, +) { + Mock::given(method("POST")) + .and(body_string(SELECT_APPLIED)) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(format!("{{\"version\":{version},\"checksum\":\"00\"}}\n")), + ) + .mount(&mock_server) + .await; + let client = Client::no_redirect_for_test(); + let connection = Connection::writer(&mock_server.uri()).expect("valid URL"); + let migrator = migrator(vec![]); + let mut adapter = ClickHouseMigrate::new( + &client, + &connection, + "trace_test", + str::to_owned, + REQUEST_TIMEOUT, + ) + .expect("valid schema"); + + migrator + .run_direct(None, &mut adapter, false) + .await + .expect("applied version parses"); +} + +#[rstest] +#[tokio::test] +async fn schema_requests_override_unsafe_connection_settings( + #[future(awt)] mock_server: MockServer, +) { + Mock::given(method("POST")) + .and(query_param("wait_end_of_query", "1")) + .and(query_param("send_progress_in_http_headers", "0")) + .and(query_param("async_insert", "0")) + .and(query_param("custom_setting", "preserved")) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&mock_server) + .await; + let url = format!( + "{}/?wait_end_of_query=0&send_progress_in_http_headers=1&async_insert=1&custom_setting=preserved", + mock_server.uri() + ); + execute_statement( + &Client::no_redirect_for_test(), + &Connection::writer(&url).expect("valid URL"), + "CREATE DATABASE IF NOT EXISTS trace_test", + REQUEST_TIMEOUT, + ) + .await + .expect("schema execution succeeds"); + let requests = mock_server + .received_requests() + .await + .expect("requests recorded"); + for name in [ + "wait_end_of_query", + "send_progress_in_http_headers", + "async_insert", + ] { + assert_eq!( + requests[0] + .url + .query_pairs() + .filter(|(key, _)| key == name) + .count(), + 1 + ); + } +} + +#[rstest] +#[tokio::test] +async fn oversized_ledger_response_is_rejected_before_migrations( + #[future(awt)] mock_server: MockServer, +) { + Mock::given(method("POST")) + .and(body_string(SELECT_APPLIED)) + .respond_with( + ResponseTemplate::new(200).set_body_string(" ".repeat(READ_LIMITS.response_bytes + 1)), + ) + .mount(&mock_server) + .await; + let client = Client::no_redirect_for_test(); + let connection = Connection::writer(&mock_server.uri()).expect("valid URL"); + let migrator = migrator(vec![]); + let mut adapter = ClickHouseMigrate::new( + &client, + &connection, + "trace_test", + str::to_owned, + REQUEST_TIMEOUT, + ) + .expect("valid schema"); + let error = migrator + .run_direct(None, &mut adapter, false) + .await + .expect_err("oversized result is rejected"); + + assert!(matches!( + storage_error(&error), + Some(Error::ResponseTooLarge) + )); + assert_eq!( + mock_server + .received_requests() + .await + .expect("requests recorded") + .len(), + 3 + ); +} diff --git a/litellm-rust/crates/traces-cache/src/reader.rs b/litellm-rust/crates/traces-cache/src/reader.rs index 5c68d6e58bd..b0f83a89118 100644 --- a/litellm-rust/crates/traces-cache/src/reader.rs +++ b/litellm-rust/crates/traces-cache/src/reader.rs @@ -15,6 +15,7 @@ use litellm_traces::{ ListTracesParams, ReadAccessParams, SpanDetailParams, SpanErrorParams, TraceIdentityParams, TraceSpansParams, }, + request::{TRACE_PAGE_SIZE_MAX, TRACE_PAGE_SIZE_MIN}, resolve_trace, to_ui_content, }; @@ -133,7 +134,7 @@ impl TraceReader { cursor: Option<&str>, page_size: u32, ) -> Result, ReadError> { - if !(1..=500).contains(&page_size) { + if !(u32::from(TRACE_PAGE_SIZE_MIN)..=u32::from(TRACE_PAGE_SIZE_MAX)).contains(&page_size) { return Err(ReadError::InvalidParameters); } let Some(trace_ref) = reference(store, access, trace_id, trace_ref).await? else { diff --git a/litellm-rust/crates/traces-clickhouse/AGENTS.md b/litellm-rust/crates/traces-clickhouse/AGENTS.md index ae0c1c8eeb9..462100b10a0 100644 --- a/litellm-rust/crates/traces-clickhouse/AGENTS.md +++ b/litellm-rust/crates/traces-clickhouse/AGENTS.md @@ -1,6 +1,7 @@ - Own trace schema, row encoding, SQL query adapters and reader provisioning; consume domain types from `litellm-traces` - Keep generic ClickHouse connections and HTTP execution in `litellm-storage-clickhouse`; keep PyO3 conversion in `python-bridge` -- Keep schema definitions only in `migrations/NNNN_description.sql`, embedded by `litellm_migrate::migrate!` +- Keep schema definitions only in `migrations/NNNN_description.sql`, embedded by `sqlx::migrate!` +- Treat retention TTLs as current configuration: change them in `RETENTION` in `src/schema.rs`, which every startup reapplies, never in a new migration - Require typed query parameters and SELECT-only readers with server-side limits and tenant isolation - Bound insert time and encoded bytes; preserve shared values and explicit retry deduplication - Test storage behavior through the public API against ClickHouse diff --git a/litellm-rust/crates/traces-clickhouse/Cargo.toml b/litellm-rust/crates/traces-clickhouse/Cargo.toml index b90ad8a7bf8..13aa6d06af6 100644 --- a/litellm-rust/crates/traces-clickhouse/Cargo.toml +++ b/litellm-rust/crates/traces-clickhouse/Cargo.toml @@ -16,7 +16,6 @@ flate2.workspace = true futures-util.workspace = true hmac = "0.12.1" litellm-http.workspace = true -litellm-migrate.workspace = true litellm-storage-clickhouse.workspace = true litellm-traces.workspace = true litellm-traces-cache.workspace = true @@ -24,6 +23,7 @@ moka.workspace = true serde.workspace = true serde_json.workspace = true sha2.workspace = true +sqlx = { workspace = true, features = ["migrate", "macros"] } strum.workspace = true thiserror.workspace = true time = { workspace = true, features = ["formatting"] } diff --git a/litellm-rust/crates/traces-clickhouse/src/error.rs b/litellm-rust/crates/traces-clickhouse/src/error.rs index 1d15c556316..e0ef127a0e9 100644 --- a/litellm-rust/crates/traces-clickhouse/src/error.rs +++ b/litellm-rust/crates/traces-clickhouse/src/error.rs @@ -16,10 +16,6 @@ pub enum Error { InvalidResponse, #[error("ClickHouse insert exceeds the encoded size limit")] InsertTooLarge, - #[error("ClickHouse schema setup failed with HTTP status {0}")] - SchemaFailed(u16), - #[error("ClickHouse schema setup transport failed")] - SchemaTransport, #[error("trace SQL queries require a configured proxy master key")] MissingSecret, #[error("invalid trace query scope")] @@ -39,5 +35,7 @@ pub enum Error { #[error(transparent)] Storage(#[from] litellm_storage_clickhouse::Error), #[error(transparent)] + Migration(#[from] sqlx::migrate::MigrateError), + #[error(transparent)] Cached(#[from] std::sync::Arc), } diff --git a/litellm-rust/crates/traces-clickhouse/src/lib.rs b/litellm-rust/crates/traces-clickhouse/src/lib.rs index 43ff0b8bf33..aed87691958 100644 --- a/litellm-rust/crates/traces-clickhouse/src/lib.rs +++ b/litellm-rust/crates/traces-clickhouse/src/lib.rs @@ -33,7 +33,8 @@ pub use query::{QueryHelp, execute_read, query_help, query_sql}; pub use query_access::QueryReaders; pub use reads::ClickHouseTraces; pub use schema::{ - NORMALIZED_FIELD_DEFINITIONS, NormalizedFieldDefinition, ensure_schema, schema_statements, + NORMALIZED_FIELD_DEFINITIONS, NormalizedFieldDefinition, apply_migrations, ensure_schema, + reconcile_retention, schema_statements, }; pub use span_row::span_rows; pub use sql::execute_named_read; diff --git a/litellm-rust/crates/traces-clickhouse/src/schema.rs b/litellm-rust/crates/traces-clickhouse/src/schema.rs index 6a39bd24041..562dbb976c2 100644 --- a/litellm-rust/crates/traces-clickhouse/src/schema.rs +++ b/litellm-rust/crates/traces-clickhouse/src/schema.rs @@ -1,15 +1,26 @@ use litellm_http::Client; -use litellm_migrate::Migration; +use litellm_storage_clickhouse::{ClickHouseMigrate, execute_statement, storage_error}; use serde::Serialize; +use sqlx::migrate::Migrator; use std::time::Duration; use super::{Connection, Error}; const SCHEMA_REQUEST_TIMEOUT: Duration = Duration::from_secs(30); -const MIGRATIONS: &[Migration] = litellm_migrate::migrate!("migrations"); +static MIGRATOR: Migrator = Migrator { + ignore_missing: true, + locking: false, + ..sqlx::migrate!("./migrations") +}; -pub fn schema_statements(database: &str, retention_days: u32) -> Result, Error> { +const RETENTION: [(&str, &str); 3] = [ + ("otel_traces", "toDateTime(Timestamp)"), + ("agent_traces_by_key", "toDateTime(StartTs)"), + ("spend_logs", "toDateTime(start_time)"), +]; + +fn validate_schema(database: &str, retention_days: u32) -> Result<(), Error> { if database.is_empty() || !database .bytes() @@ -18,19 +29,109 @@ pub fn schema_statements(database: &str, retention_days: u32) -> Result String { + sql.replace("{database}", database) + .replace("{retention_days}", &retention_days.to_string()) +} + +pub fn schema_statements(database: &str, retention_days: u32) -> Result, Error> { + validate_schema(database, retention_days)?; let database = format!("`{database}`"); Ok( std::iter::once(format!("CREATE DATABASE IF NOT EXISTS {database}")) - .chain(MIGRATIONS.iter().map(|migration| { - migration - .sql - .replace("{database}", &database) - .replace("{retention_days}", &retention_days.to_string()) - })) + .chain( + MIGRATOR + .migrations + .iter() + .map(|migration| render(migration.sql.as_str(), &database, retention_days)), + ) .collect(), ) } +pub async fn apply_migrations( + client: &Client, + connection: &Connection, + database: &str, + retention_days: u32, +) -> Result<(), Error> { + apply_migrations_with_timeout( + client, + connection, + database, + retention_days, + SCHEMA_REQUEST_TIMEOUT, + ) + .await +} + +async fn apply_migrations_with_timeout( + client: &Client, + connection: &Connection, + database: &str, + retention_days: u32, + request_timeout: Duration, +) -> Result<(), Error> { + validate_schema(database, retention_days)?; + let quoted_database = format!("`{database}`"); + let mut adapter = ClickHouseMigrate::new( + client, + connection, + database, + |sql| render(sql, "ed_database, retention_days), + request_timeout, + )?; + MIGRATOR + .run_direct(None, &mut adapter, false) + .await + .map_err(|error| match storage_error(&error) { + Some(storage_error) => Error::Storage(storage_error.clone()), + None => Error::Migration(error), + }) +} + +pub async fn reconcile_retention( + client: &Client, + connection: &Connection, + database: &str, + retention_days: u32, +) -> Result<(), Error> { + reconcile_retention_with_timeout( + client, + connection, + database, + retention_days, + SCHEMA_REQUEST_TIMEOUT, + ) + .await +} + +async fn reconcile_retention_with_timeout( + client: &Client, + connection: &Connection, + database: &str, + retention_days: u32, + request_timeout: Duration, +) -> Result<(), Error> { + validate_schema(database, retention_days)?; + let database = format!("`{database}`"); + for (table, expression) in RETENTION { + execute_statement( + client, + connection, + &format!( + "ALTER TABLE {database}.{table} MODIFY TTL {expression} + INTERVAL {retention_days} DAY" + ), + request_timeout, + ) + .await?; + } + Ok(()) +} + pub async fn ensure_schema( client: &Client, connection: &Connection, @@ -54,19 +155,22 @@ async fn ensure_schema_with_timeout( retention_days: u32, request_timeout: Duration, ) -> Result<(), Error> { - for statement in schema_statements(database, retention_days)? { - let response = client - .post(connection.url().clone()) - .timeout(request_timeout) - .body(statement) - .send() - .await - .map_err(|_| Error::SchemaTransport)?; - if !response.status().is_success() { - return Err(Error::SchemaFailed(response.status().as_u16())); - } - } - Ok(()) + apply_migrations_with_timeout( + client, + connection, + database, + retention_days, + request_timeout, + ) + .await?; + reconcile_retention_with_timeout( + client, + connection, + database, + retention_days, + request_timeout, + ) + .await } #[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)] diff --git a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs index 09fd3e78da8..8976ef208d8 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs @@ -1,11 +1,14 @@ use std::{collections::BTreeMap, time::Duration}; use litellm_http::Client; +use litellm_storage_clickhouse::Error as StorageError; use litellm_traces_clickhouse::{ Connection, Error, InsertTable, NORMALIZED_FIELD_DEFINITIONS, Parameter, ReadQuery, - encode_rows, ensure_schema, execute_named_read, execute_read, schema_statements, + apply_migrations, encode_rows, ensure_schema, execute_named_read, execute_read, + reconcile_retention, schema_statements, }; use rstest::rstest; +use sqlx::migrate::MigrateError; mod support; use support::{ClickHouseDatabase, TestResult, database}; @@ -71,6 +74,44 @@ async fn mutation_rows(database: &ClickHouseDatabase) -> TestResult { .expect("ClickHouse returns mutation counts as unsigned integers")) } +fn migration_versions() -> Vec { + let mut versions = std::fs::read_dir(concat!(env!("CARGO_MANIFEST_DIR"), "/migrations")) + .expect("migration directory exists") + .map(|entry| { + entry + .expect("migration directory entry is readable") + .file_name() + .into_string() + .expect("migration file name is UTF-8") + }) + .filter_map(|name| { + name.strip_suffix(".sql") + .and_then(|stem| stem.split('_').next()) + .and_then(|version| version.parse::().ok()) + }) + .collect::>(); + versions.sort_unstable(); + versions +} + +async fn migration_ledger_versions(database: &ClickHouseDatabase) -> TestResult> { + let response = read_json( + database, + "SELECT version FROM trace_test._sqlx_migrations GROUP BY version ORDER BY version", + ) + .await?; + Ok(response["data"] + .as_array() + .expect("ClickHouse returns version rows") + .iter() + .map(|row| { + row["version"] + .as_u64() + .expect("ClickHouse returns versions as unsigned integers") + }) + .collect()) +} + #[rstest] #[tokio::test] async fn schema_supports_span_rollups_and_spend_joins( @@ -80,6 +121,27 @@ async fn schema_supports_span_rollups_and_spend_joins( let writer = Connection::writer(&database.url)?; ensure_schema(&database.client, &writer, "trace_test", 7).await?; ensure_schema(&database.client, &writer, "trace_test", 7).await?; + let expected_versions = migration_versions(); + let ledger = read_json( + &database, + "SELECT count() AS rows, uniqExact(version) AS versions, \ + countIf(NOT match(checksum, '^[0-9a-f]{96}$')) AS invalid_checksums \ + FROM trace_test._sqlx_migrations", + ) + .await?; + assert_eq!( + ledger["data"][0]["rows"].as_u64(), + Some(expected_versions.len() as u64) + ); + assert_eq!( + ledger["data"][0]["versions"].as_u64(), + Some(expected_versions.len() as u64) + ); + assert_eq!(ledger["data"][0]["invalid_checksums"].as_u64(), Some(0)); + assert_eq!( + migration_ledger_versions(&database).await?, + expected_versions + ); let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; let span = serde_json::from_value(serde_json::json!({ "Timestamp": timestamp, "TraceId": "trace-1", "SpanId": "span-1", "ParentSpanId": "", @@ -202,6 +264,112 @@ async fn schema_supports_span_rollups_and_spend_joins( Ok(()) } +#[rstest] +#[case::quoted_versions("output_format_json_quote_64bit_integers=1")] +#[case::asynchronous_inserts( + "async_insert=1&wait_for_async_insert=0&async_insert_busy_timeout_ms=20000" +)] +#[tokio::test] +async fn schema_setup_records_migrations_synchronously_with_configured_settings( + #[future(awt)] database: TestResult, + #[case] settings: &str, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&format!("{}?{settings}", database.url))?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + assert_eq!( + migration_ledger_versions(&database).await?, + migration_versions() + ); + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + assert_eq!( + migration_ledger_versions(&database).await?, + migration_versions() + ); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn changed_migration_is_rejected( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + execute_write( + &database, + "ALTER TABLE trace_test._sqlx_migrations UPDATE checksum = '00' \ + WHERE version = 1 SETTINGS mutations_sync = 1", + ) + .await?; + + assert!(matches!( + ensure_schema(&database.client, &writer, "trace_test", 7).await, + Err(Error::Migration(MigrateError::VersionMismatch(1))) + )); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn concurrent_schema_setup_succeeds( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + let (first, second, third, fourth) = tokio::join!( + ensure_schema(&database.client, &writer, "trace_test", 7), + ensure_schema(&database.client, &writer, "trace_test", 7), + ensure_schema(&database.client, &writer, "trace_test", 7), + ensure_schema(&database.client, &writer, "trace_test", 7), + ); + for result in [first, second, third, fourth] { + result?; + } + let tables = read_json( + &database, + "SELECT count() AS tables FROM system.tables \ + WHERE database = 'trace_test' AND name IN \ + ('otel_traces', 'agent_traces_by_key', 'spend_logs')", + ) + .await?; + assert_eq!(tables["data"][0]["tables"].as_u64(), Some(3)); + assert_eq!( + migration_ledger_versions(&database).await?, + migration_versions() + ); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn existing_schema_without_ledger_is_adopted( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + for statement in schema_statements("trace_test", 7)? { + execute_write(&database, &statement).await?; + } + let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; + let span = serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp, "TraceId": "trace-adopted", "SpanId": "span-adopted", + "ParentSpanId": "", "ServiceName": "proxy", "SpanName": "request", + "Input": "existing row", "ResourceAttributes": {}, "SpanAttributes": {} + }))?; + insert_rows(&database, "otel_traces", vec![span]).await?; + + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + + assert_eq!(table_rows(&database, "otel_traces").await?, 1); + assert_eq!( + migration_ledger_versions(&database).await?, + migration_versions() + ); + Ok(()) +} + #[rstest] #[tokio::test] async fn normalized_fields_match_clickhouse_catalog( @@ -785,6 +953,60 @@ async fn retention_changes_materialize_existing_rows_and_remain_idempotent( Ok(()) } +#[rstest] +#[tokio::test] +async fn retention_reconciliation_updates_each_table_ttl( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + apply_migrations(&database.client, &writer, "trace_test", 7).await?; + reconcile_retention(&database.client, &writer, "trace_test", 7).await?; + let ttl_queries = read_json( + &database, + "SELECT name, create_table_query FROM system.tables \ + WHERE database = 'trace_test' AND name IN \ + ('otel_traces', 'agent_traces_by_key', 'spend_logs') ORDER BY name", + ) + .await?; + let ttl_queries = ttl_queries["data"].as_array().expect("retention tables"); + assert_eq!( + ttl_queries + .iter() + .map(|row| row["name"].as_str().expect("table name")) + .collect::>(), + ["agent_traces_by_key", "otel_traces", "spend_logs"] + ); + for row in ttl_queries { + let query = row["create_table_query"] + .as_str() + .expect("table creation query"); + assert!( + query.contains("toIntervalDay(7)") || query.contains("INTERVAL 7 DAY"), + "{query}" + ); + } + + reconcile_retention(&database.client, &writer, "trace_test", 3).await?; + let ttl_queries = read_json( + &database, + "SELECT name, create_table_query FROM system.tables \ + WHERE database = 'trace_test' AND name IN \ + ('otel_traces', 'agent_traces_by_key', 'spend_logs') ORDER BY name", + ) + .await?; + for row in ttl_queries["data"].as_array().expect("retention tables") { + let query = row["create_table_query"] + .as_str() + .expect("table creation query"); + assert!( + query.contains("toIntervalDay(3)") || query.contains("INTERVAL 3 DAY"), + "{query}" + ); + } + Ok(()) +} + #[rstest] #[tokio::test] async fn schema_statement_timeout_maps_to_transport_error() -> TestResult { @@ -804,7 +1026,7 @@ async fn schema_statement_timeout_maps_to_transport_error() -> TestResult { .await; server.abort(); assert!( - matches!(result, Ok(Err(Error::SchemaTransport))), + matches!(result, Ok(Err(Error::Storage(StorageError::Transport)))), "{result:?}" ); Ok(()) diff --git a/litellm-rust/crates/traces/AGENTS.md b/litellm-rust/crates/traces/AGENTS.md index a9249f31aa8..f31750f8712 100644 --- a/litellm-rust/crates/traces/AGENTS.md +++ b/litellm-rust/crates/traces/AGENTS.md @@ -4,3 +4,8 @@ - Keep ClickHouse schema, row encoding and queries in `litellm-traces-clickhouse`; keep PyO3 conversion in `python-bridge` - Test decoding and normalization through the public API - Expose one top-level `Error` enum in `src/error.rs` for decoding and normalization failures +- Own the tracing HTTP contracts: request types in `src/request.rs` and response views exported from `src/schema.rs` +- Python and the dashboard consume them only through generated code: `uv run scripts/generate_trace_types.py` writes `litellm/rust_bridge/trace/generated/`, and `npm run gen:api` in `ui/litellm-dashboard` regenerates `schema.d.ts` from the proxy's OpenAPI +- Declare each bound once as a constant and read it from both the schema attribute and the runtime check +- GET request types accept unknown fields because the routes ignore unknown query parameters; body request types use `deny_unknown_fields` +- Changing a request or response shape changes the public API; ship it in its own behavior-change PR diff --git a/litellm-rust/crates/traces/src/bin/export_schema.rs b/litellm-rust/crates/traces/src/bin/export_schema.rs index 25d1250ef12..edf89fa514f 100644 --- a/litellm-rust/crates/traces/src/bin/export_schema.rs +++ b/litellm-rust/crates/traces/src/bin/export_schema.rs @@ -1,6 +1,8 @@ fn main() { - println!( - "{}", - serde_json::to_string_pretty(&litellm_traces::schema::schemas()).unwrap() - ); + let schemas = if std::env::args().nth(1).as_deref() == Some("--requests") { + litellm_traces::schema::request_schemas() + } else { + litellm_traces::schema::schemas() + }; + println!("{}", serde_json::to_string_pretty(&schemas).unwrap()); } diff --git a/litellm-rust/crates/traces/src/lib.rs b/litellm-rust/crates/traces/src/lib.rs index a44679a59f3..7a72183de2c 100644 --- a/litellm-rust/crates/traces/src/lib.rs +++ b/litellm-rust/crates/traces/src/lib.rs @@ -15,6 +15,7 @@ mod normalize; mod otlp; pub mod query; mod query_access; +pub mod request; mod resolve; #[cfg(feature = "schema")] pub mod schema; diff --git a/litellm-rust/crates/traces/src/request.rs b/litellm-rust/crates/traces/src/request.rs new file mode 100644 index 00000000000..0ade6112f02 --- /dev/null +++ b/litellm-rust/crates/traces/src/request.rs @@ -0,0 +1,56 @@ +pub const TRACE_PAGE_SIZE_MIN: u16 = 1; +pub const TRACE_PAGE_SIZE_MAX: u16 = 500; + +#[macro_rules_attribute::apply(request_type)] +#[derive(Clone, Debug)] +pub struct TraceListRequest { + /// Window start, unix ms. Default: 24h ago + #[serde(default)] + pub start_ms: Option, + /// Window end, unix ms. Default: now + #[serde(default)] + pub end_ms: Option, + #[serde(default)] + #[cfg_attr(feature = "schema", schemars(length(max = 512)))] + pub cursor: Option, +} + +#[macro_rules_attribute::apply(request_type)] +#[derive(Clone, Debug)] +pub struct TraceDetailRequest { + #[serde(default)] + pub trace_ref: String, + #[serde(default)] + #[cfg_attr(feature = "schema", schemars(length(max = 512)))] + pub cursor: Option, + #[serde(default)] + #[cfg_attr( + feature = "schema", + schemars(range(min = TRACE_PAGE_SIZE_MIN, max = TRACE_PAGE_SIZE_MAX)) + )] + pub page_size: Option, +} + +#[macro_rules_attribute::apply(request_type)] +#[derive(Clone, Debug)] +pub struct TraceSpanRequest { + #[serde(default)] + pub trace_ref: String, +} + +#[macro_rules_attribute::apply(request_type)] +#[derive(Clone, Debug)] +pub struct TraceErrorPageRequest { + #[serde(default)] + pub trace_ref: String, + #[serde(default)] + #[cfg_attr(feature = "schema", schemars(length(max = 512)))] + pub cursor: Option, +} + +#[macro_rules_attribute::apply(request_type)] +#[derive(Clone, Debug)] +#[serde(deny_unknown_fields)] +pub struct TraceQueryRequest { + pub sql: String, +} diff --git a/litellm-rust/crates/traces/src/schema.rs b/litellm-rust/crates/traces/src/schema.rs index cfa1d8e201e..56e166019d7 100644 --- a/litellm-rust/crates/traces/src/schema.rs +++ b/litellm-rust/crates/traces/src/schema.rs @@ -36,6 +36,13 @@ fn received() -> Schema { .into_root_schema_for::() } +fn requested() -> Schema { + SchemaSettings::draft2020_12() + .for_deserialize() + .into_generator() + .into_root_schema_for::() +} + fn emitted() -> Schema { SchemaSettings::draft2020_12() .for_serialize() @@ -58,3 +65,28 @@ pub fn schemas() -> BTreeMap<&'static str, Schema> { ("SpanErrorPage", emitted::()), ]) } + +pub fn request_schemas() -> BTreeMap<&'static str, Schema> { + BTreeMap::from([ + ( + "TraceListRequest", + requested::(), + ), + ( + "TraceDetailRequest", + requested::(), + ), + ( + "TraceSpanRequest", + requested::(), + ), + ( + "TraceErrorPageRequest", + requested::(), + ), + ( + "TraceQueryRequest", + requested::(), + ), + ]) +} diff --git a/litellm-rust/crates/traces/tests/request_schema.rs b/litellm-rust/crates/traces/tests/request_schema.rs new file mode 100644 index 00000000000..598f73021fb --- /dev/null +++ b/litellm-rust/crates/traces/tests/request_schema.rs @@ -0,0 +1,93 @@ +#![cfg(feature = "schema")] + +use litellm_traces::request::{ + TraceDetailRequest, TraceErrorPageRequest, TraceListRequest, TraceQueryRequest, + TraceSpanRequest, +}; +use litellm_traces::schema::request_schemas; +use rstest::rstest; +use serde_json::json; + +#[rstest] +#[case::list("TraceListRequest")] +#[case::detail("TraceDetailRequest")] +#[case::span("TraceSpanRequest")] +#[case::error_page("TraceErrorPageRequest")] +fn get_request_schemas_ignore_unknown_fields(#[case] name: &str) { + let schemas = request_schemas(); + let schema = serde_json::to_value(&schemas[name]).unwrap(); + assert_ne!(schema["additionalProperties"], false); +} + +#[rstest] +fn query_request_schema_rejects_unknown_fields() { + let schemas = request_schemas(); + let schema = serde_json::to_value(&schemas["TraceQueryRequest"]).unwrap(); + assert_eq!(schema["additionalProperties"], false); +} + +#[rstest] +fn request_schemas_preserve_explicit_constraints() { + let schemas = request_schemas(); + let detail = serde_json::to_value(&schemas["TraceDetailRequest"]).unwrap(); + let list = serde_json::to_value(&schemas["TraceListRequest"]).unwrap(); + let span = serde_json::to_value(&schemas["TraceSpanRequest"]).unwrap(); + let error_page = serde_json::to_value(&schemas["TraceErrorPageRequest"]).unwrap(); + let query = serde_json::to_value(&schemas["TraceQueryRequest"]).unwrap(); + + assert_eq!(detail["properties"]["page_size"]["minimum"], 1); + assert_eq!(detail["properties"]["page_size"]["maximum"], 500); + assert_eq!(list["properties"]["cursor"]["maxLength"], 512); + assert_eq!(detail["properties"]["cursor"]["maxLength"], 512); + assert_eq!(error_page["properties"]["cursor"]["maxLength"], 512); + assert!(list["properties"]["start_ms"].get("minimum").is_none()); + assert!(list["properties"]["start_ms"].get("maximum").is_none()); + assert!(list["properties"]["end_ms"].get("minimum").is_none()); + assert!(list["properties"]["end_ms"].get("maximum").is_none()); + assert_eq!(detail["properties"]["trace_ref"]["default"], ""); + assert_eq!(span["properties"]["trace_ref"]["default"], ""); + assert_eq!(error_page["properties"]["trace_ref"]["default"], ""); + assert_eq!(query["required"], json!(["sql"])); + assert_eq!( + list["properties"]["start_ms"]["description"], + "Window start, unix ms. Default: 24h ago" + ); + assert_eq!( + list["properties"]["end_ms"]["description"], + "Window end, unix ms. Default: now" + ); +} + +#[rstest] +fn request_models_deserialize_defaults_and_null_cursors() { + let list: TraceListRequest = serde_json::from_value(json!({})).unwrap(); + let detail: TraceDetailRequest = serde_json::from_value(json!({})).unwrap(); + let span: TraceSpanRequest = serde_json::from_value(json!({})).unwrap(); + let error_page: TraceErrorPageRequest = serde_json::from_value(json!({})).unwrap(); + + assert!(list.start_ms.is_none()); + assert!(list.end_ms.is_none()); + assert!(list.cursor.is_none()); + assert_eq!(detail.trace_ref, ""); + assert!(detail.cursor.is_none()); + assert!(detail.page_size.is_none()); + assert_eq!(span.trace_ref, ""); + assert_eq!(error_page.trace_ref, ""); + assert!(error_page.cursor.is_none()); + + let null_cursor: TraceListRequest = serde_json::from_value(json!({"cursor": null})).unwrap(); + assert!(null_cursor.cursor.is_none()); +} + +#[rstest] +fn get_request_models_ignore_unknown_fields_and_query_model_rejects_them() { + assert!(serde_json::from_value::(json!({"unknown": true})).is_ok()); + assert!(serde_json::from_value::(json!({"unknown": true})).is_ok()); + assert!(serde_json::from_value::(json!({"unknown": true})).is_ok()); + assert!(serde_json::from_value::(json!({"unknown": true})).is_ok()); + assert!( + serde_json::from_value::(json!({"sql": "SELECT 1", "unknown": true})) + .is_err() + ); + assert!(serde_json::from_value::(json!({})).is_err()); +} diff --git a/litellm/a2a_protocol/providers/base.py b/litellm/a2a_protocol/providers/base.py index 5a5eff8cf35..6196cc789fc 100644 --- a/litellm/a2a_protocol/providers/base.py +++ b/litellm/a2a_protocol/providers/base.py @@ -4,7 +4,6 @@ Base configuration for A2A protocol providers. from abc import ABC, abstractmethod from collections.abc import AsyncIterator -from typing import Any class BaseA2AProviderConfig(ABC): @@ -19,10 +18,10 @@ class BaseA2AProviderConfig(ABC): async def handle_non_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> dict[str, Any]: + ) -> dict[str, object]: """ Handle non-streaming A2A request. @@ -40,10 +39,10 @@ class BaseA2AProviderConfig(ABC): async def handle_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> AsyncIterator[dict[str, Any]]: + ) -> AsyncIterator[dict[str, object]]: """ Handle streaming A2A request. diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/config.py b/litellm/a2a_protocol/providers/bedrock_agentcore/config.py index 2b37c0c4906..22dc3aa508a 100644 --- a/litellm/a2a_protocol/providers/bedrock_agentcore/config.py +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/config.py @@ -3,7 +3,7 @@ Bedrock AgentCore A2A provider configuration. """ from collections.abc import AsyncIterator -from typing import Any, Final +from typing import Final from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.bedrock_agentcore.handler import ( @@ -23,10 +23,10 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig): async def handle_non_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> dict[str, Any]: + ) -> dict[str, object]: """Handle non-streaming request to AgentCore A2A agent.""" litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: @@ -43,10 +43,10 @@ class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig): async def handle_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> AsyncIterator[dict[str, Any]]: + ) -> AsyncIterator[dict[str, object]]: """Handle streaming request to AgentCore A2A agent.""" litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py b/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py index da5eb522187..aecc887fbfe 100644 --- a/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py @@ -7,7 +7,7 @@ completion bridge that would otherwise strip the envelope. import json from collections.abc import AsyncIterator, Mapping -from typing import Any, Final +from typing import Final from litellm._logging import verbose_logger from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( @@ -30,9 +30,9 @@ class BedrockAgentCoreA2AHandler: async def handle_non_streaming( request_id: str, params: Mapping[str, object], - litellm_params: dict[str, Any], + litellm_params: dict[str, object], agent_extra_headers: dict[str, str] | None = None, - ) -> dict[str, Any]: + ) -> dict[str, object]: """ Handle non-streaming A2A request to AgentCore. @@ -77,9 +77,9 @@ class BedrockAgentCoreA2AHandler: async def handle_streaming( request_id: str, params: Mapping[str, object], - litellm_params: dict[str, Any], + litellm_params: dict[str, object], agent_extra_headers: dict[str, str] | None = None, - ) -> AsyncIterator[dict[str, Any]]: + ) -> AsyncIterator[dict[str, object]]: """ Handle streaming A2A request to AgentCore. diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py b/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py index c486f1f6d95..5d4e0d5fb39 100644 --- a/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py @@ -219,7 +219,7 @@ class BedrockAgentCoreA2ATransformation: return url, signed_headers, signed_body @staticmethod - async def parse_sse_events(response: _SSELineSource) -> AsyncIterator[dict[str, Any]]: + async def parse_sse_events(response: _SSELineSource) -> AsyncIterator[dict[str, object]]: """ Parse SSE events from an httpx streaming response. diff --git a/litellm/a2a_protocol/providers/langflow/config.py b/litellm/a2a_protocol/providers/langflow/config.py index 54d403f88c0..b8819dbe4bc 100644 --- a/litellm/a2a_protocol/providers/langflow/config.py +++ b/litellm/a2a_protocol/providers/langflow/config.py @@ -1,5 +1,4 @@ from collections.abc import AsyncIterator -from typing import Any from litellm.a2a_protocol.litellm_completion_bridge.handler import ( A2A_USER_API_KEY_HASH_PARAM, @@ -16,10 +15,10 @@ class LangFlowA2AConfig(BaseA2AProviderConfig): async def handle_non_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> dict[str, Any]: + ) -> dict[str, object]: litellm_params = kwargs.get("litellm_params") if not litellm_params: raise ValueError( @@ -39,10 +38,10 @@ class LangFlowA2AConfig(BaseA2AProviderConfig): async def handle_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, - ) -> AsyncIterator[dict[str, Any]]: + ) -> AsyncIterator[dict[str, object]]: litellm_params = kwargs.get("litellm_params") if not litellm_params: raise ValueError( diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py index 20404e3702b..f4c758af2d0 100644 --- a/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py +++ b/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py @@ -2,8 +2,7 @@ Pydantic AI provider configuration. """ -from collections.abc import AsyncIterator -from typing import Any +from collections.abc import AsyncIterator, Mapping from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.pydantic_ai_agents.handler import PydanticAIHandler @@ -20,9 +19,12 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig): async def handle_non_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, - **kwargs: Any, + *, + timeout: float = 60.0, + agent_extra_headers: Mapping[str, str] | None = None, + **kwargs: object, ) -> dict[str, object]: """Handle non-streaming request to Pydantic AI agent.""" if api_base is None: @@ -31,14 +33,14 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig): request_id=request_id, params=params, api_base=api_base, - timeout=kwargs.get("timeout", 60.0), - agent_extra_headers=kwargs.get("agent_extra_headers"), + timeout=timeout, + agent_extra_headers=agent_extra_headers, ) async def handle_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, **kwargs, ) -> AsyncIterator[dict[str, object]]: diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py index c083c0267f7..f3ceb899a92 100644 --- a/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py +++ b/litellm/a2a_protocol/providers/pydantic_ai_agents/handler.py @@ -5,8 +5,8 @@ Pydantic AI agents follow A2A protocol but don't support streaming natively. This handler provides fake streaming by converting non-streaming responses into streaming chunks. """ -from collections.abc import AsyncIterator -from typing import Any, Final +from collections.abc import AsyncIterator, Mapping +from typing import Final from litellm._logging import verbose_logger from litellm.a2a_protocol.providers.pydantic_ai_agents.transformation import ( @@ -26,11 +26,11 @@ class PydanticAIHandler: @staticmethod async def handle_non_streaming( request_id: str, - params: dict[str, Any], + params: Mapping[str, object], api_base: str | None = None, timeout: float = 60.0, - agent_extra_headers: dict[str, str] | None = None, - ) -> dict[str, Any]: + agent_extra_headers: Mapping[str, str] | None = None, + ) -> dict[str, object]: """ Handle non-streaming request to Pydantic AI agent. @@ -63,13 +63,13 @@ class PydanticAIHandler: @staticmethod async def handle_streaming( request_id: str, - params: dict[str, Any], + params: Mapping[str, object], api_base: str | None = None, timeout: float = 60.0, chunk_size: int = 50, delay_ms: int = 10, - agent_extra_headers: dict[str, str] | None = None, - ) -> AsyncIterator[dict[str, Any]]: + agent_extra_headers: Mapping[str, str] | None = None, + ) -> AsyncIterator[dict[str, object]]: """ Handle streaming request to Pydantic AI agent with fake streaming. diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/transformation.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/transformation.py index 024e8c179c2..35fe6e76f4d 100644 --- a/litellm/a2a_protocol/providers/pydantic_ai_agents/transformation.py +++ b/litellm/a2a_protocol/providers/pydantic_ai_agents/transformation.py @@ -7,7 +7,7 @@ This module provides fake streaming by converting non-streaming responses into s import asyncio from collections.abc import AsyncIterator, Mapping, Sequence -from typing import Any, Final, Protocol, cast, runtime_checkable +from typing import Final, Protocol, runtime_checkable from uuid import uuid4 from pydantic import TypeAdapter @@ -100,7 +100,7 @@ class PydanticAITransformation: request_id: str, max_attempts: int = 30, poll_interval: float = 0.5, - agent_extra_headers: dict[str, str] | None = None, + agent_extra_headers: Mapping[str, str] | None = None, ) -> dict[str, object]: """ Poll for task completion using tasks/get method. @@ -156,7 +156,7 @@ class PydanticAITransformation: request_id: str, params: "_SupportsModelDump | _SupportsPydanticDict | Mapping[str, object]", timeout: float = 60.0, - agent_extra_headers: dict[str, str] | None = None, + agent_extra_headers: Mapping[str, str] | None = None, ) -> dict[str, object]: """ Send a request to Pydantic AI agent and return the raw task response. @@ -200,7 +200,7 @@ class PydanticAITransformation: # Send request to Pydantic AI agent using shared async HTTP client client: Final = get_async_httpx_client( - llm_provider=cast(Any, "pydantic_ai_agent"), + llm_provider="pydantic_ai_agent", params={"timeout": timeout}, ) response: Final = await client.post( @@ -242,7 +242,7 @@ class PydanticAITransformation: request_id: str, params: "_SupportsModelDump | _SupportsPydanticDict | Mapping[str, object]", timeout: float = 60.0, - agent_extra_headers: dict[str, str] | None = None, + agent_extra_headers: Mapping[str, str] | None = None, ) -> dict[str, object]: """ Send a non-streaming A2A request to Pydantic AI agent and wait for completion. @@ -278,7 +278,7 @@ class PydanticAITransformation: request_id: str, params: "_SupportsModelDump | _SupportsPydanticDict | Mapping[str, object]", timeout: float = 60.0, - agent_extra_headers: dict[str, str] | None = None, + agent_extra_headers: Mapping[str, str] | None = None, ) -> dict[str, object]: """ Send a request to Pydantic AI agent and return the raw task response. diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py index 44873edf271..a1a6c3a5957 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py @@ -3,11 +3,11 @@ A2A provider configuration for IBM watsonx Orchestrate (WXO). """ from collections.abc import AsyncIterator -from typing import Any, Final from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import ( WatsonxOrchestrateHandler, + WXOLitellmParams, ) @@ -17,12 +17,13 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig): async def handle_non_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, - **kwargs: Any, + *, + litellm_params: WXOLitellmParams | None = None, + **kwargs: object, ) -> dict[str, object]: """Handle a non-streaming A2A request via WXO runs API.""" - litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: raise ValueError( "litellm_params is required for WatsonxOrchestrateA2AConfig " @@ -37,12 +38,13 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig): async def handle_streaming( self, request_id: str, - params: dict[str, Any], + params: dict[str, object], api_base: str | None = None, - **kwargs: Any, + *, + litellm_params: WXOLitellmParams | None = None, + **kwargs: object, ) -> AsyncIterator[dict[str, object]]: """Handle a streaming A2A request via WXO streaming runs API.""" - litellm_params: Final = kwargs.get("litellm_params") if not litellm_params: raise ValueError( "litellm_params is required for WatsonxOrchestrateA2AConfig " diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py index c66b07c321c..0a2871d25cc 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py @@ -7,7 +7,7 @@ import hashlib import json import time from collections.abc import AsyncIterator -from typing import Any, Final, NamedTuple, Protocol +from typing import Final, NamedTuple, Protocol import httpx from typing_extensions import NotRequired, ReadOnly, TypedDict @@ -228,7 +228,7 @@ class WatsonxOrchestrateHandler: return run_data @staticmethod - async def _accumulate_wxo_sse_text(response: Any) -> str: + async def _accumulate_wxo_sse_text(response: _SSELineSource) -> str: source: Final[_WXOView] = {"sse_source": response} accumulated_text = "" async for line in source["sse_source"].aiter_lines(): diff --git a/litellm/caching/_embedding_router.py b/litellm/caching/_embedding_router.py index cec25634bb8..d82e58ecc91 100644 --- a/litellm/caching/_embedding_router.py +++ b/litellm/caching/_embedding_router.py @@ -12,7 +12,7 @@ This module is dependency-injected: callers pass the proxy ``llm_router`` and from __future__ import annotations -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final import litellm @@ -39,10 +39,10 @@ def resolve_embedding_router( def build_router_embedding_metadata( - request_metadata: dict[str, Any] | None, -) -> dict[str, Any]: + request_metadata: Mapping[str, object] | None, +) -> Mapping[str, object]: """Forward the caller's full metadata, flagged as a semantic-cache embedding.""" - metadata: Final[dict[str, Any]] = dict(request_metadata or {}) + metadata: Final = dict(request_metadata or {}) metadata["semantic-cache-embedding"] = True return metadata diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 12aaa051041..4070f41fd32 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -446,11 +446,19 @@ class Cache: 2. Else if a model_group is set, then return the model_group as the model. This is used for all requests sent through the litellm.Router() 3. Else use the `model` passed in kwargs """ - metadata: Final[dict] = kwargs.get("metadata", {}) or {} litellm_params: Final[dict] = kwargs.get("litellm_params", {}) or {} - metadata_in_litellm_params: Final[dict] = litellm_params.get("metadata", {}) or {} - model_group: Final[str | None] = metadata.get("model_group") or metadata_in_litellm_params.get("model_group") - caching_group: Final = self._get_caching_group(metadata, model_group) + metadata_sources: Final[tuple[dict, ...]] = ( + kwargs.get("metadata") or {}, + kwargs.get("litellm_metadata") or {}, + litellm_params.get("metadata") or {}, + litellm_params.get("litellm_metadata") or {}, + ) + model_group: Final[str | None] = next( + (source["model_group"] for source in metadata_sources if source.get("model_group")), None + ) + caching_group: Final = next( + (group for source in metadata_sources if (group := self._get_caching_group(source, model_group))), None + ) return caching_group or model_group or kwargs["model"] def _get_caching_group(self, metadata: dict, model_group: str | None) -> str | None: diff --git a/litellm/files/main.py b/litellm/files/main.py index 723784795b0..9057211256d 100644 --- a/litellm/files/main.py +++ b/litellm/files/main.py @@ -331,7 +331,7 @@ async def afile_retrieve( else: response = init_response - return OpenAIFileObject(**response.model_dump()) + return OpenAIFileObject.model_validate(response.model_dump()) except Exception as e: raise e diff --git a/litellm/files/streaming.py b/litellm/files/streaming.py index 5d23ebf32ae..460050872f4 100644 --- a/litellm/files/streaming.py +++ b/litellm/files/streaming.py @@ -34,7 +34,7 @@ class FileContentStreamingResponse: self.custom_llm_provider = custom_llm_provider self.logging_obj = logging_obj self.standard_logging_object: StandardLoggingPayload | None = None - self._hidden_params: dict[str, Any] = {} + self._hidden_params: dict[str, object] = {} self._logging_completed = False self._close_completed = False self._start_time = ( @@ -121,7 +121,7 @@ class FileContentStreamingResponse: return response def _sync_hidden_params(self) -> None: - litellm_params: dict[str, Any] = {} + litellm_params: dict[str, object] = {} if self.logging_obj is not None: litellm_params = self.logging_obj.model_call_details.get("litellm_params", {}) or {} diff --git a/litellm/harness/context.py b/litellm/harness/context.py index eaae9aafe1f..f75ee7e49f1 100644 --- a/litellm/harness/context.py +++ b/litellm/harness/context.py @@ -4,7 +4,7 @@ from __future__ import annotations from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, TypeAlias +from typing import TYPE_CHECKING, TypeAlias from pydantic import BaseModel @@ -42,7 +42,7 @@ class SessionContext: api_base: str | None = None endpoint: ModelEndpoint | None = None instructions: str | None = None - tools: Sequence[Callable[..., Any]] = () + tools: Sequence[Callable[..., object]] = () skills: Sequence[str] = () disable_tools: Sequence[str] = () permissions: PermissionMode = "full" @@ -50,7 +50,7 @@ class SessionContext: output: type[BaseModel] | None = None max_turns: int | None = None timeout: float | None = None - metadata: Mapping[str, Any] = field(default_factory=dict) + metadata: Mapping[str, object] = field(default_factory=dict) options: HarnessOptions | None = None # Set by the handler after each turn. final_text: str = "" diff --git a/litellm/harness/handlers/cli_handler.py b/litellm/harness/handlers/cli_handler.py index 0fda16c34fe..d217d1e5bbb 100644 --- a/litellm/harness/handlers/cli_handler.py +++ b/litellm/harness/handlers/cli_handler.py @@ -13,7 +13,7 @@ import asyncio import os from collections import deque from collections.abc import AsyncIterator, Sequence -from typing import Any, Final +from typing import Final from litellm._logging import verbose_logger from litellm.constants import HARNESS_STDERR_TAIL_LINES, HARNESS_STREAM_READ_CHUNK_BYTES @@ -125,7 +125,7 @@ class CLIHarnessHandler(BaseHarnessHandler): self._proc = proc tail: Final[deque[str]] = deque(maxlen=HARNESS_STDERR_TAIL_LINES) # mutable-ok: bounded stderr ring buffer stderr_task = asyncio.ensure_future(drain_stderr(proc.stderr, tail)) - state: Any = self.config.create_stream_state() + state: Final[object] = self.config.create_stream_state() exit_code: int | None = None try: await send_stdin(proc, request.stdin) diff --git a/litellm/harness/options.py b/litellm/harness/options.py index 18865359ba8..5d7c8ae0807 100644 --- a/litellm/harness/options.py +++ b/litellm/harness/options.py @@ -9,7 +9,7 @@ from typing import Any, Literal @dataclass(frozen=True) class ClaudeCodeOptions: - config: Mapping[str, Any] = field(default_factory=dict) + config: Mapping[str, object] = field(default_factory=dict) env: Mapping[str, str] = field(default_factory=dict) @@ -17,20 +17,20 @@ class ClaudeCodeOptions: class CodexOptions: reasoning_effort: Literal["low", "medium", "high", "xhigh"] | None = None web_search: bool = False - config: Mapping[str, Any] = field(default_factory=dict) + config: Mapping[str, object] = field(default_factory=dict) env: Mapping[str, str] = field(default_factory=dict) @dataclass(frozen=True) class OpenCodeOptions: agent: str = "build" - config: Mapping[str, Any] = field(default_factory=dict) + config: Mapping[str, object] = field(default_factory=dict) env: Mapping[str, str] = field(default_factory=dict) @dataclass(frozen=True) class DeepAgentsOptions: - subagents: Sequence[Any] = () + subagents: Sequence[object] = () recursion_limit: int | None = None diff --git a/litellm/harness/runtime.py b/litellm/harness/runtime.py index da90b28a751..dc24a1f0ef7 100644 --- a/litellm/harness/runtime.py +++ b/litellm/harness/runtime.py @@ -1005,24 +1005,24 @@ def aagent( `await litellm.aagent(...)` returns a Result. With stream=True it returns an async iterator of events instead: `async for event in litellm.aagent(..., stream=True)`. """ - kwargs: dict[str, Any] = { # mutable-ok: forwarded as **kwargs to arun_agent/astream_agent - "sandbox": sandbox, - "model": model, - "api_key": api_key, - "api_base": api_base, - "instructions": instructions, - "tools": tools, - "skills": skills, - "disable_tools": disable_tools, - "permissions": permissions, - "on_approval": on_approval, - "output": output, - "max_turns": max_turns, - "timeout": timeout, - "metadata": metadata, - "options": options, - "install": install, - } - if stream: - return astream_agent(harness, prompt, **kwargs) - return arun_agent(harness, prompt, **kwargs) + call: Final = astream_agent if stream else arun_agent + return call( + harness, + prompt, + sandbox=sandbox, + model=model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ) diff --git a/litellm/harness/sync.py b/litellm/harness/sync.py index 9787800833b..18336a3d06e 100644 --- a/litellm/harness/sync.py +++ b/litellm/harness/sync.py @@ -8,7 +8,7 @@ import threading from collections.abc import AsyncIterator, Callable, Coroutine, Mapping, Sequence from concurrent.futures import Future from typing import ( - Any, + Final, TypeVar, ) @@ -417,24 +417,24 @@ def agent( Prefix the model with `litellm_proxy/` to route every model call through your LiteLLM AI Gateway. """ - kwargs: dict[str, Any] = { # mutable-ok: forwarded as **kwargs to _run/_stream - "sandbox": sandbox, - "model": model, - "api_key": api_key, - "api_base": api_base, - "instructions": instructions, - "tools": tools, - "skills": skills, - "disable_tools": disable_tools, - "permissions": permissions, - "on_approval": on_approval, - "output": output, - "max_turns": max_turns, - "timeout": timeout, - "metadata": metadata, - "options": options, - "install": install, - } - if stream: - return _stream(harness, prompt, **kwargs) - return _run(harness, prompt, **kwargs) + call: Final = _stream if stream else _run + return call( + harness, + prompt, + sandbox=sandbox, + model=model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ) diff --git a/litellm/harness/types.py b/litellm/harness/types.py index 475568a9979..6b0118f4349 100644 --- a/litellm/harness/types.py +++ b/litellm/harness/types.py @@ -7,7 +7,7 @@ import json from collections.abc import Mapping from dataclasses import dataclass, field from enum import Enum -from typing import Any, Literal +from typing import Literal from pydantic import BaseModel @@ -71,7 +71,7 @@ class ToolCall: id: str name: str native_name: str - input: Mapping[str, Any] + input: Mapping[str, object] builtin: bool = True @@ -100,7 +100,7 @@ class Approval: """A request to run a tool. The turn waits until allow() or deny() is called.""" tool: str - input: Mapping[str, Any] + input: Mapping[str, object] _decision: asyncio.Future[tuple[bool, str]] = field( default_factory=lambda: asyncio.get_event_loop().create_future(), compare=False, diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py index 0271cf1e03c..3f89ac6978f 100644 --- a/litellm/integrations/arize/_utils.py +++ b/litellm/integrations/arize/_utils.py @@ -2,6 +2,7 @@ import json from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final +from pydantic import ConfigDict, TypeAdapter from typing_extensions import ReadOnly, TypedDict, override from litellm._logging import verbose_logger @@ -29,6 +30,8 @@ from litellm.integrations._types.open_inference import ( ToolCallAttributes, ) +_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + class ArizeOTELAttributes(BaseLLMObsOTELAttributes): @staticmethod @@ -609,9 +612,7 @@ def _coerce_response_obj_for_attrs(response_obj): text: Final = getattr(response_obj, "text", None) if isinstance(text, str) and text: try: - parsed: Final = json.loads(text) - if isinstance(parsed, dict): - return parsed + return _JSON_OBJECT.validate_python(json.loads(text)) except Exception: pass return response_obj @@ -1062,9 +1063,7 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs): inner = candidate.get("response") if isinstance(inner, str): try: - parsed = json.loads(inner) - if isinstance(parsed, dict): - return parsed + return _JSON_OBJECT.validate_python(json.loads(inner)) except Exception: continue if isinstance(inner, dict): @@ -1078,9 +1077,7 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs): return original if isinstance(original, str): try: - parsed = json.loads(original) - if isinstance(parsed, dict): - return parsed + return _JSON_OBJECT.validate_python(json.loads(original)) except Exception: return None return None diff --git a/litellm/integrations/datadog/datadog_llm_obs.py b/litellm/integrations/datadog/datadog_llm_obs.py index 22bb50cd739..9565a64e69c 100644 --- a/litellm/integrations/datadog/datadog_llm_obs.py +++ b/litellm/integrations/datadog/datadog_llm_obs.py @@ -15,6 +15,7 @@ from types import MappingProxyType from typing import Any, Final, Literal import httpx +from pydantic import ConfigDict, TypeAdapter, ValidationError import litellm from litellm._logging import verbose_logger @@ -57,6 +58,7 @@ from litellm.types.utils import ( _EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({}) _EMPTY_MESSAGE: Final[Message] = {"role": "", "content": ""} _MAX_PARSED_TOOL_ARGUMENT_CHARS: Final = 256 * 1024 +_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) _SAFE_REDACTED_MESSAGE_ROLES: Final = frozenset( {"agent", "assistant", "developer", "function", "model", "system", "tool", "user"} ) @@ -282,8 +284,10 @@ def _to_dd_arguments(raw_arguments: object) -> dict[str, object] | str: return raw_arguments if isinstance(raw_arguments, dict) else str(raw_arguments) if len(raw_arguments) > _MAX_PARSED_TOOL_ARGUMENT_CHARS: return raw_arguments - parsed: Final = safe_json_loads(raw_arguments) - return parsed if isinstance(parsed, dict) else raw_arguments + try: + return _JSON_OBJECT.validate_python(safe_json_loads(raw_arguments)) + except ValidationError: + return raw_arguments def _to_dd_tool_calls(message: Mapping[str, object]) -> tuple[ToolCall, ...]: diff --git a/litellm/integrations/galileo.py b/litellm/integrations/galileo.py index 21906b0d996..9bc324256fb 100644 --- a/litellm/integrations/galileo.py +++ b/litellm/integrations/galileo.py @@ -4,12 +4,12 @@ import json import os import re import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Mapping, Sequence from datetime import datetime, timezone, tzinfo from typing import Any, Final, Protocol, cast import httpx -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter from typing_extensions import ReadOnly, TypedDict import litellm @@ -35,6 +35,11 @@ GALILEO_CLOUD_API_BASE_URL: Final = "https://api.galileo.ai" # unavailable, invalid credentials) cannot leak memory unboundedly. GALILEO_MAX_IN_MEMORY_RECORDS: Final = 1000 +_UNTYPED_VALUE: Final = TypeAdapter(object) +_ZERO_ARGUMENT_CALLABLE: Final[TypeAdapter[Callable[[], object]]] = TypeAdapter( + Callable[[], object], config=ConfigDict(hide_input_in_errors=True) +) + class _GalileoLoginBody(TypedDict): """Decoded body of the Galileo login response.""" @@ -473,9 +478,9 @@ class GalileoObserve(CustomLogger): if isinstance(value, str): return value - def _json_default(obj: Any) -> object: + def _json_default(obj: object) -> object: if hasattr(obj, "model_dump"): - return obj.model_dump() + return _ZERO_ARGUMENT_CALLABLE.validate_python(getattr(obj, "model_dump", None))() return str(obj) return json.dumps(value, default=_json_default) @@ -496,7 +501,7 @@ class GalileoObserve(CustomLogger): if hasattr(message, "json"): message_json: Final[object] = message.json() if isinstance(message_json, str): - return json.loads(message_json) + return _UNTYPED_VALUE.validate_python(json.loads(message_json)) return message_json return message return None diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 8d588896b2f..956e754944a 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -8,6 +8,8 @@ from datetime import datetime from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, TypedDict, cast +from pydantic import ConfigDict, TypeAdapter + import litellm from litellm._logging import verbose_logger from litellm.integrations._types.open_inference import ( @@ -95,6 +97,8 @@ class _ResponseWithUsageView(TypedDict, total=False): usage: "_UsageCompletionTokensView | None" +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + # Cap on credential-scoped providers held at once; each one owns an exporter thread. _MAX_DYNAMIC_TRACER_PROVIDERS: Final = 256 @@ -1633,7 +1637,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): attributes = self.config.attributes if attributes is None and self.callback_name in (None, "otel"): otel_settings: Final = (litellm.callback_settings or {}).get("otel") or {} - raw: Final = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None + raw: Final[object] = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None if raw is not None: attributes = _build_metric_attribute_filter(raw) ( @@ -2919,7 +2923,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): import json try: - _parsed: Final[Mapping[str, object]] = json.loads(_raw_response) + _parsed: Final = _JSON_OBJECT.validate_python(json.loads(_raw_response)) for param, val in _parsed.items(): self.safe_set_attribute( span=span, diff --git a/litellm/integrations/opik/opik_payload_builder/api.py b/litellm/integrations/opik/opik_payload_builder/api.py index 4334c81eb6e..d16bc71c416 100644 --- a/litellm/integrations/opik/opik_payload_builder/api.py +++ b/litellm/integrations/opik/opik_payload_builder/api.py @@ -1,12 +1,17 @@ """Public API for Opik payload building.""" +from collections.abc import Mapping from datetime import datetime from typing import Any, Final +from pydantic import ConfigDict, TypeAdapter + from litellm.integrations.opik import utils from . import extractors, payload_builders, types +_STANDARD_LOGGING_FIELDS: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + def build_opik_payload( kwargs: dict[str, Any], @@ -36,12 +41,14 @@ def build_opik_payload( - First element is TracePayload if creating a new trace, None if attaching to existing - Second element is always SpanPayload """ - standard_logging_object: Final = kwargs["standard_logging_object"] + standard_logging_object: Final = _STANDARD_LOGGING_FIELDS.validate_python(kwargs["standard_logging_object"]) # Extract litellm params and metadata litellm_params: Final = kwargs.get("litellm_params", {}) or {} litellm_metadata: Final = litellm_params.get("metadata", {}) or {} - standard_logging_metadata: Final = standard_logging_object.get("metadata", {}) or {} + standard_logging_metadata: Final = _STANDARD_LOGGING_FIELDS.validate_python( + standard_logging_object.get("metadata", {}) or {} + ) # Extract and merge Opik metadata opik_metadata: Final = extractors.extract_opik_metadata(litellm_metadata, standard_logging_metadata) diff --git a/litellm/integrations/opik/opik_payload_builder/extractors.py b/litellm/integrations/opik/opik_payload_builder/extractors.py index 4dd3d40fae3..1ceb9aa763f 100644 --- a/litellm/integrations/opik/opik_payload_builder/extractors.py +++ b/litellm/integrations/opik/opik_payload_builder/extractors.py @@ -4,8 +4,12 @@ import json from collections.abc import Mapping from typing import Any, Final +from pydantic import TypeAdapter + from litellm import _logging +_DECODED_JSON: Final = TypeAdapter(object) + def normalize_provider_name(provider: str | None) -> str | None: """ @@ -149,7 +153,7 @@ def apply_proxy_header_overrides( thread_id = value elif param_key == "tags": try: - parsed_tags: object = json.loads(value) + parsed_tags = _DECODED_JSON.validate_python(json.loads(value)) if isinstance(parsed_tags, list): tags.extend(parsed_tags) except (json.JSONDecodeError, TypeError): diff --git a/litellm/integrations/opik/opik_payload_builder/types.py b/litellm/integrations/opik/opik_payload_builder/types.py index 546ce55f840..2f418b2ef9e 100644 --- a/litellm/integrations/opik/opik_payload_builder/types.py +++ b/litellm/integrations/opik/opik_payload_builder/types.py @@ -1,7 +1,8 @@ """Type definitions for Opik payload building.""" +from collections.abc import Mapping from dataclasses import dataclass -from typing import Any, Final, Literal +from typing import Final, Literal @dataclass @@ -13,9 +14,9 @@ class TracePayload: name: str start_time: str end_time: str - input: Any - output: Any - metadata: dict[str, Any] + input: object + output: object + metadata: Mapping[str, object] tags: list[str] thread_id: str | None = None @@ -32,9 +33,9 @@ class SpanPayload: model: str start_time: str end_time: str - input: Any - output: Any - metadata: dict[str, Any] + input: object + output: object + metadata: Mapping[str, object] tags: list[str] usage: dict[str, int] parent_span_id: str | None = None diff --git a/litellm/integrations/opik/utils.py b/litellm/integrations/opik/utils.py index 2caacd1f871..79b6199423e 100644 --- a/litellm/integrations/opik/utils.py +++ b/litellm/integrations/opik/utils.py @@ -2,7 +2,8 @@ import configparser import os import time import uuid -from typing import Any, Final +from collections.abc import Mapping, Sequence +from typing import Final CONFIG_FILE_PATH_DEFAULT: Final[str] = "~/.opik.config" @@ -93,14 +94,14 @@ def create_usage_object(usage): return usage_dict -def _remove_nulls(x: dict[str, Any]) -> dict[str, Any]: +def _remove_nulls(x: Mapping[str, object]) -> dict[str, object]: """Remove None values from dict.""" return {k: v for k, v in x.items() if v is not None} def get_traces_and_spans_from_payload( - payload: list[dict[str, Any]], -) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: + payload: Sequence[Mapping[str, object]], +) -> tuple[list[dict[str, object]], list[dict[str, object]]]: """ Separate traces and spans from payload. diff --git a/litellm/integrations/otel/model/config.py b/litellm/integrations/otel/model/config.py index ddb8e127408..6e9038c0cad 100644 --- a/litellm/integrations/otel/model/config.py +++ b/litellm/integrations/otel/model/config.py @@ -2,7 +2,7 @@ from enum import Enum from functools import lru_cache -from typing import Annotated, Any, Final +from typing import Annotated, Final from pydantic import AliasChoices, BaseModel, Field, TypeAdapter, ValidationError, field_validator, model_validator from pydantic.fields import FieldInfo @@ -315,7 +315,7 @@ class OpenTelemetryV2Config(BaseSettings): mode="before", ) @classmethod - def _split_csv(cls, value: Any) -> Any: + def _split_csv(cls, value: object) -> object: """Accept a comma-separated string for list fields. Env vars are strings, but these fields are lists. Pydantic-settings would diff --git a/litellm/integrations/otel/plumbing/metrics.py b/litellm/integrations/otel/plumbing/metrics.py index e1623f4697f..76b23467679 100644 --- a/litellm/integrations/otel/plumbing/metrics.py +++ b/litellm/integrations/otel/plumbing/metrics.py @@ -307,7 +307,7 @@ class GenAIMetricRecorder: return common_attrs - def _bounded_attributes(self, kwargs: Mapping[str, Any]) -> MetricAttributes: + def _bounded_attributes(self, kwargs: Mapping[str, object]) -> MetricAttributes: """The datapoint attributes, capped at :data:`METRIC_ATTRIBUTE_CEILING`. The cap runs BEFORE the operator's include/exclude filter so the filter can @@ -322,7 +322,7 @@ class GenAIMetricRecorder: attributes = None if self._callback_name in (None, "otel"): otel_settings: Final = (litellm.callback_settings or {}).get("otel") or {} - raw: Final = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None + raw: Final[object] = otel_settings.get("attributes") if isinstance(otel_settings, dict) else None if raw is not None: attributes = _build_metric_attribute_filter(raw) # A bad filter (include_list + exclude_list both set, an unfilterable name) diff --git a/litellm/integrations/s3_v2.py b/litellm/integrations/s3_v2.py index f504292cb64..08249d33561 100644 --- a/litellm/integrations/s3_v2.py +++ b/litellm/integrations/s3_v2.py @@ -595,7 +595,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): and not (self.s3_drop_on_terminal_error and _is_terminal(response)) and attempt < max_retries - 1 ): - wait_time = 2**attempt # 1s, 2s + wait_time = 1 << attempt # 1s, 2s verbose_logger.log( logging.DEBUG if _in_flush.get() else logging.WARNING, "S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s", @@ -897,7 +897,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): and not (self.s3_drop_on_terminal_error and _is_terminal(response)) and attempt < max_retries - 1 ): - wait_time = 2**attempt # 1s, 2s + wait_time = 1 << attempt # 1s, 2s verbose_logger.warning( "S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s", response.status_code, diff --git a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py index 80530622e36..f40a166e67d 100644 --- a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py +++ b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py @@ -6,12 +6,12 @@ It searches the vector store for relevant context, runs the request's pre-call g over that context, and appends it to the messages. """ -from collections.abc import Awaitable, Callable, Mapping, Sequence +from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence from dataclasses import dataclass from itertools import chain from typing import TYPE_CHECKING, Any, Final, Protocol, cast, get_args -from pydantic import TypeAdapter, ValidationError +from pydantic import ConfigDict, TypeAdapter, ValidationError from typing_extensions import assert_never import litellm @@ -44,10 +44,12 @@ else: LiteLLMLoggingObj = Any SEARCH_FAILURES_FIELD: Final = "vector_store_search_failures" +_PROVIDER_FIELDS_ATTRIBUTE: Final = "provider_specific_fields" _DEFAULT_FAILURE_MODE: Final[VectorStoreSearchFailureMode] = "annotate" _FAILURE_MODE_ADAPTER: Final = TypeAdapter(VectorStoreSearchFailureMode) _OBJECT_ADAPTER: Final = TypeAdapter(object) _STR_KEYED_ADAPTER: Final = TypeAdapter(dict[str, object]) +_ITERABLE_ADAPTER: Final = TypeAdapter(Iterable[object], config=ConfigDict(hide_input_in_errors=True)) _GUARDRAIL_KEYS_THE_PROXY_MERGES_INTO_METADATA: Final = frozenset( {"guardrails", "guardrail_config", "policies", "include_guardrail_response"} ) @@ -475,7 +477,7 @@ class VectorStorePreCallHook(CustomLogger): async def async_post_call_streaming_deployment_hook( self, request_data: dict, - response_chunk: Any, + response_chunk: object, call_type: CallTypes | None, ) -> object | None: """ @@ -496,15 +498,17 @@ class VectorStorePreCallHook(CustomLogger): return response_chunk # Add search results to streaming chunk - if hasattr(response_chunk, "choices") and response_chunk.choices: - for choice in response_chunk.choices: - if hasattr(choice, "delta") and choice.delta: - provider_fields = getattr(choice.delta, "provider_specific_fields", None) or {} + choices: Final[object] = getattr(response_chunk, "choices", None) + if choices: + for choice in _ITERABLE_ADAPTER.validate_python(choices): + delta: object = getattr(choice, "delta", None) + if delta: + provider_fields = getattr(delta, _PROVIDER_FIELDS_ATTRIBUTE, None) or {} if search_results: provider_fields["search_results"] = search_results if search_failures: provider_fields[SEARCH_FAILURES_FIELD] = search_failures - choice.delta.provider_specific_fields = provider_fields + setattr(delta, _PROVIDER_FIELDS_ATTRIBUTE, provider_fields) # Return modified chunk return response_chunk diff --git a/litellm/llms/a2a/chat/transformation.py b/litellm/llms/a2a/chat/transformation.py index af4c6f69944..6171b341ca4 100644 --- a/litellm/llms/a2a/chat/transformation.py +++ b/litellm/llms/a2a/chat/transformation.py @@ -3,8 +3,8 @@ A2A Protocol Transformation for LiteLLM """ import uuid -from collections.abc import Iterator, Mapping -from typing import TYPE_CHECKING, Any, Final +from collections.abc import AsyncIterator, Iterator, Mapping +from typing import TYPE_CHECKING, Final import httpx @@ -72,9 +72,9 @@ class A2AConfig(BaseConfig): agent_name: str, api_base: str | None, api_key: str | None, - headers: dict[str, Any] | None, - optional_params: dict[str, Any], - ) -> tuple[str | None, str | None, dict[str, Any] | None]: + headers: dict[str, object] | None, + optional_params: dict[str, object], + ) -> tuple[str | None, str | None, dict[str, object] | None]: """ Resolve agent configuration from the registry for a registered agent. @@ -376,7 +376,7 @@ class A2AConfig(BaseConfig): def get_model_response_iterator( self, - streaming_response: Iterator | Any, + streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, json_mode: bool | None = False, ) -> BaseModelResponseIterator: @@ -397,7 +397,7 @@ class A2AConfig(BaseConfig): json_mode=json_mode, ) - def _openai_message_to_a2a_message(self, message: dict[str, Any]) -> dict[str, Any]: + def _openai_message_to_a2a_message(self, message: Mapping[str, object]) -> dict[str, object]: """ Convert OpenAI message to A2A message format. diff --git a/litellm/llms/aiohttp_openai/chat/transformation.py b/litellm/llms/aiohttp_openai/chat/transformation.py index a06c670e3f1..dba7a7405ab 100644 --- a/litellm/llms/aiohttp_openai/chat/transformation.py +++ b/litellm/llms/aiohttp_openai/chat/transformation.py @@ -7,9 +7,11 @@ https://github.com/BerriAI/litellm/issues/6592 New config to ensure we introduce this without causing breaking changes for users """ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final from aiohttp import ClientResponse +from pydantic import ConfigDict, TypeAdapter from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig from litellm.types.llms.openai import AllMessageValues @@ -23,6 +25,10 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECTS: Final = TypeAdapter( + Iterable[Mapping[str, object]], config=ConfigDict(strict=True, hide_input_in_errors=True) +) + class AiohttpOpenAIChatConfig(OpenAILikeChatConfig): def get_complete_url( @@ -73,7 +79,9 @@ class AiohttpOpenAIChatConfig(OpenAILikeChatConfig): ) -> ModelResponse: _json_response: Final = await raw_response.json() model_response.id = _json_response.get("id") - model_response.choices = [Choices(**choice) for choice in _json_response.get("choices")] + model_response.choices = [ + Choices.model_validate(choice) for choice in _JSON_OBJECTS.validate_python(_json_response.get("choices")) + ] model_response.created = _json_response.get("created") model_response.model = _json_response.get("model") model_response.object = _json_response.get("object") diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 806240c9749..578a4e47dfd 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -224,6 +224,11 @@ def _write_back_message_text(message: _WritableMessage, target: MessageTextTarge _TOOL_USE_INPUT_ADAPTER: Final = TypeAdapter(dict[str, object]) +_RELEASED_TOOL_USE_STOP: Final = ( + b"event: message_delta\n" + b'data: {"type": "message_delta", "delta": {"stop_reason": "tool_use", "stop_sequence": null}, ' + b'"usage": {"output_tokens": 0}}\n\n' +) def _rewritten_tool_use_input(arguments: str) -> Mapping[str, object] | None: @@ -1573,6 +1578,12 @@ class AnthropicMessagesHandler(BaseTranslation): tool_calls_in_flight=bool(tool_use_fingerprints) and not stream_ended, ) + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + released_key: Final = self.get_streaming_scan_key(responses_so_far) + if released_key is None or not released_key.tool_calls_in_flight: + return tuple(responses_so_far) + return (*responses_so_far, _RELEASED_TOOL_USE_STOP) + @classmethod def _streamed_tool_use_fingerprints(cls, responses_so_far: Sequence[object]) -> tuple[str, ...]: return tuple( diff --git a/litellm/llms/anthropic/count_tokens/token_counter.py b/litellm/llms/anthropic/count_tokens/token_counter.py index 920916726f2..401c09b580b 100644 --- a/litellm/llms/anthropic/count_tokens/token_counter.py +++ b/litellm/llms/anthropic/count_tokens/token_counter.py @@ -27,7 +27,7 @@ class AnthropicTokenCounter(BaseTokenCounter): self, model_to_use: str, messages: list[dict[str, Any]] | None, - contents: list[dict[str, Any]] | None, + contents: list[dict[str, object]] | None, deployment: dict[str, Any] | None = None, request_model: str = "", tools: list[dict[str, Any]] | None = None, diff --git a/litellm/llms/anthropic/files/transformation.py b/litellm/llms/anthropic/files/transformation.py index da21f130f1a..082902e8897 100644 --- a/litellm/llms/anthropic/files/transformation.py +++ b/litellm/llms/anthropic/files/transformation.py @@ -19,6 +19,7 @@ from typing import Final, cast import httpx from openai.types.file_deleted import FileDeleted +from pydantic import ConfigDict, TypeAdapter from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data from litellm.litellm_core_utils.url_utils import encode_url_path_segment @@ -47,6 +48,8 @@ ANTHROPIC_FILES_API_BASE: Final = "https://api.anthropic.com" ANTHROPIC_FILES_BETA_HEADER: Final = "files-api-2025-04-14" ANTHROPIC_MESSAGE_BATCH_ID_PREFIX: Final = "msgbatch_" +_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + class AnthropicFilesConfig(BaseFilesConfig): """ @@ -211,7 +214,7 @@ class AnthropicFilesConfig(BaseFilesConfig): "created_at": "2025-01-01T00:00:00Z" } """ - response_json: Final = raw_response.json() + response_json: Final = _JSON_OBJECT.validate_python(raw_response.json()) return self._parse_anthropic_file(response_json) def transform_retrieve_file_request( @@ -230,7 +233,7 @@ class AnthropicFilesConfig(BaseFilesConfig): logging_obj: LiteLLMLoggingObj, litellm_params: dict, ) -> OpenAIFileObject: - response_json: Final = raw_response.json() + response_json: Final = _JSON_OBJECT.validate_python(raw_response.json()) return self._parse_anthropic_file(response_json) def transform_delete_file_request( @@ -249,13 +252,9 @@ class AnthropicFilesConfig(BaseFilesConfig): logging_obj: LiteLLMLoggingObj, litellm_params: dict, ) -> FileDeleted: - response_json: Final = raw_response.json() + response_json: Final = _JSON_OBJECT.validate_python(raw_response.json()) file_id: Final = response_json.get("id", "") - return FileDeleted( - id=file_id, - deleted=True, - object="file", - ) + return FileDeleted.model_validate({"id": file_id, "deleted": True, "object": "file"}) def transform_list_files_request( self, diff --git a/litellm/llms/azure/audio_transcription/transformation.py b/litellm/llms/azure/audio_transcription/transformation.py index 44623ec778a..17597da3895 100644 --- a/litellm/llms/azure/audio_transcription/transformation.py +++ b/litellm/llms/azure/audio_transcription/transformation.py @@ -5,10 +5,12 @@ Maps OpenAI-compatible audio transcription calls to Azure Speech REST recognition for short audio. """ -from typing import Any, Final +from collections.abc import Mapping +from typing import Final from urllib.parse import urlencode, urlparse import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.llms.base_llm.audio_transcription.transformation import ( @@ -23,6 +25,8 @@ from litellm.types.llms.openai import ( ) from litellm.types.utils import FileTypes, TranscriptionResponse +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class AzureSpeechAudioTranscriptionException(BaseLLMException): pass @@ -127,7 +131,8 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig): raw_response: httpx.Response, ) -> TranscriptionResponse: response_json: Final = raw_response.json() - recognition_status: Final = response_json.get("RecognitionStatus") + payload: Final = _JSON_OBJECT.validate_python(response_json) + recognition_status: Final = payload.get("RecognitionStatus") if recognition_status is not None and recognition_status != "Success": raise AzureSpeechAudioTranscriptionException( message=(f"Azure AI Speech transcription failed with RecognitionStatus={recognition_status}."), @@ -135,7 +140,7 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig): headers=raw_response.headers, ) - text: Final = self._extract_text(response_json) + text: Final = self._extract_text(payload) response: Final = TranscriptionResponse(text=text) response._hidden_params = response_json return response @@ -194,9 +199,10 @@ class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig): return "detailed" return "simple" - def _extract_text(self, response_json: dict[str, Any]) -> str: - if isinstance(response_json.get("DisplayText"), str): - return response_json["DisplayText"] + def _extract_text(self, response_json: Mapping[str, object]) -> str: + display_text: Final = response_json.get("DisplayText") + if isinstance(display_text, str): + return display_text nbest: Final = response_json.get("NBest") if isinstance(nbest, list) and nbest: diff --git a/litellm/llms/azure_ai/anthropic/count_tokens/handler.py b/litellm/llms/azure_ai/anthropic/count_tokens/handler.py index 36d5a56db0d..6d6e10ce1dc 100644 --- a/litellm/llms/azure_ai/anthropic/count_tokens/handler.py +++ b/litellm/llms/azure_ai/anthropic/count_tokens/handler.py @@ -30,7 +30,7 @@ class AzureAIAnthropicCountTokensHandler(AzureAIAnthropicCountTokensConfig): messages: list[dict[str, Any]], api_key: str, api_base: str, - litellm_params: dict[str, Any] | None = None, + litellm_params: dict[str, object] | None = None, timeout: float | httpx.Timeout | None = None, tools: list[dict[str, Any]] | None = None, system: object = None, diff --git a/litellm/llms/base_llm/guardrail_translation/base_translation.py b/litellm/llms/base_llm/guardrail_translation/base_translation.py index 89ad67f0485..70b8291d32c 100644 --- a/litellm/llms/base_llm/guardrail_translation/base_translation.py +++ b/litellm/llms/base_llm/guardrail_translation/base_translation.py @@ -203,6 +203,11 @@ class BaseTranslation(ABC): def get_streaming_scan_key(self, responses_so_far: Sequence[object]) -> StreamingScanKey | None: return None + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + """The chunks a client left the stream with, closed the way this endpoint ends a stream, so the + end-of-stream scan also inspects tool calls the stream never finished""" + return tuple(responses_so_far) + def build_block_sse_chunks( self, exc: "ModifyResponseException", diff --git a/litellm/llms/base_llm/harness/transformation.py b/litellm/llms/base_llm/harness/transformation.py index 643f808d4d8..d6db3d2ecb1 100644 --- a/litellm/llms/base_llm/harness/transformation.py +++ b/litellm/llms/base_llm/harness/transformation.py @@ -18,7 +18,7 @@ from __future__ import annotations from abc import ABC, abstractmethod from collections.abc import Mapping, Sequence from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, ClassVar, Generic, TypeVar +from typing import TYPE_CHECKING, ClassVar, Generic, TypeVar from litellm.harness.errors import HarnessError, OptionsMismatch from litellm.harness.types import Capabilities, Event, Harness @@ -133,7 +133,7 @@ class BaseCLIHarnessConfig(BaseHarnessConfig[OptionsT], Generic[OptionsT, Stream """Fresh per-turn parser state.""" @abstractmethod - def transform_stream_line(self, line: Mapping[str, Any], state: StreamStateT) -> Sequence[Event]: + def transform_stream_line(self, line: Mapping[str, object], state: StreamStateT) -> Sequence[Event]: """One decoded JSON line from stdout to zero or more events. Pure.""" @abstractmethod diff --git a/litellm/llms/bedrock/batches/handler.py b/litellm/llms/bedrock/batches/handler.py index fd4c3dc1659..e87c5c35fed 100644 --- a/litellm/llms/bedrock/batches/handler.py +++ b/litellm/llms/bedrock/batches/handler.py @@ -1,6 +1,6 @@ from collections.abc import Mapping, Sequence from datetime import datetime -from typing import TYPE_CHECKING, Any, Final, cast +from typing import TYPE_CHECKING, Final, Literal from openai.types.batch import BatchRequestCounts from openai.types.batch import Metadata as OpenAIBatchMetadata @@ -17,7 +17,12 @@ if TYPE_CHECKING: # AWS Bedrock model-invocation-job statuses โ†’ OpenAI Batch statuses. # Mirrors the mapping used by `BedrockBatchesConfig.transform_create_batch_response` # so create / retrieve return consistent statuses. -_BEDROCK_MIJ_STATUS_TO_OPENAI: Final = { +_BEDROCK_MIJ_STATUS_TO_OPENAI: Final[ + Mapping[ + str, + Literal["validating", "failed", "in_progress", "finalizing", "completed", "expired", "cancelling", "cancelled"], + ] +] = { "Submitted": "validating", "Validating": "validating", "Scheduled": "validating", @@ -92,7 +97,7 @@ def _record_counts_from_response(response: Mapping[str, object]) -> BatchRequest ) -def _to_epoch(value: Any) -> int | None: +def _to_epoch(value: object) -> int | None: if value is None: return None if isinstance(value, (int, float)): @@ -349,10 +354,7 @@ class BedrockBatchesHandler: ) bedrock_status: Final = str(response.get("status", "")) - openai_status: Final = cast( - Any, - _BEDROCK_MIJ_STATUS_TO_OPENAI.get(bedrock_status, "in_progress"), - ) + openai_status: Final = _BEDROCK_MIJ_STATUS_TO_OPENAI.get(bedrock_status, "in_progress") input_uri: Final = response.get("inputDataConfig", {}).get("s3InputDataConfig", {}).get("s3Uri", "") output_prefix: Final = response.get("outputDataConfig", {}).get("s3OutputDataConfig", {}).get("s3Uri", "") diff --git a/litellm/llms/bedrock/image_edit/stability_transformation.py b/litellm/llms/bedrock/image_edit/stability_transformation.py index 01e25f4671e..f66dcf130b8 100644 --- a/litellm/llms/bedrock/image_edit/stability_transformation.py +++ b/litellm/llms/bedrock/image_edit/stability_transformation.py @@ -25,6 +25,7 @@ import base64 from typing import TYPE_CHECKING, Any, Final import httpx +from httpx._types import RequestFiles from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig from litellm.llms.bedrock.common_utils import BedrockError @@ -165,7 +166,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, - ) -> tuple[dict, Any]: + ) -> tuple[dict, RequestFiles]: """ Transform OpenAI-style request to Bedrock Stability request format. diff --git a/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py b/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py index e1b06791c9d..12ace32f43b 100644 --- a/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py +++ b/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py @@ -80,9 +80,7 @@ class AmazonTitanImageGenerationConfig: non_default_params: dict, optional_params: dict, ): - from typing import Any - - image_generation_config: Final[dict[str, Any]] = {} + image_generation_config: Final[dict[str, object]] = {} for k, v in non_default_params.items(): if k == "size" and v is not None: width, height = v.split("x") @@ -106,11 +104,9 @@ class AmazonTitanImageGenerationConfig: text: str, optional_params: dict, ) -> AmazonTitanImageGenerationRequestBody: - from typing import Any - image_generation_config = optional_params.pop("imageGenerationConfig", {}) negative_text: Final = optional_params.pop("negativeText", None) - text_to_image_params: Final[dict[str, Any]] = {"text": text} + text_to_image_params: Final = AmazonTitanTextToImageParams(text=text) if negative_text: text_to_image_params["negativeText"] = negative_text task_type: Final = optional_params.pop("taskType", "TEXT_IMAGE") @@ -121,7 +117,7 @@ class AmazonTitanImageGenerationConfig: } return AmazonTitanImageGenerationRequestBody( taskType=task_type, - textToImageParams=AmazonTitanTextToImageParams(**text_to_image_params), + textToImageParams=text_to_image_params, imageGenerationConfig=AmazonNovaCanvasImageGenerationConfig(**image_generation_config), ) diff --git a/litellm/llms/bedrock/rerank/handler.py b/litellm/llms/bedrock/rerank/handler.py index 8847381cbc9..4a8308bb46a 100644 --- a/litellm/llms/bedrock/rerank/handler.py +++ b/litellm/llms/bedrock/rerank/handler.py @@ -2,6 +2,7 @@ import json from typing import TYPE_CHECKING, Any, Final, cast import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging @@ -24,6 +25,8 @@ if TYPE_CHECKING: else: AWSPreparedRequest = Any +_JSON_DICT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + class BedrockRerankHandler(BaseAWSLLM): async def arerank( @@ -55,13 +58,13 @@ class BedrockRerankHandler(BaseAWSLLM): except httpx.TimeoutException: raise BedrockError(status_code=408, message="Timeout error occurred.") - return BedrockRerankConfig()._transform_response(response.json()) + return BedrockRerankConfig()._transform_response(_JSON_DICT.validate_python(response.json())) def rerank( self, model: str, query: str, - documents: list[str | dict[str, Any]], + documents: list[str | dict[str, object]], optional_params: dict, logging_obj: LitellmLogging, top_n: int | None = None, @@ -136,7 +139,7 @@ class BedrockRerankHandler(BaseAWSLLM): api_key="", ) - response_json: Final = response.json() + response_json: Final = _JSON_DICT.validate_python(response.json()) return BedrockRerankConfig()._transform_response(response_json) diff --git a/litellm/llms/bedrock/search/transformation.py b/litellm/llms/bedrock/search/transformation.py index 19ba7d5673b..565fc12aa1e 100644 --- a/litellm/llms/bedrock/search/transformation.py +++ b/litellm/llms/bedrock/search/transformation.py @@ -37,6 +37,7 @@ from collections.abc import Iterator, Mapping, Sequence from typing import Final import httpx +from pydantic import TypeAdapter from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.search.transformation import ( @@ -81,6 +82,8 @@ _SSE_EVENT_SEPARATOR: Final = re.compile(r"\r?\n[ \t]*\r?\n") _SSE_LINE_PREFIXES: Final = ("event:", "data:", ":", "id:", "retry:") +_JSON_VALUE: Final = TypeAdapter(object) + def _gateway_host_match(api_base: str) -> re.Match[str] | None: return _GATEWAY_HOST_PATTERN.fullmatch(httpx.URL(api_base).host) @@ -128,7 +131,7 @@ def _parse_result_items(raw_text: object) -> tuple[Mapping[str, object], ...]: if not isinstance(raw_text, str): return () try: - parsed: Final = json.loads(raw_text) + parsed: Final = _JSON_VALUE.validate_python(json.loads(raw_text)) except json.JSONDecodeError: return () return _result_items(parsed) @@ -147,7 +150,7 @@ def _iter_sse_events(text: str) -> Iterator[Mapping[str, object]]: if not payload: continue try: - parsed = json.loads(payload) + parsed = _JSON_VALUE.validate_python(json.loads(payload)) except json.JSONDecodeError: continue if isinstance(parsed, dict): diff --git a/litellm/llms/black_forest_labs/image_generation/transformation.py b/litellm/llms/black_forest_labs/image_generation/transformation.py index e5c2bdf59e9..9d586fba26a 100644 --- a/litellm/llms/black_forest_labs/image_generation/transformation.py +++ b/litellm/llms/black_forest_labs/image_generation/transformation.py @@ -8,9 +8,11 @@ API Reference: https://docs.bfl.ai/ """ import time +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -36,6 +38,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): """ @@ -218,7 +222,7 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): https://docs.bfl.ai/flux_models/flux_1_1_pro """ # Build request body with prompt - request_body: Final[dict[str, Any]] = { + request_body: Final[dict[str, object]] = { "prompt": prompt, } @@ -275,7 +279,7 @@ class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): message=f"Error parsing BFL response: {e}", ) - result: Final = response_data.get("result", {}) + result: Final = _JSON_OBJECT.validate_python(response_data).get("result", {}) if not model_response.data: model_response.data = [] diff --git a/litellm/llms/clarifai/chat/transformation.py b/litellm/llms/clarifai/chat/transformation.py index a0946254de0..cea6e6c7440 100644 --- a/litellm/llms/clarifai/chat/transformation.py +++ b/litellm/llms/clarifai/chat/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.openai.common_utils import OpenAIError @@ -20,6 +22,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(strict=True, hide_input_in_errors=True)) + class ClarifaiConfig(OpenAIGPTConfig): """ @@ -111,7 +115,7 @@ class ClarifaiConfig(OpenAIGPTConfig): headers=raw_response.headers, ) from e - response: Final = ModelResponse(**completion_response) + response: Final = ModelResponse(**_JSON_OBJECT.validate_python(completion_response)) if response.model is not None: response.model = "clarifai/" + model diff --git a/litellm/llms/claude_code/harness/transformation.py b/litellm/llms/claude_code/harness/transformation.py index ea3cdd67ffb..f75d86c7e63 100644 --- a/litellm/llms/claude_code/harness/transformation.py +++ b/litellm/llms/claude_code/harness/transformation.py @@ -16,6 +16,8 @@ from dataclasses import dataclass, field from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final +from pydantic import ConfigDict, TypeAdapter + from litellm.harness.errors import HarnessError, OptionsMismatch from litellm.harness.options import ClaudeCodeOptions from litellm.harness.types import ( @@ -48,6 +50,7 @@ if TYPE_CHECKING: CLAUDE_BINARY: Final = "claude" SYNTHETIC_MODEL: Final = "" +_BLOCK: Final = TypeAdapter(Mapping[object, object], config=ConfigDict(hide_input_in_errors=True)) BASE_COMMAND: Final = ("-p", "--output-format", "stream-json", "--verbose", "--input-format", "text") @@ -157,7 +160,7 @@ def _stringify_block(block: object) -> str: return json.dumps(block, ensure_ascii=False) -def _message_blocks(event: Mapping[str, object]) -> Sequence[Any]: +def _message_blocks(event: Mapping[str, object]) -> Sequence[object]: message: Final = event.get("message") content: Final = message.get("content") if isinstance(message, Mapping) else None if isinstance(content, str): @@ -211,8 +214,7 @@ def _user_events(event: Mapping[str, object]) -> Sequence[Event]: output=stringify_tool_output(block.get("content")), is_error=bool(block.get("is_error", False)), ) - for block in _message_blocks(event) - if _is_tool_result(block) + for block in map(_BLOCK.validate_python, filter(_is_tool_result, _message_blocks(event))) ) ) diff --git a/litellm/llms/codex/harness/transformation.py b/litellm/llms/codex/harness/transformation.py index ab17cf3d869..86a5f8274e6 100644 --- a/litellm/llms/codex/harness/transformation.py +++ b/litellm/llms/codex/harness/transformation.py @@ -17,6 +17,8 @@ from dataclasses import dataclass, field from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final +from pydantic import ConfigDict, TypeAdapter + from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.harness.errors import HarnessError, OptionsMismatch from litellm.harness.options import CodexOptions @@ -61,6 +63,7 @@ MANAGED_CONFIG_KEYS: Final = frozenset( ) _BARE_TOML_KEY: Final = re.compile(r"^[A-Za-z0-9_-]+$") _TOOL_ITEM_TYPES: Final = frozenset({"command_execution", "file_change", "web_search", "mcp_tool_call"}) +_CHANGES: Final = TypeAdapter(tuple[Mapping[object, object], ...], config=ConfigDict(hide_input_in_errors=True)) @dataclass @@ -92,7 +95,7 @@ def _tool_input(item: Mapping[str, Any]) -> tuple[str, str, Mapping[str, Any], b return name, tool, tool_args, False -def _tool_output(item: Mapping[str, Any]) -> tuple[str, bool]: +def _tool_output(item: Mapping[str, object]) -> tuple[str, bool]: """(output text, is_error) for a completed tool-like item.""" item_type = item.get("type") status = item.get("status") @@ -101,7 +104,8 @@ def _tool_output(item: Mapping[str, Any]) -> tuple[str, bool]: is_error = status == "failed" or (exit_code is not None and exit_code != 0) return str(item.get("aggregated_output") or ""), is_error if item_type == "file_change": - lines = (f"{c.get('kind', '')} {c.get('path', '')}".strip() for c in item.get("changes") or ()) + changes: Final = _CHANGES.validate_python(item.get("changes") or ()) + lines: Final = (f"{c.get('kind', '')} {c.get('path', '')}".strip() for c in changes) return "\n".join(lines), status == "failed" if item_type == "web_search": return "", status == "failed" @@ -118,7 +122,7 @@ def _tool_output(item: Mapping[str, Any]) -> tuple[str, bool]: def _tool_item_events( - item_id: str, item: Mapping[str, Any], completed: bool, state: CodexStreamState + item_id: str, item: Mapping[str, object], completed: bool, state: CodexStreamState ) -> Iterator[Event]: if item_id not in state.started: state.started.add(item_id) @@ -129,7 +133,7 @@ def _tool_item_events( yield ToolResult(id=item_id, output=output, is_error=is_error) -def _item_events(event_type: str, item: Mapping[str, Any], state: CodexStreamState) -> Sequence[Event]: +def _item_events(event_type: str, item: Mapping[str, object], state: CodexStreamState) -> Sequence[Event]: item_type = item.get("type") item_id = str(item.get("id") or "") completed = event_type == "item.completed" @@ -181,7 +185,7 @@ def _config_override(key: object, value: object) -> str: return f"{dotted}={toml_value(value)}" -def config_overrides(config: Mapping[str, Any]) -> Sequence[str]: +def config_overrides(config: Mapping[str, object]) -> Sequence[str]: """`-c` override strings for CodexOptions.config, rejecting managed keys.""" overrides: Final = (_config_override(key, value) for key, value in config.items()) return list(overrides) # mutable-ok: public helper; tests compare to a list @@ -310,7 +314,7 @@ class CodexHarnessConfig(BaseCLIHarnessConfig): def create_stream_state(self) -> CodexStreamState: return CodexStreamState() - def transform_stream_line(self, line: Mapping[str, Any], state: CodexStreamState) -> Sequence[Event]: + def transform_stream_line(self, line: Mapping[str, object], state: CodexStreamState) -> Sequence[Event]: """turn.completed usage is ignored on purpose: the session endpoint accounts it.""" event_type = line.get("type") if event_type == "thread.started": diff --git a/litellm/llms/cohere/rerank/transformation.py b/litellm/llms/cohere/rerank/transformation.py index a8e755406d8..b6d94d4853a 100644 --- a/litellm/llms/cohere/rerank/transformation.py +++ b/litellm/llms/cohere/rerank/transformation.py @@ -1,7 +1,8 @@ from collections.abc import Mapping -from typing import Any, Final +from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -12,6 +13,8 @@ from litellm.types.rerank import OptionalRerankParams, RerankRequest, RerankResp from ..common_utils import CohereError +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class CohereRerankConfig(BaseRerankConfig): """ @@ -51,7 +54,7 @@ class CohereRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, Any]], + documents: list[str | dict[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, rank_fields: list[str] | None = None, @@ -148,7 +151,7 @@ class CohereRerankConfig(BaseRerankConfig): except Exception: raise CohereError(message=raw_response.text, status_code=raw_response.status_code) - return RerankResponse(**raw_response_json) + return RerankResponse.model_validate(_JSON_OBJECT.validate_python(raw_response_json)) def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return CohereError(message=error_message, status_code=status_code) diff --git a/litellm/llms/cometapi/image_generation/transformation.py b/litellm/llms/cometapi/image_generation/transformation.py index 4432c151a64..567dbfd308e 100644 --- a/litellm/llms/cometapi/image_generation/transformation.py +++ b/litellm/llms/cometapi/image_generation/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -20,6 +22,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class CometAPIImageGenerationConfig(BaseImageGenerationConfig): DEFAULT_BASE_URL: str = "https://api.cometapi.com" @@ -155,7 +160,8 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig): # CometAPI returns OpenAI-compatible format # Expected format: {"created": timestamp, "data": [{"url": "...", "b64_json": "..."}]} if "data" in response_data: - for image_data in response_data["data"]: + payload: Final = _JSON_OBJECT.validate_python(response_data) + for image_data in _JSON_OBJECTS.validate_python(payload["data"]): image_obj = ImageObject( b64_json=image_data.get("b64_json"), url=image_data.get("url"), diff --git a/litellm/llms/dashscope/image_generation/transformation.py b/litellm/llms/dashscope/image_generation/transformation.py index ffa60a3d9bc..6a172e1e074 100644 --- a/litellm/llms/dashscope/image_generation/transformation.py +++ b/litellm/llms/dashscope/image_generation/transformation.py @@ -23,9 +23,11 @@ Response format: } """ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -45,6 +47,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + DEFAULT_API_BASE: Final = "https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation" CHAT_COMPATIBLE_MODE_PATH: Final = "/compatible-mode/v1" @@ -192,9 +197,10 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig): # DashScope can return API-level errors in a 200 response body. # Example: {"code": "InvalidParameter", "message": "Size not supported"} - if "code" in response_data and "output" not in response_data: + response_object: Final = _JSON_OBJECT.validate_python(response_data) + if "code" in response_object and "output" not in response_object: raise self.get_error_class( - error_message=str(response_data.get("message", response_data)), + error_message=str(response_object.get("message", response_object)), status_code=raw_response.status_code, headers=raw_response.headers, ) @@ -202,9 +208,11 @@ class DashScopeImageGenerationConfig(BaseImageGenerationConfig): if not model_response.data: model_response.data = [] - choices: Final = response_data.get("output", {}).get("choices", []) + output: Final = _JSON_OBJECT.validate_python(response_object.get("output", {})) + choices: Final = _JSON_OBJECTS.validate_python(output.get("choices", [])) for choice in choices: - content_list = choice.get("message", {}).get("content", []) + message = _JSON_OBJECT.validate_python(choice.get("message", {})) + content_list = _JSON_OBJECTS.validate_python(message.get("content", [])) for content_item in content_list: image_url = content_item.get("image") if image_url: diff --git a/litellm/llms/dashscope/rerank/transformation.py b/litellm/llms/dashscope/rerank/transformation.py index 14ad756ec9c..4ec7d697f02 100644 --- a/litellm/llms/dashscope/rerank/transformation.py +++ b/litellm/llms/dashscope/rerank/transformation.py @@ -26,10 +26,11 @@ as supported only for gte-rerank-v2 / qwen3-vl-rerank. Docs - https://help.aliyun.com/zh/model-studio/text-rerank-api """ -from collections.abc import Mapping -from typing import Any, Final +from collections.abc import Iterable, Mapping +from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm._uuid import uuid from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -48,6 +49,11 @@ from ..common_utils import DashScopeError, resolve_dashscope_family_rerank_api_b DEFAULT_RERANK_URL: Final = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks" +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_OPTIONAL_INT: Final[TypeAdapter[int | None]] = TypeAdapter(int | None) +_STR: Final = TypeAdapter(str) + class DashScopeRerankConfig(BaseRerankConfig): """ @@ -117,7 +123,7 @@ class DashScopeRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, Any]], + documents: list[str | dict[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, rank_fields: list[str] | None = None, @@ -197,7 +203,8 @@ class DashScopeRerankConfig(BaseRerankConfig): message=response_json.get("message", str(response_json)), ) - results: Final = response_json.get("results") + payload: Final = _JSON_OBJECT.validate_python(response_json) + results: Final = payload.get("results") if results is None: raise DashScopeError( status_code=raw_response.status_code, @@ -210,7 +217,7 @@ class DashScopeRerankConfig(BaseRerankConfig): # "document": {"text": "..."} # which already matches LiteLLM's RerankResponseDocument shape. transformed_results: Final[list[dict]] = [] - for r in results: + for r in _JSON_OBJECTS.validate_python(results): item: dict[str, object] = { "index": r["index"], "relevance_score": r["relevance_score"], @@ -223,14 +230,14 @@ class DashScopeRerankConfig(BaseRerankConfig): item["document"] = {"text": doc} transformed_results.append(item) - usage: Final = response_json.get("usage") or {} - total_tokens: Final = usage.get("total_tokens") + usage: Final = _JSON_OBJECT.validate_python(payload.get("usage") or {}) + total_tokens: Final = _OPTIONAL_INT.validate_python(usage.get("total_tokens")) billed_units: Final = RerankBilledUnits(total_tokens=total_tokens) tokens: Final = RerankTokens(input_tokens=total_tokens) meta: Final = RerankResponseMeta(billed_units=billed_units, tokens=tokens) return RerankResponse( - id=response_json.get("id") or str(uuid.uuid4()), + id=_STR.validate_python(payload.get("id") or str(uuid.uuid4())), results=transformed_results, meta=meta, ) diff --git a/litellm/llms/deepagents/harness/transformation.py b/litellm/llms/deepagents/harness/transformation.py index 14f87f0753b..1f8728230bc 100644 --- a/litellm/llms/deepagents/harness/transformation.py +++ b/litellm/llms/deepagents/harness/transformation.py @@ -173,7 +173,7 @@ def stream_events( ] -def tool_call_event(call: Mapping[str, Any]) -> ToolCall: +def tool_call_event(call: Mapping[str, object]) -> ToolCall: native = str(call.get("name") or "") args = call.get("args") return ToolCall( @@ -185,7 +185,7 @@ def tool_call_event(call: Mapping[str, Any]) -> ToolCall: ) -def _node_messages(update: Mapping[Any, Any]) -> Iterator[object]: +def _node_messages(update: Mapping[object, object]) -> Iterator[object]: for node, delta in update.items(): if node not in _EVENT_NODES or not isinstance(delta, Mapping): continue @@ -231,7 +231,7 @@ def interrupts_in( return list(items) # mutable-ok: list return; callers/tests compare to lists -def final_ai_text(messages: Sequence[Any]) -> str: +def final_ai_text(messages: Sequence[object]) -> str: for message in reversed(messages): if getattr(message, "type", None) == "ai": text = content_text(getattr(message, "content", "")) diff --git a/litellm/llms/deepgram/audio_transcription/transformation.py b/litellm/llms/deepgram/audio_transcription/transformation.py index 95bdd406c3f..c5c9fe12423 100644 --- a/litellm/llms/deepgram/audio_transcription/transformation.py +++ b/litellm/llms/deepgram/audio_transcription/transformation.py @@ -2,10 +2,12 @@ Translates from OpenAI's `/v1/audio/transcriptions` to Deepgram's `/v1/listen` """ +from collections.abc import Iterable, Mapping from typing import Final from urllib.parse import urlencode from httpx import Headers, Response +from pydantic import ConfigDict, TypeAdapter from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -22,6 +24,9 @@ from ...base_llm.audio_transcription.transformation import ( ) from ..common_utils import DeepgramException +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: @@ -105,7 +110,7 @@ class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig): response["task"] = "transcribe" # Use detected_language if available, otherwise default to "en" - detected_language: Final = first_channel.get("detected_language") + detected_language: Final = _JSON_OBJECT.validate_python(first_channel).get("detected_language") response["language"] = detected_language if detected_language else "en" response["duration"] = response_json["metadata"]["duration"] @@ -114,7 +119,7 @@ class DeepgramAudioTranscriptionConfig(BaseAudioTranscriptionConfig): if "words" in first_alternative: response["words"] = [ {"word": word["word"], "start": word["start"], "end": word["end"]} - for word in first_alternative["words"] + for word in _JSON_OBJECTS.validate_python(first_alternative["words"]) ] # Store full response in hidden params diff --git a/litellm/llms/e2b/sandbox/transformation.py b/litellm/llms/e2b/sandbox/transformation.py index 4928ca0c092..d2702b116dc 100644 --- a/litellm/llms/e2b/sandbox/transformation.py +++ b/litellm/llms/e2b/sandbox/transformation.py @@ -8,9 +8,11 @@ Talks to e2b's REST API directly over httpx (no e2b SDK dependency): """ import json +from collections.abc import Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.sandbox.transformation import ( SANDBOX_MAX_OUTPUT_BYTES, @@ -32,6 +34,12 @@ JUPYTER_PORT: Final = 49999 DEFAULT_SANDBOX_TIMEOUT: Final = 300 MAX_OUTPUT_BYTES: Final = SANDBOX_MAX_OUTPUT_BYTES +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_MESSAGE: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter( + Mapping[str, object] | None, config=ConfigDict(hide_input_in_errors=True) +) +_TEXT: Final = TypeAdapter(str, config=ConfigDict(strict=True, hide_input_in_errors=True)) + class E2BSandboxConfig(BaseSandboxConfig): def _http(self, client: AsyncHTTPHandler | None) -> AsyncHTTPHandler: @@ -73,12 +81,14 @@ class E2BSandboxConfig(BaseSandboxConfig): headers={"X-API-Key": key, "Content-Type": "application/json"}, json=body, ) - data: Final = response.json() + data: Final = _JSON_OBJECT.validate_python(response.json()) - handle: Final = ContainerHandle( - id=data["sandboxID"], - provider="e2b", - domain=data.get("domain") or E2B_DEFAULT_DOMAIN, + handle: Final = ContainerHandle.model_validate( + { + "id": data["sandboxID"], + "provider": "e2b", + "domain": data.get("domain") or E2B_DEFAULT_DOMAIN, + } ) handle._hidden_params = { "envd_access_token": data.get("envdAccessToken"), @@ -158,7 +168,7 @@ class E2BSandboxConfig(BaseSandboxConfig): def _parse_lines(lines: list[str]) -> CodeExecutionResult: def _try_parse(stripped: str): try: - return json.loads(stripped) + return _MESSAGE.validate_python(json.loads(stripped)) except json.JSONDecodeError: return None @@ -178,10 +188,12 @@ class E2BSandboxConfig(BaseSandboxConfig): None, ) - return CodeExecutionResult( - stdout="".join(m.get("text", "") for m in of_type("stdout")), - stderr="".join(m.get("text", "") for m in of_type("stderr")), - results=[{k: v for k, v in m.items() if k != "type"} for m in of_type("result")], - error=error, - execution_count=execution_count, + return CodeExecutionResult.model_validate( + { + "stdout": "".join(_TEXT.validate_python(m.get("text", "")) for m in of_type("stdout")), + "stderr": "".join(_TEXT.validate_python(m.get("text", "")) for m in of_type("stderr")), + "results": [{k: v for k, v in m.items() if k != "type"} for m in of_type("result")], + "error": error, + "execution_count": execution_count, + } ) diff --git a/litellm/llms/elevenlabs/audio_transcription/transformation.py b/litellm/llms/elevenlabs/audio_transcription/transformation.py index a2b25760aac..0dbf61917e2 100644 --- a/litellm/llms/elevenlabs/audio_transcription/transformation.py +++ b/litellm/llms/elevenlabs/audio_transcription/transformation.py @@ -2,9 +2,11 @@ Translates from OpenAI's `/v1/audio/transcriptions` to ElevenLabs's `/v1/speech-to-text` """ +from collections.abc import Iterable, Mapping from typing import Final from httpx import Headers, Response +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.litellm_core_utils.audio_utils.utils import process_audio_file @@ -22,6 +24,9 @@ from ...base_llm.audio_transcription.transformation import ( ) from ..common_utils import ElevenLabsException +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class ElevenLabsAudioTranscriptionConfig(BaseAudioTranscriptionConfig): @property @@ -115,21 +120,22 @@ class ElevenLabsAudioTranscriptionConfig(BaseAudioTranscriptionConfig): """ try: response_json: Final = raw_response.json() + response_object: Final = _JSON_OBJECT.validate_python(response_json) # Extract the main transcript text - text: Final = response_json.get("text", "") + text: Final = response_object.get("text", "") # Create TranscriptionResponse object response: Final = TranscriptionResponse(text=text) # Add additional metadata matching OpenAI format response["task"] = "transcribe" - response["language"] = response_json.get("language_code", "unknown") + response["language"] = response_object.get("language_code", "unknown") # Map ElevenLabs words to OpenAI format - if "words" in response_json: + if "words" in response_object: response["words"] = [] - for word_data in response_json["words"]: + for word_data in _JSON_OBJECTS.validate_python(response_object["words"]): # Only include actual words, skip spacing and audio events if word_data.get("type") == "word": response["words"].append( diff --git a/litellm/llms/fal_ai/image_generation/flux_pro_v11_ultra_transformation.py b/litellm/llms/fal_ai/image_generation/flux_pro_v11_ultra_transformation.py index 7c63e1077f1..d4d2941db43 100644 --- a/litellm/llms/fal_ai/image_generation/flux_pro_v11_ultra_transformation.py +++ b/litellm/llms/fal_ai/image_generation/flux_pro_v11_ultra_transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams from litellm.types.utils import ImageResponse @@ -15,6 +17,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class FalAIFluxProV11UltraConfig(FalAIBaseConfig): """ @@ -228,16 +232,17 @@ class FalAIFluxProV11UltraConfig(FalAIBaseConfig): if not model_response.data: model_response.data = [] - images: Final = response_data.get("images", []) + response_object: Final = _JSON_OBJECT.validate_python(response_data) + images: Final = response_object.get("images", []) model_response.data.extend(fal_images_to_image_objects(images)) # Add additional metadata from Flux Pro response if hasattr(model_response, "_hidden_params"): - if "seed" in response_data: - model_response._hidden_params["seed"] = response_data["seed"] - if "timings" in response_data: - model_response._hidden_params["timings"] = response_data["timings"] - if "has_nsfw_concepts" in response_data: - model_response._hidden_params["has_nsfw_concepts"] = response_data["has_nsfw_concepts"] + if "seed" in response_object: + model_response._hidden_params["seed"] = response_object["seed"] + if "timings" in response_object: + model_response._hidden_params["timings"] = response_object["timings"] + if "has_nsfw_concepts" in response_object: + model_response._hidden_params["has_nsfw_concepts"] = response_object["has_nsfw_concepts"] return model_response diff --git a/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py b/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py index ad1852a622b..970addbd5a2 100644 --- a/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py +++ b/litellm/llms/fal_ai/image_generation/ideogram_v3_transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams from litellm.types.utils import ImageObject, ImageResponse @@ -15,6 +17,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class FalAIIdeogramV3Config(FalAIBaseConfig): """ @@ -169,7 +173,8 @@ class FalAIIdeogramV3Config(FalAIBaseConfig): if not model_response.data: model_response.data = [] - images: Final = response_data.get("images", []) + response_object: Final = _JSON_OBJECT.validate_python(response_data) + images: Final = response_object.get("images", []) if isinstance(images, list): for image_entry in images: if isinstance(image_entry, dict): @@ -184,7 +189,7 @@ class FalAIIdeogramV3Config(FalAIBaseConfig): ) ) - if hasattr(model_response, "_hidden_params") and "seed" in response_data: - model_response._hidden_params["seed"] = response_data["seed"] + if hasattr(model_response, "_hidden_params") and "seed" in response_object: + model_response._hidden_params["seed"] = response_object["seed"] return model_response diff --git a/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py b/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py index 934ce420d53..038a06c5749 100644 --- a/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py +++ b/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams from litellm.types.utils import ImageObject, ImageResponse @@ -15,6 +17,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class FalAIRecraftV3Config(FalAIBaseConfig): """ @@ -89,7 +93,7 @@ class FalAIRecraftV3Config(FalAIBaseConfig): return optional_params - def _map_image_size(self, size: str) -> Any: + def _map_image_size(self, size: str) -> str | Mapping[str, int]: """ Map OpenAI size format to Recraft v3 image_size format. @@ -203,7 +207,7 @@ class FalAIRecraftV3Config(FalAIBaseConfig): model_response.data = [] # Handle Recraft v3 response format - images: Final = response_data.get("images", []) + images: Final = _JSON_OBJECT.validate_python(response_data).get("images", []) if isinstance(images, list): for image_data in images: if isinstance(image_data, dict): diff --git a/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py b/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py index 79d8800773b..659011ae537 100644 --- a/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py +++ b/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams from litellm.types.utils import ImageObject, ImageResponse @@ -15,6 +17,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class FalAIStableDiffusionConfig(FalAIBaseConfig): """ @@ -124,7 +128,7 @@ class FalAIStableDiffusionConfig(FalAIBaseConfig): return optional_params - def _map_image_size(self, size: str) -> Any: + def _map_image_size(self, size: str) -> str | Mapping[str, int]: """ Map OpenAI size format to Stable Diffusion image_size format. @@ -243,7 +247,8 @@ class FalAIStableDiffusionConfig(FalAIBaseConfig): model_response.data = [] # Handle Stable Diffusion response format - images: Final = response_data.get("images", []) + response_object: Final = _JSON_OBJECT.validate_python(response_data) + images: Final = response_object.get("images", []) if isinstance(images, list): for image_data in images: if isinstance(image_data, dict): @@ -264,11 +269,11 @@ class FalAIStableDiffusionConfig(FalAIBaseConfig): # Add additional metadata from Stable Diffusion response if hasattr(model_response, "_hidden_params"): - if "seed" in response_data: - model_response._hidden_params["seed"] = response_data["seed"] - if "timings" in response_data: - model_response._hidden_params["timings"] = response_data["timings"] - if "has_nsfw_concepts" in response_data: - model_response._hidden_params["has_nsfw_concepts"] = response_data["has_nsfw_concepts"] + if "seed" in response_object: + model_response._hidden_params["seed"] = response_object["seed"] + if "timings" in response_object: + model_response._hidden_params["timings"] = response_object["timings"] + if "has_nsfw_concepts" in response_object: + model_response._hidden_params["has_nsfw_concepts"] = response_object["has_nsfw_concepts"] return model_response diff --git a/litellm/llms/fireworks_ai/rerank/transformation.py b/litellm/llms/fireworks_ai/rerank/transformation.py index 509dbd5ff24..3d9d813c6e9 100644 --- a/litellm/llms/fireworks_ai/rerank/transformation.py +++ b/litellm/llms/fireworks_ai/rerank/transformation.py @@ -4,10 +4,11 @@ Fireworks AI Rerank API transformation Reference: https://docs.fireworks.ai/inference-api-reference/rerank """ -from collections.abc import Mapping -from typing import Any, Final +from collections.abc import Iterable, Mapping, Sequence +from typing import Final import httpx +from pydantic import BaseModel, ConfigDict, TypeAdapter from litellm._uuid import uuid from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -23,6 +24,26 @@ from litellm.types.rerank import ( ) +class _FireworksAIUsageFields(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True) + + total_tokens: int | None = 0 + prompt_tokens: int | None = 0 + completion_tokens: int | None = 0 + + +class _FireworksAIResultFields(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True, hide_input_in_errors=True) + + index: int | float | str + relevance_score: int | float | str + + +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_STR: Final = TypeAdapter(str) + + class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): """ Fireworks AI Rerank API configuration @@ -59,7 +80,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, Any]], + documents: Sequence[str | Mapping[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, rank_fields: list[str] | None = None, @@ -204,23 +225,18 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): # } # Extract usage information - usage: Final = raw_response_json.get("usage", {}) - _billed_units: Final = RerankBilledUnits(search_units=usage.get("total_tokens", 0)) - _tokens: Final = RerankTokens( - input_tokens=usage.get("prompt_tokens", 0), - output_tokens=usage.get("completion_tokens", 0), - ) - rerank_meta: Final = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) + response_json: Final = _JSON_OBJECT.validate_python(raw_response_json) + usage: Final = _JSON_OBJECT.validate_python(response_json.get("usage", {})) # Extract results - Fireworks AI uses "data" instead of "results" - _results: Final[list[dict] | None] = raw_response_json.get("data") or raw_response_json.get("results") + _results: Final = response_json.get("data") or response_json.get("results") if _results is None: raise ValueError(f"No results found in the response={raw_response_json}") rerank_results: Final[list[RerankResponseResult]] = [] - for result in _results: + for result in _JSON_OBJECTS.validate_python(_results): # Validate required fields exist if not all(key in result for key in ["index", "relevance_score"]): raise ValueError(f"Missing required fields in the result={result}") @@ -239,9 +255,10 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): document = RerankResponseDocument(text=str(text)) # Create typed result + fields = _FireworksAIResultFields.model_validate(result) rerank_result = RerankResponseResult( - index=int(result["index"]), - relevance_score=float(result["relevance_score"]), + index=int(fields.index), + relevance_score=float(fields.relevance_score), ) # Only add document if it exists @@ -250,7 +267,15 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig): rerank_results.append(rerank_result) - response_id: Final = raw_response_json.get("id") or str(uuid.uuid4()) + usage_fields: Final = _FireworksAIUsageFields.model_validate(usage) + _billed_units: Final = RerankBilledUnits(search_units=usage_fields.total_tokens) + _tokens: Final = RerankTokens( + input_tokens=usage_fields.prompt_tokens, + output_tokens=usage_fields.completion_tokens, + ) + rerank_meta: Final = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) + + response_id: Final = _STR.validate_python(response_json.get("id") or str(uuid.uuid4())) return RerankResponse( id=response_id, diff --git a/litellm/llms/gemini/image_generation/transformation.py b/litellm/llms/gemini/image_generation/transformation.py index bb8d7455031..4d13d0d2e6b 100644 --- a/litellm/llms/gemini/image_generation/transformation.py +++ b/litellm/llms/gemini/image_generation/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -31,6 +33,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_JSON_DICT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + class GoogleImageGenConfig(BaseImageGenerationConfig): DEFAULT_BASE_URL: str = "https://generativelanguage.googleapis.com/v1beta" @@ -216,13 +221,15 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): # Extract usage metadata for Gemini models if "usageMetadata" in response_data: - model_response.usage = transform_gemini_image_usage(response_data["usageMetadata"]) + model_response.usage = transform_gemini_image_usage( + _JSON_DICT.validate_python(response_data["usageMetadata"]) + ) web_search_requests: Final = get_gemini_image_web_search_requests(response_data) if web_search_requests and model_response.usage is not None: setattr(model_response.usage, "web_search_requests", web_search_requests) else: # Original Imagen format - predictions with generated images - predictions: Final = response_data.get("predictions", []) + predictions: Final = _JSON_OBJECTS.validate_python(response_data.get("predictions", [])) for prediction in predictions: # Google AI returns base64 encoded images in the prediction model_response.data.append( diff --git a/litellm/llms/gigachat/embedding/transformation.py b/litellm/llms/gigachat/embedding/transformation.py index 927f5e944b6..321f5525433 100644 --- a/litellm/llms/gigachat/embedding/transformation.py +++ b/litellm/llms/gigachat/embedding/transformation.py @@ -8,9 +8,11 @@ API Documentation: https://developers.sber.ru/docs/ru/gigachat/api/reference/res from __future__ import annotations import types +from collections.abc import Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm import LlmProviders from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -22,6 +24,8 @@ from litellm.types.utils import EmbeddingResponse from ..authenticator import get_access_token +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class GigaChatEmbeddingError(BaseLLMException): """GigaChat Embedding API error.""" @@ -165,7 +169,7 @@ class GigaChatEmbeddingConfig(BaseEmbeddingConfig): "total_tokens": total_tokens, } - return EmbeddingResponse(**response_json) + return EmbeddingResponse.model_validate(_JSON_OBJECT.validate_python(response_json)) def validate_environment( self, diff --git a/litellm/llms/github_copilot/embedding/transformation.py b/litellm/llms/github_copilot/embedding/transformation.py index 7ea7a89b4ca..c1acf6fccaa 100644 --- a/litellm/llms/github_copilot/embedding/transformation.py +++ b/litellm/llms/github_copilot/embedding/transformation.py @@ -11,6 +11,7 @@ import os from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm._logging import verbose_logger from litellm.exceptions import AuthenticationError @@ -33,6 +34,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_RESPONSE_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): """ @@ -154,7 +157,7 @@ class GithubCopilotEmbeddingConfig(BaseEmbeddingConfig): logging_obj.post_call(original_response=raw_response.text) # GitHub Copilot returns standard OpenAI-compatible embedding response - response_json: Final = raw_response.json() + response_json: Final = _RESPONSE_OBJECT.validate_python(raw_response.json()) return convert_to_model_response_object( response_object=response_json, diff --git a/litellm/llms/jina_ai/embedding/transformation.py b/litellm/llms/jina_ai/embedding/transformation.py index 260d9e6e494..e184ad2628b 100644 --- a/litellm/llms/jina_ai/embedding/transformation.py +++ b/litellm/llms/jina_ai/embedding/transformation.py @@ -7,9 +7,11 @@ Docs - https://jina.ai/embeddings/ """ import types +from collections.abc import Mapping from typing import Final, cast import httpx +from pydantic import ConfigDict, TypeAdapter from litellm import LlmProviders from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -22,6 +24,8 @@ from litellm.utils import is_base64_encoded from ..common_utils import JinaAIError +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class JinaAIEmbeddingConfig(BaseEmbeddingConfig): """ @@ -139,7 +143,7 @@ class JinaAIEmbeddingConfig(BaseEmbeddingConfig): additional_args={"complete_input_dict": request_data}, original_response=response_json, ) - return EmbeddingResponse(**response_json) + return EmbeddingResponse.model_validate(_JSON_OBJECT.validate_python(response_json)) def validate_environment( self, diff --git a/litellm/llms/jina_ai/rerank/transformation.py b/litellm/llms/jina_ai/rerank/transformation.py index 2a81c38fe34..f4a3a9f0f09 100644 --- a/litellm/llms/jina_ai/rerank/transformation.py +++ b/litellm/llms/jina_ai/rerank/transformation.py @@ -6,10 +6,11 @@ Why separate file? Make it easy to see how transformation works Docs - https://jina.ai/reranker """ -from collections.abc import Mapping, Sequence +from collections.abc import Iterable, Mapping, Sequence from typing import Final from httpx import URL, Response +from pydantic import ConfigDict, TypeAdapter from litellm._uuid import uuid from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj @@ -23,6 +24,12 @@ from litellm.types.rerank import ( ) from litellm.types.utils import ModelInfo +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_BILLED_UNITS: Final = TypeAdapter(RerankBilledUnits) +_TOKENS: Final = TypeAdapter(RerankTokens) +_STR: Final = TypeAdapter(str) + class JinaAIRerankConfig(BaseRerankConfig): def get_supported_cohere_rerank_params(self, model: str) -> list: @@ -100,13 +107,11 @@ class JinaAIRerankConfig(BaseRerankConfig): logging_obj.post_call(original_response=raw_response.text) - _json_response: Final = raw_response.json() + _json_response: Final = _JSON_OBJECT.validate_python(raw_response.json()) - _billed_units: Final = RerankBilledUnits(**_json_response.get("usage", {})) - _tokens: Final = RerankTokens(**_json_response.get("usage", {})) - rerank_meta: Final = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) + usage: Final = _JSON_OBJECT.validate_python(_json_response.get("usage", {})) - _results: Final[list[dict] | None] = _json_response.get("results") + _results: Final = _json_response.get("results") if _results is None: raise ValueError(f"No results found in the response={_json_response}") @@ -115,7 +120,7 @@ class JinaAIRerankConfig(BaseRerankConfig): # Jina AI returns: {"index": 0, "relevance_score": 0.72, "document": "hello"} # LiteLLM expects: {"index": 0, "relevance_score": 0.72, "document": {"text": "hello"}} transformed_results: Final = [] - for result in _results: + for result in _JSON_OBJECTS.validate_python(_results): transformed_result = { "index": result["index"], "relevance_score": result["relevance_score"], @@ -128,8 +133,12 @@ class JinaAIRerankConfig(BaseRerankConfig): transformed_result["document"] = result["document"] transformed_results.append(transformed_result) + _billed_units: Final = _BILLED_UNITS.validate_python(usage) + _tokens: Final = _TOKENS.validate_python(usage) + rerank_meta: Final = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) + return RerankResponse( - id=_json_response.get("id") or str(uuid.uuid4()), + id=_STR.validate_python(_json_response.get("id") or str(uuid.uuid4())), results=transformed_results, meta=rerank_meta, ) # Return response diff --git a/litellm/llms/manus/files/transformation.py b/litellm/llms/manus/files/transformation.py index 8c7b02f8d92..94b0e62199d 100644 --- a/litellm/llms/manus/files/transformation.py +++ b/litellm/llms/manus/files/transformation.py @@ -11,10 +11,12 @@ Reference: https://open.manus.im/docs/openai-compatibility#file-management """ import time -from typing import Any, Final +from collections.abc import Iterable, Mapping +from typing import Final import httpx from openai.types.file_deleted import FileDeleted +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._logging import verbose_logger @@ -39,6 +41,10 @@ from litellm.types.utils import LlmProviders MANUS_API_BASE: Final = "https://api.manus.im" +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(strict=True, hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_TEXT: Final = TypeAdapter(str, config=ConfigDict(strict=True, hide_input_in_errors=True)) + class ManusFilesConfig(BaseFilesConfig): """ @@ -337,8 +343,7 @@ class ManusFilesConfig(BaseFilesConfig): litellm_params: dict, ) -> FileDeleted: """Transform delete file response.""" - response_json: Final = raw_response.json() - return FileDeleted(**response_json) + return FileDeleted.model_validate(_JSON_OBJECT.validate_python(raw_response.json())) def transform_list_files_request( self, @@ -366,19 +371,20 @@ class ManusFilesConfig(BaseFilesConfig): litellm_params: dict, ) -> list[OpenAIFileObject]: """Transform list files response.""" - response_json: Final = raw_response.json() - files_data: Final = response_json.get("data", []) + response_json: Final = _JSON_OBJECT.validate_python(raw_response.json()) + files_data: Final = _JSON_OBJECTS.validate_python(response_json.get("data", [])) return [self._parse_file_dict(f) for f in files_data] - def _parse_file_dict(self, file_dict: dict[str, Any]) -> OpenAIFileObject: + def _parse_file_dict(self, file_dict: Mapping[str, object]) -> OpenAIFileObject: """Parse a file dict into OpenAIFileObject.""" created_at_str: Final = file_dict.get("created_at", "") if created_at_str: + created_at_text: Final = _TEXT.validate_python(created_at_str) try: created_at = int( time.mktime( time.strptime( - created_at_str.replace("Z", "+00:00")[:19], + created_at_text.replace("Z", "+00:00")[:19], "%Y-%m-%dT%H:%M:%S", ) ) @@ -388,15 +394,17 @@ class ManusFilesConfig(BaseFilesConfig): else: created_at = int(time.time()) - return OpenAIFileObject( - id=file_dict.get("id", ""), - bytes=file_dict.get("bytes", 0), - created_at=created_at, - filename=file_dict.get("filename", ""), - object="file", - purpose=file_dict.get("purpose", "assistants"), - status=file_dict.get("status", "uploaded"), - status_details=file_dict.get("status_details"), + return OpenAIFileObject.model_validate( + { + "id": file_dict.get("id", ""), + "bytes": file_dict.get("bytes", 0), + "created_at": created_at, + "filename": file_dict.get("filename", ""), + "object": "file", + "purpose": file_dict.get("purpose", "assistants"), + "status": file_dict.get("status", "uploaded"), + "status_details": file_dict.get("status_details"), + } ) def transform_file_content_request( diff --git a/litellm/llms/minimax/text_to_speech/transformation.py b/litellm/llms/minimax/text_to_speech/transformation.py index e38a8a2c3a3..1cabc43eb1c 100644 --- a/litellm/llms/minimax/text_to_speech/transformation.py +++ b/litellm/llms/minimax/text_to_speech/transformation.py @@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any, Final import httpx from httpx import Headers +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -26,6 +27,9 @@ else: LiteLLMLoggingObj = Any HttpxBinaryResponseContent = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_STR: Final = TypeAdapter(str, config=ConfigDict(strict=True, hide_input_in_errors=True)) + class MinimaxException(BaseLLMException): """Custom exception for MiniMax API errors""" @@ -299,7 +303,7 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): try: # Parse JSON response - response_json: Final = raw_response.json() + response_json: Final = _JSON_OBJECT.validate_python(raw_response.json()) # MiniMax API response format check # The API can return different structures: @@ -320,7 +324,7 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): # Extract audio data # MiniMax returns audio in "data" field - data: Final = response_json.get("data", {}) + data: Final = _JSON_OBJECT.validate_python(response_json.get("data", {})) # Check if response contains a URL (output_format='url') audio_url: Final = data.get("audio_url", None) @@ -334,7 +338,7 @@ class MinimaxTextToSpeechConfig(BaseTextToSpeechConfig): ) # Get hex-encoded audio data - audio_hex: Final = data.get("audio", "") or response_json.get("audio_file", "") + audio_hex: Final = _STR.validate_python(data.get("audio", "") or response_json.get("audio_file", "") or "") if not audio_hex: raise MinimaxException( diff --git a/litellm/llms/modelscope/image_generation/transformation.py b/litellm/llms/modelscope/image_generation/transformation.py index 3a8a37307d6..08cdbacdad0 100644 --- a/litellm/llms/modelscope/image_generation/transformation.py +++ b/litellm/llms/modelscope/image_generation/transformation.py @@ -6,9 +6,11 @@ Handles transformation between OpenAI-compatible format and ModelScope API forma API Reference: https://modelscope.cn/docs/model-service/API-Inference/intro """ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Final import httpx +from pydantic import ConfigDict, TypeAdapter from typing_extensions import override from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -29,6 +31,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = object +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class ModelScopeImageGenerationConfig(BaseImageGenerationConfig): """ @@ -177,8 +182,10 @@ class ModelScopeImageGenerationConfig(BaseImageGenerationConfig): ) # Check for errors in response - if "error" in response_data: - error_msg: Final = response_data["error"].get("message", str(response_data["error"])) + response_object: Final = _JSON_OBJECT.validate_python(response_data) + if "error" in response_object: + error: Final = _JSON_OBJECT.validate_python(response_object["error"]) + error_msg: Final = error.get("message", str(error)) raise self.get_error_class( error_message=f"ModelScope error: {error_msg}", status_code=raw_response.status_code, @@ -186,7 +193,7 @@ class ModelScopeImageGenerationConfig(BaseImageGenerationConfig): ) # Extract images from response - data_list: Final = response_data.get("data", []) + data_list: Final = _JSON_OBJECTS.validate_python(response_object.get("data", [])) if not model_response.data: model_response.data = [] diff --git a/litellm/llms/nvidia_nim/rerank/transformation.py b/litellm/llms/nvidia_nim/rerank/transformation.py index 15cfdb6bece..c8e1a331c8c 100644 --- a/litellm/llms/nvidia_nim/rerank/transformation.py +++ b/litellm/llms/nvidia_nim/rerank/transformation.py @@ -1,7 +1,8 @@ from collections.abc import Mapping -from typing import Any, Final, Literal +from typing import Final, Literal import httpx +from pydantic import ConfigDict, TypeAdapter from typing_extensions import Required, TypedDict import litellm @@ -44,6 +45,14 @@ class NvidiaNimRerankResponse(TypedDict): rankings: Required[list[NvidiaNimRankingResult]] +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_NUMBER: Final[TypeAdapter[bool | int | float]] = TypeAdapter( + bool | int | float, config=ConfigDict(strict=True, hide_input_in_errors=True) +) +_BILLED_UNITS: Final = TypeAdapter(RerankBilledUnits) +_STR: Final = TypeAdapter(str) + + class NvidiaNimRerankConfig(BaseRerankConfig): """ Reference: https://docs.api.nvidia.com/nim/reference/nvidia-llama-3_2-nv-rerankqa-1b-v2-infer @@ -115,7 +124,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, Any]], + documents: list[str | dict[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, rank_fields: list[str] | None = None, @@ -133,7 +142,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig): Nvidia NIM specific params (passed through as-is from non_default_params): - truncate: How to truncate input if too long (NONE, END) """ - optional_nvidia_nim_rerank_params: Final[dict[str, Any]] = { + optional_nvidia_nim_rerank_params: Final[dict[str, object]] = { "query": query, "documents": documents, } @@ -327,15 +336,18 @@ class NvidiaNimRerankConfig(BaseRerankConfig): # Construct metadata with billed_units # Nvidia NIM uses "usage" field with "total_tokens" - usage: Final = raw_response_json.get("usage", {}) - total_tokens: Final = usage.get("total_tokens", 0) + payload: Final = _JSON_OBJECT.validate_python(raw_response_json) + usage: Final = _JSON_OBJECT.validate_python(payload.get("usage", {})) + total_tokens: Final = _NUMBER.validate_python(usage.get("total_tokens", 0)) - billed_units: Final[RerankBilledUnits] = {"total_tokens": total_tokens if total_tokens > 0 else len(results)} + billed_units: Final = _BILLED_UNITS.validate_python( + {"total_tokens": total_tokens if total_tokens > 0 else len(results)} + ) meta: Final[RerankResponseMeta] = {"billed_units": billed_units} return RerankResponse( - id=raw_response_json.get("id") or str(uuid.uuid4()), + id=_STR.validate_python(payload.get("id") or str(uuid.uuid4())), results=results, meta=meta, ) diff --git a/litellm/llms/nvidia_riva/audio_transcription/transformation.py b/litellm/llms/nvidia_riva/audio_transcription/transformation.py index cc78f5b99ef..262dfb7bd78 100644 --- a/litellm/llms/nvidia_riva/audio_transcription/transformation.py +++ b/litellm/llms/nvidia_riva/audio_transcription/transformation.py @@ -102,7 +102,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): if endpointing_config is not None: recognition_config["endpointing_config"] = endpointing_config - request_payload: Final[dict[str, Any]] = { + request_payload: Final[dict[str, object]] = { "recognition_config": recognition_config, "response_format": optional_params.get("response_format") or "json", "timestamp_granularities": optional_params.get("timestamp_granularities"), @@ -135,7 +135,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): # gRPC auth is constructed in the handler, not via HTTP headers. return headers - def _build_recognition_config_dict(self, model: str, optional_params: dict) -> dict[str, Any]: + def _build_recognition_config_dict(self, model: str, optional_params: dict) -> dict[str, object]: """ Build the Riva ``RecognitionConfig`` shape as a plain dict. @@ -159,7 +159,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): "profanity_filter": optional_params.get("profanity_filter", False), } - def _build_endpointing_config_dict(self, optional_params: dict) -> dict[str, Any] | None: + def _build_endpointing_config_dict(self, optional_params: dict) -> dict[str, object] | None: """ Translate an OpenAI-style ``chunking_strategy`` into Riva's ``EndpointingConfig`` shape, or pass through an explicit @@ -177,7 +177,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): return None if isinstance(chunking, dict) and chunking.get("type") == "server_vad": - config: Final[dict[str, Any]] = {} + config: Final[dict[str, object]] = {} if "threshold" in chunking: threshold: Final = float(chunking["threshold"]) config["start_threshold"] = threshold @@ -245,7 +245,7 @@ class NvidiaRivaAudioTranscriptionConfig(BaseAudioTranscriptionConfig): response["task"] = "transcribe" if response_format == "verbose_json": - words: Final[list[dict[str, Any]]] = [] + words: Final[list[dict[str, object]]] = [] if timestamp_granularities and "word" in timestamp_granularities: for item in final_results: for word in item.get("words", []) or []: diff --git a/litellm/llms/oci/embed/transformation.py b/litellm/llms/oci/embed/transformation.py index dec43717387..edadc293771 100644 --- a/litellm/llms/oci/embed/transformation.py +++ b/litellm/llms/oci/embed/transformation.py @@ -22,9 +22,11 @@ Supported models: Reference: https://docs.oracle.com/en-us/iaas/api/#/en/generative-ai-inference/latest/EmbedTextResult/EmbedText """ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -52,6 +54,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + # OCI sends up to 96 texts per embedText request (Cohere limit). OCI_EMBED_BATCH_LIMIT: Final = 96 @@ -275,7 +279,7 @@ class OCIEmbedConfig(BaseEmbeddingConfig): ) try: - parsed: Final = OCIEmbedResponse(**json_response) + parsed: Final = OCIEmbedResponse.model_validate(_JSON_OBJECT.validate_python(json_response)) except Exception as e: raise OCIError( status_code=500, diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index aa175733582..aee439862db 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -17,7 +17,7 @@ This pattern can be replicated for other message formats (e.g., Anthropic). import json import time import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Iterator, Mapping, Sequence from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Union, cast @@ -54,6 +54,7 @@ from litellm.types.utils import ( ChatCompletionDeltaToolCall, ChatCompletionMessageToolCall, Choices, + Delta, GenericGuardrailAPIInputs, ModelResponse, ModelResponseStream, @@ -837,6 +838,18 @@ class OpenAIChatCompletionsHandler(BaseTranslation): tool_calls_in_flight=bool(tool_call_fingerprints) and not stream_ended, ) + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + released_key: Final = self.get_streaming_scan_key(responses_so_far) + if released_key is None or not released_key.tool_calls_in_flight: + return tuple(responses_so_far) + terminator: Final = ModelResponseStream( + choices=[ + StreamingChoices(index=index, delta=Delta(), finish_reason="tool_calls") + for index in _choice_indices_with_tool_calls(responses_so_far) + ] + ) + return (*responses_so_far, terminator) + @staticmethod def _streamed_tool_call_fingerprints(responses_so_far: Sequence[object]) -> tuple[str, ...]: return tuple( @@ -1388,6 +1401,20 @@ def _streamed_delta_tool_calls(delta: object) -> tuple[object, ...]: return stream_item_items(delta, "tool_calls") + legacy +def _released_choices(responses_so_far: Sequence[object]) -> Iterator[object]: + for chunk in responses_so_far: + yield from _stream_chunk_choices(chunk) + + +def _choice_indices_with_tool_calls(responses_so_far: Sequence[object]) -> tuple[int, ...]: + indices: Final = ( + index if isinstance(index := stream_item_field(choice, "index"), int) else 0 + for choice in _released_choices(responses_so_far) + if _streamed_delta_tool_calls(stream_item_field(choice, "delta")) + ) + return tuple(dict.fromkeys(indices)) + + def _blocked_stream_identity( exc: "ModifyResponseException", responses_so_far: Sequence[object] ) -> tuple[str, int, str]: diff --git a/litellm/llms/openai/responses/count_tokens/handler.py b/litellm/llms/openai/responses/count_tokens/handler.py index 66782a83a19..0d7a553d633 100644 --- a/litellm/llms/openai/responses/count_tokens/handler.py +++ b/litellm/llms/openai/responses/count_tokens/handler.py @@ -5,6 +5,7 @@ Uses httpx for HTTP requests to OpenAI's /v1/responses/input_tokens endpoint. """ import json +from collections.abc import Sequence from typing import Any, Final import httpx @@ -26,11 +27,11 @@ class OpenAICountTokensHandler(OpenAICountTokensConfig): async def handle_count_tokens_request( self, model: str, - input: str | list[Any], + input: str | Sequence[object], api_key: str, api_base: str | None = None, timeout: float | httpx.Timeout | None = None, - tools: list[dict[str, Any]] | None = None, + tools: list[dict[str, object]] | None = None, instructions: str | None = None, ) -> dict[str, Any]: """ diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 90cdef87ec7..ad68924ce20 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -293,6 +293,51 @@ def _is_tool_call_output_item(item: object) -> bool: return _tool_call_output_item_mapping(item) is not None +def _released_tool_call_payload(responses_so_far: Sequence[object], item_id: object) -> str | None: + events: Final = tuple(event for event in responses_so_far if stream_item_field(event, "item_id") == item_id) + finished: Final = tuple( + payload + for event in events + if isinstance(event_type := stream_item_field(event, "type"), str) + and event_type in _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS + and isinstance(payload := stream_item_field(event, _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS[event_type]), str) + ) + if finished: + return finished[-1] + deltas: Final = tuple( + delta + for event in events + if stream_item_field(event, "type") in _TOOL_CALL_PAYLOAD_DELTA_EVENT_TYPES + and isinstance(delta := stream_item_field(event, "delta"), str) + ) + return "".join(deltas) if deltas else None + + +def _with_released_payload(item: Mapping[str, object], responses_so_far: Sequence[object]) -> Mapping[str, object]: + field: Final = _TOOL_CALL_PAYLOAD_FIELDS[str(item.get("type"))] + payload: Final = _released_tool_call_payload(responses_so_far, item.get("id")) + return item if payload is None else {**item, field: payload} + + +def _released_message_item(text: str) -> Mapping[str, object]: + content: Final = [{"type": "output_text", "text": text}] + return {"type": "message", "role": "assistant", "content": content} + + +def _released_tool_call_items(responses_so_far: Sequence[object]) -> tuple[Mapping[str, object], ...]: + announced: Final = tuple( + item + for item in ( + _tool_call_output_item_mapping(stream_item_field(event, "item")) + for event in responses_so_far + if stream_item_field(event, "type") in _OUTPUT_ITEM_EVENT_TYPES + ) + if item is not None + ) + latest_by_id: Final = MappingProxyType({item.get("id"): item for item in announced}) + return tuple(_with_released_payload(item, responses_so_far) for item in latest_by_id.values()) + + def _last_message_role(messages: Sequence[object]) -> str | None: if not messages: return None @@ -1324,6 +1369,27 @@ class OpenAIResponsesHandler(BaseTranslation): tool_calls_in_flight=self._has_streamed_tool_call_events(responses_so_far), ) + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + if self._check_streaming_has_ended(responses_so_far): + return tuple(responses_so_far) + ends_on_finished_item: Final = ( + bool(responses_so_far) + and stream_item_field(responses_so_far[-1], "type") == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE.value + ) + if not ends_on_finished_item and not self._has_streamed_tool_call_events(responses_so_far): + return tuple(responses_so_far) + text_events: Final = tuple( + event for event in responses_so_far if stream_item_field(event, "type") in _OUTPUT_TEXT_EVENT_TYPES + ) + released_text: Final = self.get_streaming_string_so_far(text_events) + message_items: Final = (_released_message_item(released_text),) if released_text else () + tool_items: Final = _released_tool_call_items(responses_so_far) + output: Final = [*message_items, *tool_items] + response: Final = {"status": "incomplete", "output": output} + incomplete: Final = ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE.value + envelope: Final = {"type": incomplete, "response": response} + return (*responses_so_far, envelope) + @staticmethod def _has_streamed_tool_call_events(responses_so_far: Sequence[object]) -> bool: return any( diff --git a/litellm/llms/opencode/harness/transformation.py b/litellm/llms/opencode/harness/transformation.py index af5fa1ae71d..158af0fae0a 100644 --- a/litellm/llms/opencode/harness/transformation.py +++ b/litellm/llms/opencode/harness/transformation.py @@ -165,7 +165,7 @@ def _as_dict(value: object) -> Mapping[str, Any]: return value if isinstance(value, dict) else MappingProxyType({}) -def _tool_events(part: Mapping[str, Any]) -> Sequence[Event]: +def _tool_events(part: Mapping[str, object]) -> Sequence[Event]: native = str(part.get("tool") or "") call_id = str(part.get("callID") or part.get("id") or "") state = _as_dict(part.get("state")) @@ -195,7 +195,7 @@ def _error_message(error: object) -> str: return str(error.get("name") or "opencode reported an error") -def validate_user_config(config: Mapping[str, Any]) -> None: +def validate_user_config(config: Mapping[str, object]) -> None: """Reject OpenCodeOptions.config keys LiteLLM manages (or that bypass permissions).""" for key in config: if key in MANAGED_CONFIG_KEYS: @@ -241,7 +241,7 @@ def build_opencode_config( user_config: Mapping[str, Any] | None = None, instructions_path: str | None = None, skills_path: str | None = None, -) -> Mapping[str, Any]: +) -> Mapping[str, object]: """The full opencode config: user config underneath, LiteLLM-managed keys on top.""" user: Final = user_config or MappingProxyType({}) validate_user_config(user) @@ -379,7 +379,7 @@ class OpenCodeHarnessConfig(BaseCLIHarnessConfig): def create_stream_state(self) -> OpenCodeStreamState: return OpenCodeStreamState() - def transform_stream_line(self, line: Mapping[str, Any], state: OpenCodeStreamState) -> Sequence[Event]: + def transform_stream_line(self, line: Mapping[str, object], state: OpenCodeStreamState) -> Sequence[Event]: """step_finish token counts are ignored on purpose: the session endpoint accounts usage.""" session_id = line.get("sessionID") if session_id and state.session_id is None: diff --git a/litellm/llms/openrouter/image_generation/transformation.py b/litellm/llms/openrouter/image_generation/transformation.py index 67d90d027ec..cfca853e280 100644 --- a/litellm/llms/openrouter/image_generation/transformation.py +++ b/litellm/llms/openrouter/image_generation/transformation.py @@ -27,9 +27,11 @@ Response format: } """ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -55,6 +57,11 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_DICT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_STR: Final = TypeAdapter(str, config=ConfigDict(strict=True, hide_input_in_errors=True)) + class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): """ @@ -355,15 +362,16 @@ class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): model_response.data = [] try: - choices: Final = response_json.get("choices", []) + response_object: Final = _JSON_DICT.validate_python(response_json) + choices: Final = _JSON_OBJECTS.validate_python(response_object.get("choices", [])) for choice in choices: - message = choice.get("message", {}) - images = message.get("images", []) + message = _JSON_OBJECT.validate_python(choice.get("message", {})) + images = _JSON_OBJECTS.validate_python(message.get("images", [])) for image_data in images: - image_url_obj = image_data.get("image_url", {}) - image_url = image_url_obj.get("url") + image_url_obj = _JSON_OBJECT.validate_python(image_data.get("image_url", {})) + image_url = _STR.validate_python(image_url_obj.get("url") or "") if image_url: if image_url.startswith("data:"): @@ -389,7 +397,7 @@ class OpenRouterImageGenerationConfig(BaseImageGenerationConfig): ) # Extract and set usage and cost information - self._set_usage_and_cost(model_response, response_json, model) + self._set_usage_and_cost(model_response, response_object, model) return model_response diff --git a/litellm/llms/ovhcloud/audio_transcription/transformation.py b/litellm/llms/ovhcloud/audio_transcription/transformation.py index 6fc56ebb61f..086c1afbcd4 100644 --- a/litellm/llms/ovhcloud/audio_transcription/transformation.py +++ b/litellm/llms/ovhcloud/audio_transcription/transformation.py @@ -5,9 +5,11 @@ Our unified API follows the OpenAI standard. More information on our website: https://endpoints.ai.cloud.ovh.net """ +from collections.abc import Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.llms.base_llm.audio_transcription.transformation import ( @@ -24,6 +26,8 @@ from litellm.types.utils import FileTypes, TranscriptionResponse from ..utils import OVHCloudException +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class OVHCloudAudioTranscriptionConfig(BaseAudioTranscriptionConfig): def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]: @@ -145,7 +149,8 @@ class OVHCloudAudioTranscriptionConfig(BaseAudioTranscriptionConfig): headers=raw_response.headers, ) - text: Final = response_json.get("text") or response_json.get("transcript") or "" + payload: Final = _JSON_OBJECT.validate_python(response_json) + text: Final = payload.get("text") or payload.get("transcript") or "" response: Final = TranscriptionResponse(text=text) # OVHCloud field migration (deadline: 2026-05-11): @@ -153,9 +158,7 @@ class OVHCloudAudioTranscriptionConfig(BaseAudioTranscriptionConfig): # Prefer `seconds`, fall back to `duration`, normalize to `duration` # so downstream consumers see a consistent key. duration: Final = ( - response_json["seconds"] - if "seconds" in response_json and response_json["seconds"] is not None - else response_json.get("duration") + payload["seconds"] if "seconds" in payload and payload["seconds"] is not None else payload.get("duration") ) if duration is not None: response_json["duration"] = duration diff --git a/litellm/llms/recraft/image_generation/transformation.py b/litellm/llms/recraft/image_generation/transformation.py index f65bf1e7292..0fcfafb79fb 100644 --- a/litellm/llms/recraft/image_generation/transformation.py +++ b/litellm/llms/recraft/image_generation/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -21,6 +23,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class RecraftImageGenerationConfig(BaseImageGenerationConfig): DEFAULT_BASE_URL: str = "https://external.api.recraft.ai" @@ -141,7 +146,8 @@ class RecraftImageGenerationConfig(BaseImageGenerationConfig): if not model_response.data: model_response.data = [] - for image_data in response_data["data"]: + payload: Final = _JSON_OBJECT.validate_python(response_data) + for image_data in _JSON_OBJECTS.validate_python(payload["data"]): model_response.data.append( ImageObject( url=image_data.get("url", None), diff --git a/litellm/llms/replicate/chat/transformation.py b/litellm/llms/replicate/chat/transformation.py index f7e09b7bec0..e8dd8816997 100644 --- a/litellm/llms/replicate/chat/transformation.py +++ b/litellm/llms/replicate/chat/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Iterable from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.constants import REPLICATE_MODEL_NAME_WITH_ID_LENGTH @@ -26,6 +28,8 @@ if TYPE_CHECKING: else: LoggingClass = Any +_TEXTS: Final = TypeAdapter(Iterable[str], config=ConfigDict(strict=True, hide_input_in_errors=True)) + class ReplicateConfig(BaseConfig): """ @@ -253,7 +257,7 @@ class ReplicateConfig(BaseConfig): message=f"LiteLLM Error - prediction not succeeded - {raw_response_json}", headers=raw_response.headers, ) - outputs: Final = raw_response_json.get("output", []) + outputs: Final = _TEXTS.validate_python(raw_response_json.get("output", [])) response_str = "".join(outputs) if len(response_str) == 0: # edge case, where result from replicate is empty response_str = " " diff --git a/litellm/llms/sagemaker/common_utils.py b/litellm/llms/sagemaker/common_utils.py index 8e8f7ea61aa..0d86752fc5f 100644 --- a/litellm/llms/sagemaker/common_utils.py +++ b/litellm/llms/sagemaker/common_utils.py @@ -1,9 +1,10 @@ import functools import json -from collections.abc import AsyncIterator, Iterator +from collections.abc import AsyncIterator, Iterator, Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm import verbose_logger @@ -11,6 +12,8 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.types.utils import GenericStreamingChunk as GChunk from litellm.types.utils import StreamingChatCompletionChunk +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + def _load_sagemaker_response_stream_shape(): try: @@ -18,7 +21,7 @@ def _load_sagemaker_response_stream_shape(): from botocore.model import ServiceModel loader: Final = Loader() - service_dict: Final = loader.load_service_model("sagemaker-runtime", "service-2") + service_dict: Final = _JSON_OBJECT.validate_python(loader.load_service_model("sagemaker-runtime", "service-2")) return ServiceModel(service_dict).shape_for("InvokeEndpointWithResponseStreamOutput") except Exception as e: verbose_logger.warning( diff --git a/litellm/llms/snowflake/embedding/transformation.py b/litellm/llms/snowflake/embedding/transformation.py index 75f35a8f379..6aa66de1db5 100644 --- a/litellm/llms/snowflake/embedding/transformation.py +++ b/litellm/llms/snowflake/embedding/transformation.py @@ -1,6 +1,8 @@ +from collections.abc import Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -10,6 +12,8 @@ from litellm.types.utils import EmbeddingResponse from ..utils import SnowflakeBaseConfig, SnowflakeException +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class SnowflakeEmbeddingConfig(SnowflakeBaseConfig, BaseEmbeddingConfig): """ @@ -53,7 +57,7 @@ class SnowflakeEmbeddingConfig(SnowflakeBaseConfig, BaseEmbeddingConfig): # convert embeddings to 1d array for item in response_json["data"]: item["embedding"] = item["embedding"][0] - returned_response: Final = EmbeddingResponse(**response_json) + returned_response: Final = EmbeddingResponse.model_validate(_JSON_OBJECT.validate_python(response_json)) returned_response.model = "snowflake/" + (returned_response.model or "") diff --git a/litellm/llms/stability/image_generation/transformation.py b/litellm/llms/stability/image_generation/transformation.py index 656ffe395c8..1856523ce49 100644 --- a/litellm/llms/stability/image_generation/transformation.py +++ b/litellm/llms/stability/image_generation/transformation.py @@ -6,9 +6,11 @@ Handles transformation between OpenAI-compatible format and Stability AI API for API Reference: https://platform.stability.ai/docs/api-reference """ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -33,6 +35,8 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class StabilityImageGenerationConfig(BaseImageGenerationConfig): """ @@ -234,7 +238,8 @@ class StabilityImageGenerationConfig(BaseImageGenerationConfig): ) # Check finish_reason - finish_reason: Final = response_data.get("finish_reason", "") + payload: Final = _JSON_OBJECT.validate_python(response_data) + finish_reason: Final = payload.get("finish_reason", "") if finish_reason == "CONTENT_FILTERED": raise self.get_error_class( error_message="Content was filtered by Stability AI safety systems", @@ -246,7 +251,7 @@ class StabilityImageGenerationConfig(BaseImageGenerationConfig): model_response.data = [] # Extract image from response - image_b64: Final = response_data.get("image") + image_b64: Final = payload.get("image") if image_b64: model_response.data.append( ImageObject( diff --git a/litellm/llms/tinyfish/search/transformation.py b/litellm/llms/tinyfish/search/transformation.py index d5ed7da3815..c7976456d87 100644 --- a/litellm/llms/tinyfish/search/transformation.py +++ b/litellm/llms/tinyfish/search/transformation.py @@ -26,6 +26,7 @@ from litellm.secret_managers.main import get_secret_str _UrlEncodableParams: Final = TypeAdapter(dict[str, str | int | float | bool]) _StrList: Final = TypeAdapter(list[str]) _StrFrozenSet: Final = TypeAdapter(frozenset[str]) +_DecodedJson: Final = TypeAdapter(object) _TINYFISH_PARAMS_KEY: Final = "_tinyfish_params" _TINYFISH_DOCS_URL: Final = "https://docs.tinyfish.ai/search-api" @@ -218,7 +219,7 @@ class TinyfishSearchConfig(BaseSearchConfig): ) try: - raw_json: Final[object] = raw_response.json() # any-ok: httpx Response.json() -> Any + raw_json: Final = _DecodedJson.validate_python(raw_response.json()) except json.JSONDecodeError: raise self._wrap_error( error_message=f"Expected JSON response, got: {raw_response.text[:200]}", @@ -276,7 +277,7 @@ class TinyfishSearchConfig(BaseSearchConfig): # for other envelope shapes (CDN HTML pages, other JSON envelopes, plain text). inner_message = error_message try: - body: Final[object] = json.loads(error_message) # any-ok: json.loads -> Any + body: Final = _DecodedJson.validate_python(json.loads(error_message)) if isinstance(body, dict): error_obj: Final[object] = body.get("error") # any-ok: untyped dict if isinstance(error_obj, dict): diff --git a/litellm/llms/together_ai/rerank/handler.py b/litellm/llms/together_ai/rerank/handler.py index b8079e52c97..2fff0854f10 100644 --- a/litellm/llms/together_ai/rerank/handler.py +++ b/litellm/llms/together_ai/rerank/handler.py @@ -4,7 +4,9 @@ Re rank api LiteLLM supports the re rank API format, no paramter transformation occurs """ -from typing import Any, Final +from typing import Final + +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.llms.base import BaseLLM @@ -15,6 +17,8 @@ from litellm.llms.custom_httpx.http_handler import ( from litellm.llms.together_ai.rerank.transformation import TogetherAIRerankConfig from litellm.types.rerank import RerankRequest, RerankResponse +_JSON_DICT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) + def _rerank_url(api_base: str) -> str: return f"{api_base.rstrip('/')}/rerank" @@ -27,7 +31,7 @@ class TogetherAIRerank(BaseLLM): api_key: str, api_base: str, query: str, - documents: list[str | dict[str, Any]], + documents: list[str | dict[str, object]], top_n: int | None = None, rank_fields: list[str] | None = None, return_documents: bool | None = True, @@ -66,13 +70,13 @@ class TogetherAIRerank(BaseLLM): if response.status_code != 200: raise Exception(response.text) - _json_response: Final = response.json() + _json_response: Final = _JSON_DICT.validate_python(response.json()) return TogetherAIRerankConfig()._transform_response(_json_response) async def async_rerank( # New async method self, - request_data_dict: dict[str, Any], + request_data_dict: dict[str, object], api_key: str, api_base: str, ) -> RerankResponse: @@ -91,6 +95,6 @@ class TogetherAIRerank(BaseLLM): if response.status_code != 200: raise Exception(response.text) - _json_response: Final = response.json() + _json_response: Final = _JSON_DICT.validate_python(response.json()) return TogetherAIRerankConfig()._transform_response(_json_response) diff --git a/litellm/llms/vertex_ai/image_generation/image_generation_handler.py b/litellm/llms/vertex_ai/image_generation/image_generation_handler.py index 6a5bb484540..7ba7b060991 100644 --- a/litellm/llms/vertex_ai/image_generation/image_generation_handler.py +++ b/litellm/llms/vertex_ai/image_generation/image_generation_handler.py @@ -1,8 +1,10 @@ import json +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final import httpx from openai.types.image import Image +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.llms.custom_httpx.http_handler import ( @@ -17,11 +19,13 @@ from litellm.types.utils import ImageResponse if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +_PREDICTIONS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) + class VertexImageGeneration(VertexLLM): def process_image_generation_response( self, - json_response: dict[str, Any], + json_response: Mapping[str, object], model_response: ImageResponse, model: str | None = None, ) -> ImageResponse: @@ -32,12 +36,11 @@ class VertexImageGeneration(VertexLLM): model=model, ) - predictions: Final = json_response["predictions"] + predictions: Final = _PREDICTIONS.validate_python(json_response["predictions"]) response_data: Final[list[Image]] = [] for prediction in predictions: - bytes_base64_encoded = prediction["bytesBase64Encoded"] - image_object = Image(b64_json=bytes_base64_encoded) + image_object = Image.model_validate({"b64_json": prediction["bytesBase64Encoded"]}) response_data.append(image_object) model_response.data = response_data diff --git a/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py b/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py index 17ccf16837b..9f3126fd4ff 100644 --- a/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py +++ b/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py @@ -1,7 +1,9 @@ import os +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._logging import verbose_logger @@ -31,6 +33,9 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +_JSON_DICT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): """ @@ -321,10 +326,11 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): ) ) - if usage_metadata := response_data.get("usageMetadata", None): - model_response.usage = self._transform_image_usage(usage_metadata) + response_object: Final = _JSON_OBJECT.validate_python(response_data) + if usage_metadata := response_object.get("usageMetadata", None): + model_response.usage = self._transform_image_usage(_JSON_DICT.validate_python(usage_metadata)) - web_search_requests: Final = get_gemini_image_web_search_requests(response_data) + web_search_requests: Final = get_gemini_image_web_search_requests(response_object) if web_search_requests and model_response.usage is not None: setattr(model_response.usage, "web_search_requests", web_search_requests) diff --git a/litellm/llms/vertex_ai/text_to_speech/transformation.py b/litellm/llms/vertex_ai/text_to_speech/transformation.py index a7b079fb89c..a78145fdcf4 100644 --- a/litellm/llms/vertex_ai/text_to_speech/transformation.py +++ b/litellm/llms/vertex_ai/text_to_speech/transformation.py @@ -6,11 +6,12 @@ Reference: https://cloud.google.com/text-to-speech/docs/reference/rest/v1/text/s """ import base64 -from collections.abc import Coroutine +from collections.abc import Coroutine, Mapping from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, TypeAlias, Union import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.exceptions import UnsupportedParamsError @@ -45,6 +46,9 @@ else: _LyriaVoice: TypeAlias = str | dict | None +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_STR: Final = TypeAdapter(str, config=ConfigDict(hide_input_in_errors=True)) + class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): """ @@ -465,14 +469,14 @@ class VertexAITextToSpeechConfig(BaseTextToSpeechConfig, VertexBase): from litellm.types.llms.openai import HttpxBinaryResponseContent # Parse JSON response - _json_response: Final = raw_response.json() + _json_response: Final = _JSON_OBJECT.validate_python(raw_response.json()) # Get base64-encoded audio content response_content: Final = _json_response.get("audioContent") if not response_content: raise ValueError("No audioContent in Vertex AI TTS response") - binary_data: Final = base64.b64decode(response_content) + binary_data: Final = base64.b64decode(_STR.validate_python(response_content)) media_type: Final = speech_media_type_from_audio_bytes(binary_data) response: Final = httpx.Response( status_code=200, diff --git a/litellm/llms/voyage/rerank/transformation.py b/litellm/llms/voyage/rerank/transformation.py index b48efe229b3..3acb2f2ed58 100644 --- a/litellm/llms/voyage/rerank/transformation.py +++ b/litellm/llms/voyage/rerank/transformation.py @@ -4,10 +4,11 @@ Transformation logic for Voyage AI's /v1/rerank endpoint. Docs - https://docs.voyageai.com/docs/reranker """ -from collections.abc import Mapping, Sequence +from collections.abc import Iterable, Mapping, Sequence from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm._uuid import uuid from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj @@ -23,6 +24,11 @@ from litellm.types.utils import ModelInfo from ..embedding.transformation import VoyageError +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_OPTIONAL_INT: Final[TypeAdapter[int | None]] = TypeAdapter(int | None) +_STR: Final = TypeAdapter(str) + class VoyageRerankConfig(BaseRerankConfig): def get_supported_cohere_rerank_params(self, model: str) -> list: @@ -103,13 +109,14 @@ class VoyageRerankConfig(BaseRerankConfig): ) # Voyage AI returns results in "data" key, not "results" - _results: Final[list[dict] | None] = _json_response.get("data") + payload: Final = _JSON_OBJECT.validate_python(_json_response) + _results: Final = payload.get("data") if _results is None: raise ValueError(f"No results found in the response={_json_response}") # Transform to LiteLLM format transformed_results: Final = [] - for result in _results: + for result in _JSON_OBJECTS.validate_python(_results): transformed_result: dict[str, object] = { "index": result["index"], "relevance_score": result["relevance_score"], @@ -121,14 +128,14 @@ class VoyageRerankConfig(BaseRerankConfig): transformed_result["document"] = result["document"] transformed_results.append(transformed_result) - usage: Final = _json_response.get("usage", {}) - total_tokens: Final = usage.get("total_tokens", 0) + usage: Final = _JSON_OBJECT.validate_python(payload.get("usage", {})) + total_tokens: Final = _OPTIONAL_INT.validate_python(usage.get("total_tokens", 0)) _billed_units: Final = RerankBilledUnits(total_tokens=total_tokens) _tokens: Final = RerankTokens(input_tokens=total_tokens, output_tokens=0) rerank_meta: Final = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) return RerankResponse( - id=_json_response.get("id") or str(uuid.uuid4()), + id=_STR.validate_python(payload.get("id") or str(uuid.uuid4())), results=transformed_results, meta=rerank_meta, ) diff --git a/litellm/llms/watsonx/audio_transcription/transformation.py b/litellm/llms/watsonx/audio_transcription/transformation.py index 2169b9bf49a..77ef03b8e10 100644 --- a/litellm/llms/watsonx/audio_transcription/transformation.py +++ b/litellm/llms/watsonx/audio_transcription/transformation.py @@ -4,9 +4,11 @@ Translates from OpenAI's `/v1/audio/transcriptions` to IBM WatsonX's `/ml/v1/aud WatsonX follows the OpenAI spec for audio transcription. """ +from collections.abc import Mapping from typing import Final from httpx import Response +from pydantic import ConfigDict, TypeAdapter import litellm from litellm.litellm_core_utils.audio_utils.utils import process_audio_file @@ -25,6 +27,8 @@ from ...openai.transcriptions.whisper_transformation import ( ) from ..common_utils import IBMWatsonXMixin +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class IBMWatsonXAudioTranscriptionConfig(IBMWatsonXMixin, OpenAIWhisperAudioTranscriptionConfig): """ @@ -174,8 +178,9 @@ class IBMWatsonXAudioTranscriptionConfig(IBMWatsonXMixin, OpenAIWhisperAudioTran # Extract only valid fields for TranscriptionResponse.__init__() # TranscriptionResponse only accepts 'text' and 'usage' in __init__() - text: Final = raw_response_json.get("text") - usage: Final = raw_response_json.get("usage") + response_object: Final = _JSON_OBJECT.validate_python(raw_response_json) + text: Final = response_object.get("text") + usage: Final = response_object.get("usage") # Create response with only valid fields response_kwargs: Final = {} @@ -187,14 +192,14 @@ class IBMWatsonXAudioTranscriptionConfig(IBMWatsonXMixin, OpenAIWhisperAudioTran if not response_kwargs: raise ValueError( "Invalid response format. Received response does not match the expected format. Got: ", - raw_response_json, + response_object, ) response: Final = TranscriptionResponse(**response_kwargs) # Add other fields using dictionary-style assignment (like duration, task, etc.) # Skip fields that TranscriptionResponse doesn't accept in __init__() - for key, value in raw_response_json.items(): + for key, value in response_object.items(): if key not in [ "text", "usage", diff --git a/litellm/llms/watsonx/embed/transformation.py b/litellm/llms/watsonx/embed/transformation.py index 80cd28a058d..609eeba90c0 100644 --- a/litellm/llms/watsonx/embed/transformation.py +++ b/litellm/llms/watsonx/embed/transformation.py @@ -2,9 +2,11 @@ Translates from OpenAI's `/v1/embeddings` to IBM's `/text/embeddings` route. """ +from collections.abc import Iterable, Mapping from typing import Final import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.embedding.transformation import ( BaseEmbeddingConfig, @@ -16,6 +18,9 @@ from litellm.types.utils import EmbeddingResponse, Usage from ..common_utils import IBMWatsonXMixin, _get_api_params +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_TOKEN_COUNT: Final = TypeAdapter(int) + class IBMWatsonXEmbeddingConfig(IBMWatsonXMixin, BaseEmbeddingConfig): def get_supported_openai_params(self, model: str) -> list: @@ -95,7 +100,7 @@ class IBMWatsonXEmbeddingConfig(IBMWatsonXMixin, BaseEmbeddingConfig): json_resp: Final = raw_response.json() if model_response is None: model_response = EmbeddingResponse(model=json_resp.get("model_id", None)) - results: Final = json_resp.get("results", []) + results: Final = _JSON_OBJECTS.validate_python(json_resp.get("results", [])) embedding_response: Final = [] for idx, result in enumerate(results): embedding_response.append( @@ -107,7 +112,7 @@ class IBMWatsonXEmbeddingConfig(IBMWatsonXMixin, BaseEmbeddingConfig): ) model_response.object = "list" model_response.data = embedding_response - input_tokens: Final = json_resp.get("input_token_count", 0) + input_tokens: Final = _TOKEN_COUNT.validate_python(json_resp.get("input_token_count", 0) or 0) setattr( model_response, "usage", diff --git a/litellm/llms/watsonx/rerank/transformation.py b/litellm/llms/watsonx/rerank/transformation.py index ff95f14951a..c2cfd6597f3 100644 --- a/litellm/llms/watsonx/rerank/transformation.py +++ b/litellm/llms/watsonx/rerank/transformation.py @@ -5,10 +5,11 @@ Docs - https://cloud.ibm.com/apidocs/watsonx-ai#text-rerank """ import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Iterable, Mapping, Sequence from typing import Final, cast import httpx +from pydantic import ConfigDict, TypeAdapter from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig @@ -24,6 +25,11 @@ from litellm.types.rerank import ( from ..common_utils import IBMWatsonXMixin, _generate_watsonx_token, _get_api_params +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) +_JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) +_OPTIONAL_INT: Final[TypeAdapter[int | None]] = TypeAdapter(int | None) +_STR: Final = TypeAdapter(str) + class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): """ @@ -171,13 +177,14 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): headers=raw_response.headers, ) - _results: Final[list[dict] | None] = raw_response_json.get("results") + payload: Final = _JSON_OBJECT.validate_python(raw_response_json) + _results: Final = payload.get("results") if _results is None: raise ValueError(f"No results found in the response={raw_response_json}") transformed_results: Final = [] - for result in _results: + for result in _JSON_OBJECTS.validate_python(_results): transformed_result: dict[str, object] = { "index": result["index"], "relevance_score": result["score"], @@ -191,11 +198,11 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): transformed_results.append(transformed_result) - response_id: Final = raw_response_json.get("id") or str(uuid.uuid4()) + response_id: Final = _STR.validate_python(payload.get("id") or str(uuid.uuid4())) # Extract usage information _tokens: Final = RerankTokens( - input_tokens=raw_response_json.get("input_token_count", 0), + input_tokens=_OPTIONAL_INT.validate_python(payload.get("input_token_count", 0)), ) rerank_meta: Final = RerankResponseMeta(tokens=_tokens) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 2d0985bcb64..28ad23fd372 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -4663,6 +4663,7 @@ "azure/eu/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2.2e-05, "cache_read_input_token_cost": 2.75e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.00011, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", @@ -4672,6 +4673,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.00022, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -6306,6 +6308,7 @@ "azure/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2e-05, "cache_read_input_token_cost": 2.5e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.0001, "input_cost_per_token": 5e-06, "litellm_provider": "azure", @@ -6315,6 +6318,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.0002, "output_cost_per_token": 2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -6350,7 +6354,7 @@ "supports_tool_choice": true }, "azure/gpt-4o-transcribe": { - "deprecation_date": "2026-12-31", + "deprecation_date": "2026-10-15", "input_cost_per_audio_token": 2.5e-06, "input_cost_per_token": 2.5e-06, "litellm_provider": "azure", @@ -6358,7 +6362,7 @@ "max_output_tokens": 2000, "mode": "audio_transcription", "output_cost_per_token": 1e-05, - "source": "https://management.azure.com/subscriptions/c873328e-b572-4770-8dff-aaeb6f1f0e79/providers/Microsoft.CognitiveServices/locations/eastus2/models?api-version=2024-10-01", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule", "supported_endpoints": [ "/v1/audio/transcriptions" ] @@ -11093,6 +11097,7 @@ "azure/us/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2.2e-05, "cache_read_input_token_cost": 2.75e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.00011, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", @@ -11102,6 +11107,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.00022, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -12556,6 +12562,7 @@ "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models" }, "azure_ai/jamba-instruct": { + "deprecation_date": "2025-03-01", "input_cost_per_token": 5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 70000, @@ -12563,6 +12570,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 7e-07, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true }, "azure_ai/kimi-k2.5": { diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 21a6113bc4e..94fc109b33c 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -8,7 +8,7 @@ from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypedDict, cast from fastapi import HTTPException -from pydantic import TypeAdapter +from pydantic import ConfigDict, TypeAdapter from typing_extensions import ReadOnly from litellm._logging import verbose_proxy_logger @@ -20,7 +20,9 @@ from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( credential_binding_matches, 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, +) from litellm.proxy._types import ( LiteLLM_MCPServerTable, MCPApprovalStatus, @@ -105,7 +107,31 @@ _AUTH_FLOW_SCOPED_FIELDS: Final["frozenset[str]"] = frozenset( ) -def _blank_to_none(value: str | None) -> str | None: +_OAUTH_CLIENT_CREDENTIAL_FIELDS: Final = frozenset( + { + "client_id", + "client_secret", + "token_endpoint_auth_method", + "redirect_uris", + "dcr_issuer", + "dcr_server_url", + "access_token", + "refresh_token", + "expires_in", + "scope", + } +) + + +def stale_mcp_auth_fields(submitted: Mapping[str, object], previous_value: Callable[[str], object]) -> dict[str, None]: + return { + field: None + for field in _AUTH_FLOW_SCOPED_FIELDS + if field not in submitted or submitted[field] == previous_value(field) + } + + +def _blank_to_none(value: object) -> str | None: if not isinstance(value, str): return None return value.strip() or None @@ -136,6 +162,58 @@ _CLIENT_FORWARDED_AUTH_TYPES: Final["frozenset[str]"] = frozenset({"true_passthr # Minted token material that must never survive a client rotation on a persisted row. _MINTED_TOKEN_CREDENTIAL_FIELDS: Final["frozenset[str]"] = frozenset({"access_token", "refresh_token", "expires_in"}) +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + + +def _bind_submitted_oauth_client( + credentials: Mapping[str, object], issuer: str | None, url: str | None +) -> dict[str, object]: + if not credentials.get("client_id") or credentials.get("dcr_issuer") or credentials.get("dcr_server_url"): + return dict(credentials) + return {**credentials, "dcr_issuer": issuer, "dcr_server_url": url} + + +def is_resubmitted_oauth_client(supplied: dict[str, object], existing: dict[str, object]) -> bool: + """Recognize the saved client without conflating same-ID clients from different issuers.""" + return bool( + supplied.get("client_id") + and _decrypted_credential_field(supplied, "client_id") == _decrypted_credential_field(existing, "client_id") + and all( + _decrypted_credential_field(supplied, field) == _decrypted_credential_field(existing, field) + for field in ("client_secret", "token_endpoint_auth_method", "dcr_issuer", "dcr_server_url") + if field in supplied + ) + ) + + +def oauth_credentials_for_upstream_edit( + credentials: Mapping[str, object], + previous_issuer: str | None, + previous_url: str | None, + *, + issuer_changed: bool, +) -> dict[str, object]: + """Bind an existing client on an explicit edit, so discovery can verify reuse at the new URL. + + No first-use backfill: unchanged legacy rows never enter this path. Without a known previous + issuer, or after a known issuer change, the old client cannot be carried to the new resource. + """ + registered_issuer: Final = _blank_to_none(credentials.get("dcr_issuer")) or previous_issuer + keep_client: Final = bool(registered_issuer and credentials.get("client_id") and not issuer_changed) + removed: Final = ( + _MINTED_TOKEN_CREDENTIAL_FIELDS | {"auth_value"} + if keep_client + else _OAUTH_CLIENT_CREDENTIAL_FIELDS | {"auth_value"} + ) + retained: Final = {key: value for key, value in credentials.items() if key not in removed} + if keep_client: + return { + **retained, + "dcr_issuer": registered_issuer, + "dcr_server_url": credentials.get("dcr_server_url") or previous_url, + } + return retained + class _OAuthCredentialAccessToken(TypedDict): access_token: str @@ -386,7 +464,11 @@ def _prepare_mcp_server_data( if blob_value is not None and te_field not in data_dict: data_dict[te_field] = blob_value data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=_get_salt_key()) - data_dict["credentials"] = safe_dumps(data_dict["credentials"]) + data_dict["credentials"] = safe_dumps( + _bind_submitted_oauth_client(data_dict["credentials"], data.issuer, data.url) + if not exclude_unset and data.auth_type == "oauth2" + else data_dict["credentials"] + ) # Serialize JSON fields from ``data_dict`` (not ``data``) so the # exclude_unset filter is respected. Reading back from ``data`` would @@ -1132,6 +1214,7 @@ async def _update_mcp_server_row( *, server_id: str, data_dict: Mapping[str, object], + expected_updated_at: datetime | None = None, ) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None": identifier_write: Final = any(field in data_dict for field in ("server_name", "alias")) protocol_write: Final = bool({"transport", "mcp_info"}.intersection(data_dict)) @@ -1144,6 +1227,11 @@ async def _update_mcp_server_row( if stored is None: return None _validate_mcp_protocol_write(stored, data_dict) + if expected_updated_at is not None: + changed: Final = await table.update_many( + where={"server_id": server_id, "updated_at": expected_updated_at}, data=data_dict + ) + return await table.find_unique(where={"server_id": server_id}) if changed else None return await table.update( where={"server_id": server_id}, data=data_dict, @@ -1180,6 +1268,7 @@ async def update_mcp_server( data: UpdateMCPServerRequest, touched_by: str, fields_set: set[str] | None = None, + expected_updated_at: datetime | None = None, ) -> LiteLLM_MCPServerTable | McpIdentifierConflict | None: """ Update a new mcp server record in the db @@ -1193,16 +1282,16 @@ async def update_mcp_server( # exclude_unset=True makes this a true partial update: fields the caller did # not provide are not written, so they keep their existing DB value instead # of being reset to a schema default (transport=sse, allow_all_keys=False...). - data_dict: Final = _prepare_mcp_server_data(data, exclude_unset=True, fields_set=fields_set) + prepared_data: Final = _prepare_mcp_server_data(data, exclude_unset=True, fields_set=fields_set) # Pre-fetch existing record once if we need it for auth_type, url, or credential logic existing = None - has_credentials: Final = "credentials" in data_dict and data_dict["credentials"] is not None + has_credentials: Final = "credentials" in prepared_data and prepared_data["credentials"] is not None # An explicit token-exchange column write (set or clear) also migrates the # legacy blob copies below, so the existing row is needed for those updates. - explicit_te_write: Final = bool(_TOKEN_EXCHANGE_COLUMN_FIELDS & data_dict.keys()) - url_provided: Final = "url" in data_dict and data_dict["url"] is not None - issuer_provided: Final = "issuer" in data_dict + explicit_te_write: Final = bool(_TOKEN_EXCHANGE_COLUMN_FIELDS & prepared_data.keys()) + url_provided: Final = "url" in prepared_data and prepared_data["url"] is not None + issuer_provided: Final = "issuer" in prepared_data if data.auth_type or has_credentials or explicit_te_write or url_provided or issuer_provided: existing = await _db_find_mcp_server_row(prisma_client, data.server_id) @@ -1213,29 +1302,60 @@ async def update_mcp_server( ) # A url change re-points the server at a potentially different upstream, so any discovered or # trust-on-first-use OAuth endpoints/issuer belong to the old upstream and must re-discover. - url_changed: Final = bool(url_provided and existing and existing.url != data_dict["url"]) + url_changed: Final = bool(url_provided and existing and existing.url != prepared_data["url"]) old_issuer: Final = _blank_to_none(getattr(existing, "issuer", None)) if existing else None issuer_changed: Final = bool( - issuer_provided and old_issuer is not None and _blank_to_none(data_dict.get("issuer")) != old_issuer + issuer_provided and existing is not None and _blank_to_none(prepared_data.get("issuer")) != old_issuer ) + oauth_upstream_edited: Final = bool(existing and existing.auth_type == "oauth2" and (url_changed or issuer_changed)) + existing_credentials: Final = _credentials_blob_to_mutable_dict((existing.credentials or {}) if existing else {}) + retained_credentials: Final = ( + oauth_credentials_for_upstream_edit( + existing_credentials, + old_issuer, + existing.url if existing else None, + issuer_changed=issuer_changed or auth_type_changed, + ) + if oauth_upstream_edited + else existing_credentials + ) + edited_credentials: Final = {"credentials": safe_dumps(retained_credentials)} if oauth_upstream_edited else {} + cleared_auth_fields: Final = ( + stale_mcp_auth_fields(prepared_data, lambda field: getattr(existing, field, None)) + if auth_type_changed or url_changed or issuer_changed + else {} + ) + supplied: Final = _credentials_blob_to_mutable_dict(prepared_data.get("credentials") or {}) + safe_submitted: Final = ( + oauth_credentials_for_upstream_edit( + supplied, + old_issuer, + existing.url if existing else None, + issuer_changed=issuer_changed or auth_type_changed, + ) + if oauth_upstream_edited and is_resubmitted_oauth_client(supplied, existing_credentials) + else supplied + ) + submitted_credentials: Final = ( + { + "credentials": safe_dumps( + _bind_submitted_oauth_client( + safe_submitted, + _blank_to_none(cleared_auth_fields.get("issuer", prepared_data.get("issuer", old_issuer))), + _blank_to_none(prepared_data.get("url", existing.url if existing else None)), + ) + ) + } + if has_credentials and (data.auth_type or (existing.auth_type if existing else None)) == "oauth2" + else {} + ) + data_dict: Final = {**edited_credentials, **prepared_data, **submitted_credentials, **cleared_auth_fields} + # Clear stale credentials when auth_type changes but no new credentials provided if auth_type_changed and "credentials" not in data_dict: data_dict["credentials"] = None - if auth_type_changed or url_changed or issuer_changed: - # Clear each auth-flow-scoped field that the caller either omitted (partial update) or - # resubmitted unchanged. The edit form re-sends every field, so a stale issuer/endpoint - # belonging to the old upstream would otherwise survive a url/auth_type change and win in the - # resolution merge; only a genuinely new submitted value is kept. - data_dict.update( - { - field: None - for field in _AUTH_FLOW_SCOPED_FIELDS - if field not in data_dict or data_dict[field] == getattr(existing, field, None) - } - ) - # An explicit column write that does not touch credentials must still migrate # the row's legacy blob copies: lift values for columns the caller left # untouched, strip every copy from the blob. Without this, clearing a column @@ -1263,11 +1383,10 @@ async def update_mcp_server( # within the client-forwarded class (true_passthrough โ†” oauth_delegate) keeps # the same declared app and so must merge, not replace. if not auth_type_changed: - existing_creds = _credentials_blob_to_mutable_dict(existing.credentials) new_creds: Final = _credentials_blob_to_mutable_dict(data_dict["credentials"]) # New values override existing; existing keys not in update are preserved. A client # rotation additionally drops the previous app's stale minted token keys. - merged: Final = _drop_stale_minted_on_client_rotation({**existing_creds, **new_creds}, new_creds) + merged: Final = _drop_stale_minted_on_client_rotation({**retained_credentials, **new_creds}, new_creds) # Migrate-on-write for legacy rows: token-exchange settings the # old blob shape carried move to their dedicated columns (unless # the caller set the column this update, or the row already has @@ -1298,6 +1417,7 @@ async def update_mcp_server( prisma_client, server_id=data.server_id, data_dict=data_dict, + expected_updated_at=expected_updated_at, ) if isinstance(updated_mcp_server, McpIdentifierConflict): @@ -1798,7 +1918,7 @@ def mcp_oauth_token_identity(server: object) -> tuple[object, ...]: creds: Final = getattr(server, "credentials", None) if isinstance(creds, str): try: - parsed: dict[str, object] | None = json.loads(creds) + parsed: Mapping[str, object] | None = _JSON_OBJECT.validate_python(json.loads(creds)) except ValueError: parsed = None else: @@ -2264,13 +2384,10 @@ def _decode_user_env_vars(stored: str) -> dict[str, str]: "re-enter them rather than silently forwarding ciphertext" ) return {} - parsed: dict[str, object] | None try: - parsed = json.loads(decrypted) + parsed: Final = _JSON_OBJECT.validate_python(json.loads(decrypted)) except (ValueError, TypeError): return {} - if not isinstance(parsed, dict): - return {} return {str(k): str(v) for k, v in parsed.items()} diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 7a0f59c3c2b..342c650d2a6 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -12,7 +12,7 @@ from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse import httpx from fastapi import APIRouter, Depends, Form, HTTPException, Request from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response -from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError +from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError from litellm._logging import verbose_logger from litellm.caching.in_memory_cache import InMemoryCache @@ -75,6 +75,7 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import ( TOKEN_NO_CACHE_HEADERS, build_upstream_oauth2_token_request, get_request_base_url, + oauth_client_registration_matches, resolve_upstream_resource, validate_trusted_redirect_uri, well_known_root_suffix, @@ -200,6 +201,8 @@ def encode_state_with_base_url( dcr_client_id: str | None = None, dcr_client_secret: str | None = None, dcr_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None, + expected_issuer: str | None = None, + authorization_response_iss_parameter_supported: bool = False, oauth_nonce: str | None = None, ) -> str: """ @@ -225,6 +228,8 @@ def encode_state_with_base_url( response granted the minted client, sealed alongside the credentials so the exchange authenticates the way the upstream expects instead of falling back to the server row's configured method + expected_issuer: Issuer identifier of the authorization server this flow is being sent to, + sealed so /callback can hold the RFC 9207 ``iss`` of the response against it Returns: An encrypted string that encodes all values @@ -241,6 +246,8 @@ def encode_state_with_base_url( "dcr_client_id": dcr_client_id, "dcr_client_secret": dcr_client_secret, "dcr_token_endpoint_auth_method": dcr_token_endpoint_auth_method, + "expected_issuer": expected_issuer, + "authorization_response_iss_parameter_supported": authorization_response_iss_parameter_supported, } state_json: Final = json.dumps(state_data, sort_keys=True) encrypted_state: Final = encrypt_value_helper(state_json) @@ -807,7 +814,16 @@ def _dcr_bridge_relays_client_registration(mcp_server: MCPServer) -> bool: returns directly to the client's redirect URI without transiting the gateway. Gateway-side redirect trust and the ``/callback`` state relay therefore only apply to the short-circuit arm, where the upstream only knows the gateway's own callback.""" - return mcp_server.is_dcr_bridge and bool(mcp_server.effective_registration_url) and not mcp_server.client_id + return bool( + mcp_server.is_dcr_bridge + and mcp_server.effective_registration_url + and not ( + mcp_server.client_id + and oauth_client_registration_matches( + mcp_server.dcr_issuer, mcp_server.dcr_server_url, mcp_server.issuer, mcp_server.url + ) + ) + ) def _require_s256_pkce( @@ -932,6 +948,10 @@ async def authorize_with_server( ): _raise_if_not_oauth2(mcp_server) resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _register_flow_needed_endpoint) + if not oauth_client_registration_matches( + resolved_server.dcr_issuer, resolved_server.dcr_server_url, resolved_server.issuer, resolved_server.url + ): + raise HTTPException(status_code=400, detail="OAuth client belongs to a different issuer; register a new client") if resolved_server.effective_authorization_url is None: raise HTTPException( status_code=400, @@ -1003,6 +1023,8 @@ async def authorize_with_server( dcr_token_endpoint_auth_method=ephemeral_dcr_client.token_endpoint_auth_method if ephemeral_dcr_client else None, + expected_issuer=resolved_server.issuer, + authorization_response_iss_parameter_supported=resolved_server.authorization_response_iss_parameter_supported, ) relay_state: Final = secrets.token_urlsafe(_OAUTH_STATE_HANDLE_BYTES) @@ -1065,6 +1087,10 @@ async def exchange_token_with_server( raise HTTPException(status_code=400, detail="Unsupported grant_type") resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _token_flow_needed_endpoint) + if not oauth_client_registration_matches( + resolved_server.dcr_issuer, resolved_server.dcr_server_url, resolved_server.issuer, resolved_server.url + ): + raise HTTPException(status_code=400, detail="OAuth client belongs to a different issuer; register a new client") token_url: Final = resolved_server.effective_token_url if token_url is None: raise HTTPException( @@ -1364,6 +1390,8 @@ class _DcrClientRegistration(BaseModel): class _PersistedDcrCredentials(BaseModel): + dcr_issuer: str | None = None + dcr_server_url: str | None = None client_id: str | None = None client_secret: str | None = None token_endpoint_auth_method: str | None = None @@ -1410,19 +1438,28 @@ def _decrypt_persisted_dcr_credential(value: str | None, key: str) -> str | None def _apply_persisted_dcr_credentials(mcp_server: MCPServer, credentials: _PersistedDcrCredentials) -> bool: + if not oauth_client_registration_matches( + credentials.dcr_issuer, credentials.dcr_server_url, mcp_server.issuer, mcp_server.url + ): + return False client_id: Final = _decrypt_persisted_dcr_credential(credentials.client_id, "client_id") if not client_id: return False + mcp_server.dcr_issuer = credentials.dcr_issuer # rebind-ok: publish through the existing bool hydration contract + mcp_server.dcr_server_url = credentials.dcr_server_url # rebind-ok: retain binding in database-free shared state mcp_server.client_id = client_id mcp_server.client_secret = _decrypt_persisted_dcr_credential(credentials.client_secret, "client_secret") mcp_server.token_endpoint_auth_method = credentials.token_endpoint_auth_method return True -async def _load_store_dcr_credentials(mcp_server: MCPServer) -> _PersistedDcrCredentials | None: +async def _load_store_dcr_credentials( + mcp_server: MCPServer, *, raise_on_error: bool = False +) -> _PersistedDcrCredentials | None: """DCR client persisted in the server-scoped OAuth-client store for a config-declared server (which has no LiteLLM_MCPServerTable row). Returns None when the store has no usable client_id - or the DB is unreachable.""" + or the DB is unreachable. Registration writes request strict reads so a failed lookup cannot + be mistaken for an absent client.""" from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # avoids circular import get_mcp_server_oauth_client_credentials, ) @@ -1434,6 +1471,8 @@ async def _load_store_dcr_credentials(mcp_server: MCPServer) -> _PersistedDcrCre prisma_client=prisma_client, server_id=mcp_server.server_id ) except Exception as exc: # noqa: BLE001 # best-effort read; DB may be unreachable + if raise_on_error: + raise verbose_logger.debug( "register_client_with_server: failed to read stored DCR client for server_id=%s: %s", mcp_server.server_id, @@ -1463,6 +1502,8 @@ async def hydrate_config_server_dcr_client(mcp_server: MCPServer) -> bool: async def _resolve_persisted_dcr_client( mcp_server: MCPServer, + *, + raise_on_error: bool = False, ) -> tuple[Optional["LiteLLM_MCPServerTable"], _PersistedDcrCredentials | None]: """Resolve a server's persisted DCR client using the same two-level rule the write path uses, so read and write always agree. First, whether the server HAS a LiteLLM_MCPServerTable row: a row is @@ -1483,6 +1524,8 @@ async def _resolve_persisted_dcr_client( prisma_client = get_prisma_client_or_throw("Database not connected. Cannot read MCP OAuth client registration.") row: Final = await get_mcp_server(prisma_client=prisma_client, server_id=mcp_server.server_id) except Exception as exc: # noqa: BLE001 # best-effort read; DB may be unreachable + if raise_on_error: + raise verbose_logger.debug( "register_client_with_server: failed to read persisted DCR client for server_id=%s: %s", mcp_server.server_id, @@ -1491,12 +1534,14 @@ async def _resolve_persisted_dcr_client( return None, None if row is not None: + if row.url != mcp_server.url or (row.issuer and mcp_server.issuer and row.issuer != mcp_server.issuer): + return row, None credentials: Final = _get_persisted_dcr_credentials(row.credentials) if credentials is not None and credentials.client_id: return row, credentials return row, None if global_mcp_server_manager.is_config_declared_server(mcp_server.server_id): - return None, await _load_store_dcr_credentials(mcp_server) + return None, await _load_store_dcr_credentials(mcp_server, raise_on_error=raise_on_error) return None, None @@ -1519,20 +1564,24 @@ async def _reuse_persisted_dcr_client_if_available( if not _apply_persisted_dcr_credentials(mcp_server, credentials): return False - if persisted_mcp_server is not None: + await _refresh_persisted_dcr_server(persisted_mcp_server) + return bool(mcp_server.client_id) + + +async def _refresh_persisted_dcr_server(persisted_server: Optional["LiteLLM_MCPServerTable"]) -> None: + if persisted_server is not None and persisted_server.approval_status != "draft": from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids circular import global_mcp_server_manager, ) try: - await global_mcp_server_manager.update_server(persisted_mcp_server) + await global_mcp_server_manager.update_server(persisted_server) except Exception as exc: # noqa: BLE001 # best-effort registry refresh verbose_logger.warning( "register_client_with_server: failed to refresh persisted DCR client registration for server_id=%s: %s", - mcp_server.server_id, + persisted_server.server_id, exc, ) - return bool(mcp_server.client_id) async def _persisted_dcr_redirect_uri_is_stale(mcp_server: MCPServer, current_redirect_uri: str) -> bool: @@ -1575,7 +1624,7 @@ async def _persist_dcr_client_registration( ``refresh_token`` grant has no client identity, so an expired access token forces a full re-authorization instead of a silent refresh. Mirrors the ``encrypt_credentials`` write that ``client_credentials`` and token exchange already use. Failures are logged, - never raised: registration still returns to the caller even when persistence fails. + never raised here: the registration endpoint rejects a failed persistence result. The client-forwarded token modes (``true_passthrough`` / ``oauth_delegate``) are skipped unconditionally: the caller holds the upstream token and the gateway must hold no OAuth @@ -1604,9 +1653,6 @@ async def _persist_dcr_client_registration( ) return "failed" - if await _reuse_persisted_dcr_client_if_available(mcp_server, current_redirect_uri=current_redirect_uri): - return "reused" - token_endpoint_auth_method: Final = ( "client_secret_basic" if registration.token_endpoint_auth_method == "client_secret_basic" else None ) @@ -1615,6 +1661,8 @@ async def _persist_dcr_client_registration( "client_secret": registration.client_secret, "token_endpoint_auth_method": token_endpoint_auth_method, "redirect_uris": [current_redirect_uri], + "dcr_issuer": mcp_server.issuer, + "dcr_server_url": mcp_server.url, } from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # avoids circular import @@ -1632,6 +1680,27 @@ async def _persist_dcr_client_registration( prisma_client: Final = get_prisma_client_or_throw( "Database not connected. Cannot persist MCP OAuth client registration." ) + except HTTPException: + # This getter only raises when no database is configured. The existing single-process + # temporary-session mode keeps registrations in memory; database write errors below fail. + _apply_persisted_dcr_credentials(mcp_server, _PersistedDcrCredentials.model_validate(credentials)) + return "persisted" + + try: + stored, latest_credentials = await _resolve_persisted_dcr_client(mcp_server, raise_on_error=True) + if stored is not None and ( + stored.url != mcp_server.url + or stored.auth_type != mcp_server.auth_type + or (stored.issuer and mcp_server.issuer and stored.issuer != mcp_server.issuer) + ): + return "failed" + if ( + latest_credentials is not None + and not _redirect_uri_not_registered(latest_credentials, current_redirect_uri) + and _apply_persisted_dcr_credentials(mcp_server, latest_credentials) + ): + await _refresh_persisted_dcr_server(stored) + return "reused" updated_row: Final = await update_mcp_server( prisma_client=prisma_client, data=( @@ -1649,19 +1718,25 @@ async def _persist_dcr_client_registration( ) ), touched_by="mcp_oauth_dcr", + expected_updated_at=stored.updated_at if stored else None, ) if updated_row is not None and not isinstance(updated_row, McpIdentifierConflict): - await global_mcp_server_manager.update_server(updated_row) + await _refresh_persisted_dcr_server(updated_row) + _apply_persisted_dcr_credentials(mcp_server, _PersistedDcrCredentials.model_validate(credentials)) return "persisted" + if stored is not None: + return ( + "reused" + if await _reuse_persisted_dcr_client_if_available(mcp_server, current_redirect_uri) + else "failed" + ) if global_mcp_server_manager.is_config_declared_server(mcp_server.server_id): await upsert_mcp_server_oauth_client_credentials( prisma_client=prisma_client, server_id=mcp_server.server_id, credentials=credentials, ) - mcp_server.client_id = registration.client_id - mcp_server.client_secret = registration.client_secret - mcp_server.token_endpoint_auth_method = token_endpoint_auth_method + _apply_persisted_dcr_credentials(mcp_server, _PersistedDcrCredentials.model_validate(credentials)) return "persisted" except Exception as exc: # noqa: BLE001 verbose_logger.warning( @@ -1859,20 +1934,27 @@ async def register_client_with_server( "redirect_uris": client_facing_redirect_uris, } - if mcp_server.client_id and not ( - persist_credentials - and mcp_server.registration_url - and await _persisted_dcr_redirect_uri_is_stale(mcp_server, current_redirect_uri) + resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _register_flow_needed_endpoint) + client_matches: Final = oauth_client_registration_matches( + resolved_server.dcr_issuer, resolved_server.dcr_server_url, resolved_server.issuer, resolved_server.url + ) + if ( + resolved_server.client_id + and client_matches + and not ( + persist_credentials + and resolved_server.registration_url + and await _persisted_dcr_redirect_uri_is_stale(resolved_server, current_redirect_uri) + ) ): return dummy_return if await _reuse_persisted_dcr_client_if_available( - mcp_server, + resolved_server, current_redirect_uri=current_redirect_uri if persist_credentials else None, ): return dummy_return - resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _register_flow_needed_endpoint) if resolved_server.effective_authorization_url is None: raise HTTPException( status_code=400, @@ -1908,19 +1990,32 @@ async def register_client_with_server( server_id=resolved_server.server_id, ) - token_response = response.json() - - if persist_credentials and not bridge_relay: - persistence_result = await _persist_dcr_client_registration( - resolved_server, token_response, current_redirect_uri - ) - if persistence_result == "reused": - return dummy_return - - if client_redirect_uris and not bridge_relay and isinstance(token_response, dict): - token_response = {**token_response, "redirect_uris": client_facing_redirect_uris} - - return JSONResponse(token_response) + token_response: Final = response.json() + persistence_result: Final = ( + await _persist_dcr_client_registration(resolved_server, token_response, current_redirect_uri) + if persist_credentials and not bridge_relay + else None + ) + if persistence_result == "reused": + return dummy_return + if persistence_result == "failed": + raise HTTPException(status_code=503, detail="OAuth client registration could not be saved; retry authorization") + bound_response: Final = ( + { + **token_response, + "dcr_issuer": resolved_server.issuer, + "dcr_server_url": resolved_server.url, + "dcr_redirect_uris": [current_redirect_uri], + } + if persistence_result == "persisted" + else token_response + ) + client_response: Final = ( + {**bound_response, "redirect_uris": client_facing_redirect_uris} + if client_redirect_uris and not bridge_relay and isinstance(bound_response, dict) + else bound_response + ) + return JSONResponse(client_response) @router.get("/authorize/mcp-session") @@ -2229,11 +2324,19 @@ def _render_oauth_error_html(error: str, description: str | None) -> HTMLRespons return HTMLResponse(body, status_code=400) +def _authorization_response_issuer_is_trusted(response_issuer: str | None, state_data: Mapping[str, object]) -> bool: + expected_issuer: Final = state_data.get("expected_issuer") + if response_issuer is None: + return state_data.get("authorization_response_iss_parameter_supported") is not True + return isinstance(expected_issuer, str) and bool(expected_issuer) and response_issuer == expected_issuer + + @router.get("/callback") async def callback( request: Request, code: str | None = None, state: str | None = None, + iss: str | None = None, error: str | None = None, error_description: str | None = None, error_uri: str | None = None, @@ -2244,7 +2347,9 @@ async def callback( - A successful authorization response (``code`` + ``state``), which is forwarded back to the validated client ``redirect_uri`` with the - original (un-wrapped) ``state``. + original (un-wrapped) ``state``, once the RFC 9207 ``iss`` (when the + authorization server sent one) matches the issuer /authorize sealed + into the state. - An error response (``error``[+``error_description``/``error_uri``]), per RFC 6749 ยง4.1.2.1. When ``state`` is present and decodes to a trusted ``redirect_uri``, the error params are propagated back to the client so @@ -2262,6 +2367,13 @@ async def callback( encoded_state = _resolve_encoded_oauth_state(request, state) try: state_data = decode_state_hash(encoded_state) + error_issuer_state: Final = TypeAdapter(dict[str, object]).validate_python(state_data) + if not _authorization_response_issuer_is_trusted(iss, error_issuer_state): + rejected_error: Final = _render_oauth_error_html( + "invalid_issuer", "Unexpected authorization issuer" + ) + _clear_oauth_state_cookie(rejected_error, request, state) + return rejected_error original_state = state_data.get("original_state") redirect_uri = _get_validated_client_redirect_uri(request, state_data) except Exception: @@ -2311,6 +2423,22 @@ async def callback( # states while permitting same-origin / allowlisted clients. redirect_uri = _get_validated_client_redirect_uri(request, state_data) + issuer_state: Final = TypeAdapter(dict[str, object]).validate_python(state_data) + if not _authorization_response_issuer_is_trusted(iss, issuer_state): + verbose_logger.warning( + "MCP /callback rejected an authorization response: RFC 9207 iss=%r does not match the " + "issuer this flow was sent to (%r)", + iss, + issuer_state.get("expected_issuer"), + ) + issuer_error_response: Final = _render_oauth_error_html( + "invalid_issuer", + "This authorization response came from a different identity provider than the one this " + "MCP server is configured to use.", + ) + _clear_oauth_state_cookie(issuer_error_response, request, state) + return issuer_error_response + # Interactive dcr_bridge oauth_delegate: the state carries the litellm user the authorize step # captured. Instead of forwarding the raw upstream code (which the client would present at the # token endpoint with no way to prove who signed in), seal the user and the upstream code into a diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 49cd8cf6fb9..abfb9dc56fe 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -719,13 +719,12 @@ def _normalized_authorize_endpoint(url: str) -> str: def _issuer_matches(claimed_issuer: object, configured_issuer: str) -> bool: """RFC 8414 ยง3.3 issuer equality between the metadata document's self-attested ``issuer`` and the - admin-configured issuer, tolerant only of URL-insignificant differences (scheme/host case, the - default port, a trailing slash). A non-string or empty claimed issuer never matches, so a + admin-configured issuer. A non-string or empty claimed issuer never matches, so a document that omits ``issuer`` fails closed under issuer-anchored discovery. """ if not isinstance(claimed_issuer, str) or not claimed_issuer: return False - return _normalized_authorize_endpoint(claimed_issuer) == _normalized_authorize_endpoint(configured_issuer) + return claimed_issuer == configured_issuer def _flow_endpoints_missing( @@ -856,6 +855,11 @@ def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serv may_carry: Final = _endpoints_corroborate_authorization_url( previous_server.authorization_url, new_server.authorization_url ) + 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 + previous_server.authorization_response_iss_parameter_supported + ) if new_server.authorization_url is None and previous_server.authorization_url: new_server.authorization_url = previous_server.authorization_url if may_carry and new_server.token_url is None and previous_server.token_url: @@ -2110,13 +2114,20 @@ class MCPServerManager: if metadata is None: return server discovered_issuer: Final = metadata.discovered_issuer if not metadata.from_origin_fallback else None - resolved: Final = server.model_copy() - resolved.scopes = server.scopes or metadata.scopes - resolved.issuer = server.issuer or discovered_issuer - resolved.authorization_url = server.authorization_url or metadata.authorization_url - resolved.token_url = server.token_url or metadata.token_url - resolved.registration_url = server.registration_url or metadata.registration_url - return resolved + return server.model_copy( + update={ + "scopes": server.scopes or metadata.scopes, + "issuer": server.issuer or discovered_issuer, + "authorization_response_iss_parameter_supported": ( + metadata.authorization_response_iss_parameter_supported + if discovered_issuer is not None + else server.authorization_response_iss_parameter_supported + ), + "authorization_url": server.authorization_url or metadata.authorization_url, + "token_url": server.token_url or metadata.token_url, + "registration_url": server.registration_url or metadata.registration_url, + } + ) def _oauth_discovery_slot_is_current(self, server_id: str, generation: int) -> bool: slot: Final = self._oauth_discovery_slot(server_id) @@ -2641,6 +2652,11 @@ class MCPServerManager: scopes=resolved_scopes, configured_scopes=tuple(configured_scopes) if configured_scopes else None, issuer=effective_issuer, + authorization_response_iss_parameter_supported=( + gated_oauth_metadata.authorization_response_iss_parameter_supported + if gated_oauth_metadata + else False + ), issuer_is_anchored=use_issuer_anchor, authorization_url=resolved_authorization_url, token_url=resolved_token_url, @@ -3202,12 +3218,17 @@ class MCPServerManager: extra_headers=getattr(mcp_server, "extra_headers", None), static_headers=static_headers_dict, env_vars=env_vars_list, + dcr_issuer=credentials_dict.get("dcr_issuer") if credentials_dict else None, + dcr_server_url=credentials_dict.get("dcr_server_url") if credentials_dict else None, client_id=client_id_value or getattr(mcp_server, "client_id", None), client_secret=client_secret_value or getattr(mcp_server, "client_secret", None), oauth2_flow=self._explicit_oauth2_flow(getattr(mcp_server, "oauth2_flow", None)), scopes=resolved_scopes, configured_scopes=configured_scopes, issuer=effective_issuer, + authorization_response_iss_parameter_supported=( + gated_oauth_metadata.authorization_response_iss_parameter_supported if gated_oauth_metadata else False + ), issuer_is_anchored=use_issuer_anchor, authorization_url=manual_authorization_url or getattr(gated_oauth_metadata, "authorization_url", None), token_url=manual_token_url or getattr(gated_oauth_metadata, "token_url", None), @@ -3801,7 +3822,7 @@ class MCPServerManager: if server is None: verbose_logger.warning("MCP Server %s not found", server_id) return [] - return await self._get_tools_from_server(server) + return list(await self._get_tools_from_server(server)) except Exception as e: verbose_logger.warning("Failed to get tools from server %s: %s", server_id, e) return [] @@ -3838,11 +3859,13 @@ class MCPServerManager: server_auth_header: Final = _server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header) try: - tools: Final = await self._get_tools_from_server( - server=server, - mcp_auth_header=server_auth_header, - user_api_key_auth=user_api_key_auth, - record_listing=True, + tools: Final = list( + await self._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + user_api_key_auth=user_api_key_auth, + record_listing=True, + ) ) return tools except Exception as e: @@ -4350,12 +4373,11 @@ class MCPServerManager: # applied (e.g. "test_petstore-getinventory"). Do NOT pass them # through _create_prefixed_tools โ€” that would add the prefix a second # time producing "test_petstore-test_petstore-getinventory". - unprefixed_tools: Final = guarded_openapi - self._record_listed_tools( - server, unprefixed_tools, listed_caller, listed_generation, record_listing=record_listing + self.record_listed_tools( + server, guarded_openapi, listed_caller, listed_generation, record_listing=record_listing ) if not add_prefix: - return unprefixed_tools + return guarded_openapi return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi] else: tools = await self._fetch_tools_with_timeout(client, server.name) @@ -4371,7 +4393,7 @@ class MCPServerManager: prefixed_or_original_tools: Final = self._create_prefixed_tools( guarded_tools, server, add_prefix=add_prefix ) - self._record_listed_tools( + self.record_listed_tools( server, guarded_tools, listed_caller, listed_generation, record_listing=record_listing ) @@ -4478,17 +4500,6 @@ class MCPServerManager: return self._listed_tools_generations.get(server_id, 0) def record_listed_tools( - self, - server: MCPServer, - tools: Sequence[MCPTool], - caller: ListedToolsCaller | None, - generation: int, - *, - record_listing: bool = True, - ) -> None: - self._record_listed_tools(server, tools, caller, generation, record_listing=record_listing) - - def _record_listed_tools( self, server: MCPServer, tools: Sequence[MCPTool], @@ -5184,6 +5195,10 @@ class MCPServerManager: token_url=data.get("token_endpoint"), registration_url=data.get("registration_endpoint"), 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" + ) + is True, ) if any( diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 39865a35ec6..3b5eda6cff2 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -658,6 +658,17 @@ def canonicalize_url_identity(url: str) -> str: return urlunparse((scheme, netloc, parsed.path.rstrip("/"), "", "", "")) +def oauth_client_registration_matches( + registered_issuer: str | None, + registered_url: str | None, + current_issuer: str | None, + current_url: str | None, +) -> bool: + if registered_issuer and current_issuer: + return registered_issuer == current_issuer + return not registered_url or registered_url == current_url + + def canonical_resource_uri(url: str) -> str | None: """Canonicalize an upstream MCP server URL into an RFC 8707 resource identifier. diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 1dea44a84f0..42cca179cd8 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1125,18 +1125,20 @@ async def _get_tools_from_mcp_servers( from litellm.proxy.proxy_server import proxy_logging_obj listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id) - tools: Final = await global_mcp_server_manager._get_tools_from_server( - server=server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - client_ip=client_ip, - user_api_key_auth=user_api_key_auth, - oauth2_headers=oauth2_headers, - proxy_logging_obj=proxy_logging_obj, - catalog_auth_header=catalog_auth_header, - record_listing=False, + tools: Final = list( + await global_mcp_server_manager._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, # Always add server prefix + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, + oauth2_headers=oauth2_headers, + proxy_logging_obj=proxy_logging_obj, + catalog_auth_header=catalog_auth_header, + record_listing=False, + ) ) filtered_tools = filter_tools_by_allowed_tools(tools, server) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 845dcaa1f16..534eda07292 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -705,16 +705,18 @@ if MCP_AVAILABLE: *, record_listing: bool, ) -> list[MCPTool]: - return await global_mcp_server_manager._get_tools_from_server( - server=server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=False, - raw_headers=raw_headers, - client_ip=client_ip, - user_api_key_auth=user_api_key_auth, - proxy_logging_obj=proxy_logging_obj, - record_listing=record_listing, + return list( + await global_mcp_server_manager._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=False, + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, + proxy_logging_obj=proxy_logging_obj, + record_listing=record_listing, + ) ) async def _get_tools_for_single_server( diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index cb6d47ca4c8..4a7ca257d9e 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -139,11 +139,6 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( module_path="litellm.proxy.management_endpoints.model_insights_endpoints", path_prefixes=("/model-insights",), ), - LazyFeature( - name="roi_calculator", - module_path="litellm.proxy.management_endpoints.roi_calculator_endpoints", - path_prefixes=("/roi-calculator",), - ), LazyFeature( name="search_tools", module_path="litellm.proxy.search_endpoints.search_tool_management", diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index a927d4b415f..03408043f8e 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -31307,6 +31307,28 @@ ], "title": "Client Secret" }, + "dcr_issuer": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Issuer" + }, + "dcr_server_url": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Server Url" + }, "id_jag_resource": { "anyOf": [ { @@ -34023,6 +34045,28 @@ ], "title": "Client Secret" }, + "dcr_issuer": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Issuer" + }, + "dcr_server_url": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Server Url" + }, "id_jag_resource": { "anyOf": [ { @@ -35696,7 +35740,7 @@ }, "/callback": { "get": { - "description": "OAuth 2.0 authorization response handler for MCP loopback clients.\n\nAccepts either:\n\n- A successful authorization response (``code`` + ``state``), which is\n forwarded back to the validated client ``redirect_uri`` with the\n original (un-wrapped) ``state``.\n- An error response (``error``[+``error_description``/``error_uri``]), per\n RFC 6749 \u00a74.1.2.1. When ``state`` is present and decodes to a trusted\n ``redirect_uri``, the error params are propagated back to the client so\n its OAuth library can surface them. Otherwise we render an HTML error\n page so the user is not left on an opaque 422 / blank screen.", + "description": "OAuth 2.0 authorization response handler for MCP loopback clients.\n\nAccepts either:\n\n- A successful authorization response (``code`` + ``state``), which is\n forwarded back to the validated client ``redirect_uri`` with the\n original (un-wrapped) ``state``, once the RFC 9207 ``iss`` (when the\n authorization server sent one) matches the issuer /authorize sealed\n into the state.\n- An error response (``error``[+``error_description``/``error_uri``]), per\n RFC 6749 \u00a74.1.2.1. When ``state`` is present and decodes to a trusted\n ``redirect_uri``, the error params are propagated back to the client so\n its OAuth library can surface them. Otherwise we render an HTML error\n page so the user is not left on an opaque 422 / blank screen.", "operationId": "callback_callback_get", "parameters": [ { @@ -35731,6 +35775,22 @@ "title": "State" } }, + { + "in": "query", + "name": "iss", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Iss" + } + }, { "in": "query", "name": "error", @@ -37486,6 +37546,28 @@ ], "title": "Client Secret" }, + "dcr_issuer": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Issuer" + }, + "dcr_server_url": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Server Url" + }, "id_jag_resource": { "anyOf": [ { @@ -41480,6 +41562,28 @@ ], "title": "Client Secret" }, + "dcr_issuer": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Issuer" + }, + "dcr_server_url": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Server Url" + }, "id_jag_resource": { "anyOf": [ { @@ -49003,2976 +49107,6 @@ } } }, - "roi_calculator": { - "components": { - "schemas": { - "HTTPValidationError": { - "properties": { - "detail": { - "items": { - "$ref": "#/components/schemas/ValidationError" - }, - "title": "Detail", - "type": "array" - } - }, - "title": "HTTPValidationError", - "type": "object" - }, - "ObservedAccount": { - "properties": { - "connection_id": { - "title": "Connection Id", - "type": "string" - }, - "login": { - "title": "Login", - "type": "string" - } - }, - "required": [ - "connection_id", - "login" - ], - "title": "ObservedAccount", - "type": "object" - }, - "ObservedApp": { - "properties": { - "api_url": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Api Url" - }, - "callback_url": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Callback Url" - }, - "can_install": { - "default": false, - "title": "Can Install", - "type": "boolean" - }, - "configured": { - "title": "Configured", - "type": "boolean" - } - }, - "required": [ - "configured" - ], - "title": "ObservedApp", - "type": "object" - }, - "ObservedApps": { - "properties": { - "github": { - "$ref": "#/components/schemas/ObservedApp" - }, - "gitlab": { - "$ref": "#/components/schemas/ObservedApp" - } - }, - "required": [ - "github", - "gitlab" - ], - "title": "ObservedApps", - "type": "object" - }, - "ObservedAuthorization": { - "properties": { - "url": { - "title": "Url", - "type": "string" - } - }, - "required": [ - "url" - ], - "title": "ObservedAuthorization", - "type": "object" - }, - "ObservedConnection": { - "properties": { - "api_url": { - "title": "Api Url", - "type": "string" - }, - "connection_type": { - "enum": [ - "token", - "app" - ], - "title": "Connection Type", - "type": "string" - }, - "has_token": { - "title": "Has Token", - "type": "boolean" - }, - "id": { - "default": "", - "title": "Id", - "type": "string" - }, - "ready": { - "title": "Ready", - "type": "boolean" - }, - "repos": { - "items": { - "type": "string" - }, - "title": "Repos", - "type": "array" - }, - "source_provider": { - "enum": [ - "github", - "gitlab" - ], - "title": "Source Provider", - "type": "string" - }, - "update_interval_minutes": { - "title": "Update Interval Minutes", - "type": "number" - } - }, - "required": [ - "source_provider", - "api_url", - "repos", - "has_token", - "update_interval_minutes", - "ready", - "connection_type" - ], - "title": "ObservedConnection", - "type": "object" - }, - "ObservedConnectionIdentities": { - "properties": { - "api_url": { - "title": "Api Url", - "type": "string" - }, - "id": { - "title": "Id", - "type": "string" - }, - "identity_map": { - "additionalProperties": { - "type": "string" - }, - "title": "Identity Map", - "type": "object" - }, - "repos": { - "items": { - "type": "string" - }, - "title": "Repos", - "type": "array" - }, - "source_provider": { - "enum": [ - "github", - "gitlab" - ], - "title": "Source Provider", - "type": "string" - }, - "unmatched_logins": { - "items": { - "type": "string" - }, - "title": "Unmatched Logins", - "type": "array" - } - }, - "required": [ - "id", - "source_provider", - "api_url", - "repos", - "identity_map", - "unmatched_logins" - ], - "title": "ObservedConnectionIdentities", - "type": "object" - }, - "ObservedHumanSummary": { - "properties": { - "median_merge_hours": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Median Merge Hours" - } - }, - "required": [ - "median_merge_hours" - ], - "title": "ObservedHumanSummary", - "type": "object" - }, - "ObservedIdentities": { - "properties": { - "connections": { - "default": [], - "items": { - "$ref": "#/components/schemas/ObservedConnectionIdentities" - }, - "title": "Connections", - "type": "array" - }, - "gateway_emails": { - "items": { - "type": "string" - }, - "title": "Gateway Emails", - "type": "array" - }, - "identity_map": { - "additionalProperties": { - "type": "string" - }, - "title": "Identity Map", - "type": "object" - }, - "unmatched_logins": { - "items": { - "type": "string" - }, - "title": "Unmatched Logins", - "type": "array" - } - }, - "required": [ - "gateway_emails", - "identity_map", - "unmatched_logins" - ], - "title": "ObservedIdentities", - "type": "object" - }, - "ObservedIdentityUpdate": { - "additionalProperties": false, - "properties": { - "accounts": { - "anyOf": [ - { - "items": { - "$ref": "#/components/schemas/ObservedAccount" - }, - "maxItems": 500, - "type": "array" - }, - { - "type": "null" - } - ], - "title": "Accounts" - }, - "email": { - "title": "Email", - "type": "string" - }, - "logins": { - "default": [], - "items": { - "type": "string" - }, - "maxItems": 100, - "title": "Logins", - "type": "array" - } - }, - "required": [ - "email" - ], - "title": "ObservedIdentityUpdate", - "type": "object" - }, - "ObservedPeriod": { - "properties": { - "agent_authored": { - "title": "Agent Authored", - "type": "integer" - }, - "agents_without_requester": { - "title": "Agents Without Requester", - "type": "integer" - }, - "explicitly_titled_revert_prs": { - "title": "Explicitly Titled Revert Prs", - "type": "integer" - }, - "human_authored": { - "title": "Human Authored", - "type": "integer" - }, - "human_summary": { - "$ref": "#/components/schemas/ObservedHumanSummary" - }, - "matched_internal_prs": { - "title": "Matched Internal Prs", - "type": "integer" - }, - "matched_users_recorded_spend": { - "title": "Matched Users Recorded Spend", - "type": "number" - }, - "median_merge_hours": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Median Merge Hours" - }, - "merged_prs": { - "title": "Merged Prs", - "type": "integer" - }, - "missing_author": { - "title": "Missing Author", - "type": "integer" - }, - "new_bug_labeled_issues": { - "anyOf": [ - { - "type": "integer" - }, - { - "type": "null" - } - ], - "title": "New Bug Labeled Issues" - }, - "new_regression_labeled_issues": { - "anyOf": [ - { - "type": "integer" - }, - { - "type": "null" - } - ], - "title": "New Regression Labeled Issues" - }, - "spend_observation": { - "enum": [ - "records_present", - "no_records" - ], - "title": "Spend Observation", - "type": "string" - }, - "window": { - "$ref": "#/components/schemas/ObservedWindow" - } - }, - "required": [ - "window", - "merged_prs", - "median_merge_hours", - "human_authored", - "agent_authored", - "missing_author", - "agents_without_requester", - "matched_internal_prs", - "new_bug_labeled_issues", - "new_regression_labeled_issues", - "explicitly_titled_revert_prs", - "matched_users_recorded_spend", - "spend_observation", - "human_summary" - ], - "title": "ObservedPeriod", - "type": "object" - }, - "ObservedPeriods": { - "properties": { - "current": { - "$ref": "#/components/schemas/ObservedPeriod" - }, - "last_year": { - "$ref": "#/components/schemas/ObservedPeriod" - }, - "previous": { - "$ref": "#/components/schemas/ObservedPeriod" - } - }, - "required": [ - "current", - "previous", - "last_year" - ], - "title": "ObservedPeriods", - "type": "object" - }, - "ObservedPerson": { - "properties": { - "accounts": { - "default": [], - "items": { - "$ref": "#/components/schemas/ObservedAccount" - }, - "title": "Accounts", - "type": "array" - }, - "email": { - "title": "Email", - "type": "string" - }, - "logins": { - "items": { - "type": "string" - }, - "title": "Logins", - "type": "array" - }, - "name": { - "title": "Name", - "type": "string" - }, - "periods": { - "$ref": "#/components/schemas/ObservedPersonPeriods" - } - }, - "required": [ - "name", - "email", - "logins", - "periods" - ], - "title": "ObservedPerson", - "type": "object" - }, - "ObservedPersonPeriod": { - "properties": { - "declared_agent_owned": { - "title": "Declared Agent Owned", - "type": "integer" - }, - "direct_authored": { - "title": "Direct Authored", - "type": "integer" - }, - "gateway_recorded_spend": { - "title": "Gateway Recorded Spend", - "type": "number" - }, - "median_merge_hours": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Median Merge Hours" - }, - "merged_prs": { - "title": "Merged Prs", - "type": "integer" - }, - "pr_urls": { - "items": { - "type": "string" - }, - "title": "Pr Urls", - "type": "array" - }, - "prs_per_week": { - "title": "Prs Per Week", - "type": "number" - }, - "recorded_spend_per_attributed_pr": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Recorded Spend Per Attributed Pr" - }, - "spend_observation": { - "enum": [ - "records_present", - "no_records" - ], - "title": "Spend Observation", - "type": "string" - } - }, - "required": [ - "merged_prs", - "prs_per_week", - "median_merge_hours", - "direct_authored", - "declared_agent_owned", - "gateway_recorded_spend", - "recorded_spend_per_attributed_pr", - "spend_observation", - "pr_urls" - ], - "title": "ObservedPersonPeriod", - "type": "object" - }, - "ObservedPersonPeriods": { - "properties": { - "current": { - "$ref": "#/components/schemas/ObservedPersonPeriod" - }, - "last_year": { - "$ref": "#/components/schemas/ObservedPersonPeriod" - }, - "previous": { - "$ref": "#/components/schemas/ObservedPersonPeriod" - } - }, - "required": [ - "current", - "previous", - "last_year" - ], - "title": "ObservedPersonPeriods", - "type": "object" - }, - "ObservedPullPeriods": { - "properties": { - "current": { - "items": { - "$ref": "#/components/schemas/ObservedPullResponse" - }, - "title": "Current", - "type": "array" - }, - "last_year": { - "items": { - "$ref": "#/components/schemas/ObservedPullResponse" - }, - "title": "Last Year", - "type": "array" - }, - "previous": { - "items": { - "$ref": "#/components/schemas/ObservedPullResponse" - }, - "title": "Previous", - "type": "array" - } - }, - "required": [ - "current", - "previous", - "last_year" - ], - "title": "ObservedPullPeriods", - "type": "object" - }, - "ObservedPullResponse": { - "properties": { - "agent": { - "default": false, - "title": "Agent", - "type": "boolean" - }, - "author": { - "title": "Author", - "type": "string" - }, - "branch_cost": { - "$ref": "#/components/schemas/ROIBranchAttribution" - }, - "connection_id": { - "default": "", - "title": "Connection Id", - "type": "string" - }, - "created_at": { - "anyOf": [ - { - "format": "date-time", - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Created At" - }, - "merge_hours": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Merge Hours" - }, - "merged_at": { - "format": "date-time", - "title": "Merged At", - "type": "string" - }, - "number": { - "title": "Number", - "type": "integer" - }, - "profile_email": { - "default": "", - "title": "Profile Email", - "type": "string" - }, - "repo": { - "title": "Repo", - "type": "string" - }, - "requester": { - "default": "", - "title": "Requester", - "type": "string" - }, - "source_branch": { - "default": "", - "title": "Source Branch", - "type": "string" - }, - "source_repo": { - "default": "", - "title": "Source Repo", - "type": "string" - }, - "title": { - "title": "Title", - "type": "string" - }, - "url": { - "title": "Url", - "type": "string" - } - }, - "required": [ - "repo", - "number", - "title", - "url", - "author", - "merged_at", - "merge_hours", - "branch_cost" - ], - "title": "ObservedPullResponse", - "type": "object" - }, - "ObservedReport": { - "properties": { - "captured_at": { - "format": "date-time", - "title": "Captured At", - "type": "string" - }, - "connections": { - "default": [], - "items": { - "$ref": "#/components/schemas/ObservedSource" - }, - "title": "Connections", - "type": "array" - }, - "people": { - "items": { - "$ref": "#/components/schemas/ObservedPerson" - }, - "title": "People", - "type": "array" - }, - "periods": { - "$ref": "#/components/schemas/ObservedPeriods" - }, - "pulls": { - "$ref": "#/components/schemas/ObservedPullPeriods" - }, - "repos": { - "items": { - "type": "string" - }, - "title": "Repos", - "type": "array" - }, - "source_provider": { - "enum": [ - "github", - "gitlab", - "mixed" - ], - "title": "Source Provider", - "type": "string" - }, - "unlinked_branches": { - "items": { - "$ref": "#/components/schemas/ROIBranchSpend" - }, - "title": "Unlinked Branches", - "type": "array" - }, - "unmatched_logins": { - "items": { - "type": "string" - }, - "title": "Unmatched Logins", - "type": "array" - } - }, - "required": [ - "source_provider", - "repos", - "captured_at", - "periods", - "people", - "pulls", - "unlinked_branches", - "unmatched_logins" - ], - "title": "ObservedReport", - "type": "object" - }, - "ObservedReportResponse": { - "properties": { - "report": { - "anyOf": [ - { - "$ref": "#/components/schemas/ObservedReport" - }, - { - "type": "null" - } - ] - } - }, - "required": [ - "report" - ], - "title": "ObservedReportResponse", - "type": "object" - }, - "ObservedSettings": { - "properties": { - "api_url": { - "title": "Api Url", - "type": "string" - }, - "connection_type": { - "enum": [ - "token", - "app" - ], - "title": "Connection Type", - "type": "string" - }, - "connections": { - "default": [], - "items": { - "$ref": "#/components/schemas/ObservedConnection" - }, - "title": "Connections", - "type": "array" - }, - "has_token": { - "title": "Has Token", - "type": "boolean" - }, - "id": { - "default": "", - "title": "Id", - "type": "string" - }, - "ready": { - "title": "Ready", - "type": "boolean" - }, - "repos": { - "items": { - "type": "string" - }, - "title": "Repos", - "type": "array" - }, - "source_provider": { - "enum": [ - "github", - "gitlab" - ], - "title": "Source Provider", - "type": "string" - }, - "update_interval_minutes": { - "title": "Update Interval Minutes", - "type": "number" - } - }, - "required": [ - "source_provider", - "api_url", - "repos", - "has_token", - "update_interval_minutes", - "ready", - "connection_type" - ], - "title": "ObservedSettings", - "type": "object" - }, - "ObservedSettingsUpdate": { - "additionalProperties": false, - "properties": { - "api_url": { - "title": "Api Url", - "type": "string" - }, - "connection_id": { - "anyOf": [ - { - "maxLength": 100, - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Connection Id" - }, - "repos": { - "items": { - "type": "string" - }, - "title": "Repos", - "type": "array" - }, - "source_provider": { - "enum": [ - "github", - "gitlab" - ], - "title": "Source Provider", - "type": "string" - }, - "token": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Token" - }, - "update_interval_minutes": { - "anyOf": [ - { - "maximum": 43200.0, - "minimum": 0.0, - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Update Interval Minutes" - } - }, - "required": [ - "source_provider", - "api_url", - "repos" - ], - "title": "ObservedSettingsUpdate", - "type": "object" - }, - "ObservedSource": { - "properties": { - "api_url": { - "title": "Api Url", - "type": "string" - }, - "id": { - "title": "Id", - "type": "string" - }, - "repos": { - "items": { - "type": "string" - }, - "title": "Repos", - "type": "array" - }, - "source_provider": { - "enum": [ - "github", - "gitlab" - ], - "title": "Source Provider", - "type": "string" - } - }, - "required": [ - "id", - "source_provider", - "api_url", - "repos" - ], - "title": "ObservedSource", - "type": "object" - }, - "ObservedWindow": { - "properties": { - "end": { - "format": "date", - "title": "End", - "type": "string" - }, - "start": { - "format": "date", - "title": "Start", - "type": "string" - } - }, - "required": [ - "start", - "end" - ], - "title": "ObservedWindow", - "type": "object" - }, - "ROIBranchAttribution": { - "properties": { - "branch": { - "title": "Branch", - "type": "string" - }, - "repo": { - "title": "Repo", - "type": "string" - }, - "requests": { - "default": 0, - "title": "Requests", - "type": "integer" - }, - "spend": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Spend" - }, - "status": { - "default": "unattributed", - "enum": [ - "matched", - "unattributed", - "ambiguous", - "unavailable" - ], - "title": "Status", - "type": "string" - } - }, - "required": [ - "repo", - "branch" - ], - "title": "ROIBranchAttribution", - "type": "object" - }, - "ROIBranchMetrics": { - "properties": { - "cost_per_hour": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Cost Per Hour" - }, - "hours": { - "default": 0, - "title": "Hours", - "type": "number" - }, - "matched_pulls": { - "default": 0, - "title": "Matched Pulls", - "type": "integer" - }, - "spend": { - "default": 0, - "title": "Spend", - "type": "number" - }, - "total_tagged_spend": { - "default": 0, - "title": "Total Tagged Spend", - "type": "number" - }, - "unlinked_spend": { - "default": 0, - "title": "Unlinked Spend", - "type": "number" - } - }, - "title": "ROIBranchMetrics", - "type": "object" - }, - "ROIBranchSpend": { - "properties": { - "branch": { - "title": "Branch", - "type": "string" - }, - "repo": { - "title": "Repo", - "type": "string" - }, - "requests": { - "title": "Requests", - "type": "integer" - }, - "spend": { - "title": "Spend", - "type": "number" - } - }, - "required": [ - "repo", - "branch", - "spend", - "requests" - ], - "title": "ROIBranchSpend", - "type": "object" - }, - "ROIEstimateResponse": { - "properties": { - "cached": { - "default": false, - "title": "Cached", - "type": "boolean" - }, - "effort_basis": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Effort Basis" - }, - "evidence_source": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Evidence Source" - }, - "hours": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Hours" - }, - "model": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Model" - }, - "reasoning": { - "title": "Reasoning", - "type": "string" - }, - "status": { - "enum": [ - "estimated", - "needs_review", - "error" - ], - "title": "Status", - "type": "string" - } - }, - "required": [ - "status", - "hours", - "reasoning" - ], - "title": "ROIEstimateResponse", - "type": "object" - }, - "ROIEstimatorModel": { - "properties": { - "model_name": { - "title": "Model Name", - "type": "string" - }, - "provider_models": { - "items": { - "type": "string" - }, - "title": "Provider Models", - "type": "array" - } - }, - "required": [ - "model_name", - "provider_models" - ], - "title": "ROIEstimatorModel", - "type": "object" - }, - "ROIIdentityMapResponse": { - "properties": { - "identity_map": { - "additionalProperties": { - "type": "string" - }, - "title": "Identity Map", - "type": "object" - }, - "report": { - "anyOf": [ - { - "$ref": "#/components/schemas/ROISummaryResponse" - }, - { - "type": "null" - } - ] - } - }, - "required": [ - "report", - "identity_map" - ], - "title": "ROIIdentityMapResponse", - "type": "object" - }, - "ROIIdentityMapUpdate": { - "properties": { - "email": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Email" - }, - "github_login": { - "title": "Github Login", - "type": "string" - } - }, - "required": [ - "github_login", - "email" - ], - "title": "ROIIdentityMapUpdate", - "type": "object" - }, - "ROIMetricsResponse": { - "properties": { - "cohort_people": { - "title": "Cohort People", - "type": "integer" - }, - "cost_per_hour": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Cost Per Hour" - }, - "estimated_prs": { - "title": "Estimated Prs", - "type": "integer" - }, - "excluded_spend": { - "title": "Excluded Spend", - "type": "number" - }, - "hours_per_dollar": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Hours Per Dollar" - }, - "matched_prs": { - "title": "Matched Prs", - "type": "integer" - }, - "matched_spend": { - "title": "Matched Spend", - "type": "number" - }, - "merged_prs": { - "title": "Merged Prs", - "type": "integer" - }, - "output_hours": { - "title": "Output Hours", - "type": "number" - }, - "pending_prs": { - "title": "Pending Prs", - "type": "integer" - }, - "people_with_prs": { - "title": "People With Prs", - "type": "integer" - }, - "total_output_hours": { - "title": "Total Output Hours", - "type": "number" - }, - "total_spend": { - "title": "Total Spend", - "type": "number" - } - }, - "required": [ - "matched_spend", - "output_hours", - "total_spend", - "total_output_hours", - "excluded_spend", - "cost_per_hour", - "hours_per_dollar", - "merged_prs", - "estimated_prs", - "matched_prs", - "cohort_people", - "people_with_prs", - "pending_prs" - ], - "title": "ROIMetricsResponse", - "type": "object" - }, - "ROIPersonResponse": { - "properties": { - "cost_per_hour": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Cost Per Hour" - }, - "eligible": { - "title": "Eligible", - "type": "boolean" - }, - "email": { - "title": "Email", - "type": "string" - }, - "estimated_prs": { - "title": "Estimated Prs", - "type": "integer" - }, - "hours": { - "title": "Hours", - "type": "number" - }, - "id": { - "title": "Id", - "type": "string" - }, - "logins": { - "items": { - "type": "string" - }, - "title": "Logins", - "type": "array" - }, - "match_methods": { - "items": { - "type": "string" - }, - "title": "Match Methods", - "type": "array" - }, - "pending_prs": { - "title": "Pending Prs", - "type": "integer" - }, - "prs": { - "title": "Prs", - "type": "integer" - }, - "spend": { - "anyOf": [ - { - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Spend" - } - }, - "required": [ - "id", - "email", - "logins", - "spend", - "hours", - "prs", - "estimated_prs", - "pending_prs", - "match_methods", - "eligible", - "cost_per_hour" - ], - "title": "ROIPersonResponse", - "type": "object" - }, - "ROIPullResponse": { - "properties": { - "additions": { - "title": "Additions", - "type": "integer" - }, - "branch_cost": { - "$ref": "#/components/schemas/ROIBranchAttribution" - }, - "cache_key": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Cache Key" - }, - "changed_files": { - "title": "Changed Files", - "type": "integer" - }, - "commit_count": { - "title": "Commit Count", - "type": "integer" - }, - "deletions": { - "title": "Deletions", - "type": "integer" - }, - "email": { - "title": "Email", - "type": "string" - }, - "emails": { - "items": { - "type": "string" - }, - "title": "Emails", - "type": "array" - }, - "estimate": { - "$ref": "#/components/schemas/ROIEstimateResponse" - }, - "head_sha": { - "title": "Head Sha", - "type": "string" - }, - "incomplete_metadata": { - "title": "Incomplete Metadata", - "type": "boolean" - }, - "login": { - "title": "Login", - "type": "string" - }, - "match_method": { - "title": "Match Method", - "type": "string" - }, - "matched": { - "title": "Matched", - "type": "boolean" - }, - "merged_at": { - "title": "Merged At", - "type": "string" - }, - "number": { - "title": "Number", - "type": "integer" - }, - "profile_email": { - "title": "Profile Email", - "type": "string" - }, - "repo": { - "title": "Repo", - "type": "string" - }, - "source_branch": { - "default": "", - "title": "Source Branch", - "type": "string" - }, - "source_repo": { - "default": "", - "title": "Source Repo", - "type": "string" - }, - "title": { - "title": "Title", - "type": "string" - }, - "url": { - "title": "Url", - "type": "string" - } - }, - "required": [ - "repo", - "number", - "title", - "url", - "login", - "emails", - "profile_email", - "merged_at", - "head_sha", - "additions", - "deletions", - "changed_files", - "commit_count", - "incomplete_metadata", - "estimate", - "email", - "match_method", - "matched" - ], - "title": "ROIPullResponse", - "type": "object" - }, - "ROIReportResponse": { - "properties": { - "report": { - "anyOf": [ - { - "$ref": "#/components/schemas/ROISummaryResponse" - }, - { - "type": "null" - } - ] - } - }, - "required": [ - "report" - ], - "title": "ROIReportResponse", - "type": "object" - }, - "ROIRepositoriesResponse": { - "properties": { - "has_more": { - "title": "Has More", - "type": "boolean" - }, - "page": { - "title": "Page", - "type": "integer" - }, - "repositories": { - "items": { - "$ref": "#/components/schemas/ROIRepository" - }, - "title": "Repositories", - "type": "array" - } - }, - "required": [ - "repositories", - "page", - "has_more" - ], - "title": "ROIRepositoriesResponse", - "type": "object" - }, - "ROIRepository": { - "properties": { - "archived": { - "title": "Archived", - "type": "boolean" - }, - "name": { - "title": "Name", - "type": "string" - }, - "visibility": { - "title": "Visibility", - "type": "string" - } - }, - "required": [ - "name", - "visibility", - "archived" - ], - "title": "ROIRepository", - "type": "object" - }, - "ROISettingsResponse": { - "properties": { - "available_models": { - "items": { - "type": "string" - }, - "title": "Available Models", - "type": "array" - }, - "backfill_days": { - "title": "Backfill Days", - "type": "integer" - }, - "default_prompt": { - "title": "Default Prompt", - "type": "string" - }, - "estimator_model": { - "title": "Estimator Model", - "type": "string" - }, - "estimator_models": { - "default": [], - "items": { - "$ref": "#/components/schemas/ROIEstimatorModel" - }, - "title": "Estimator Models", - "type": "array" - }, - "estimator_prompt": { - "title": "Estimator Prompt", - "type": "string" - }, - "github_api_url": { - "title": "Github Api Url", - "type": "string" - }, - "gitlab_api_url": { - "default": "https://gitlab.com/api/v4", - "title": "Gitlab Api Url", - "type": "string" - }, - "has_estimator_key": { - "title": "Has Estimator Key", - "type": "boolean" - }, - "has_github_token": { - "title": "Has Github Token", - "type": "boolean" - }, - "has_gitlab_token": { - "default": false, - "title": "Has Gitlab Token", - "type": "boolean" - }, - "identity_map": { - "additionalProperties": { - "type": "string" - }, - "title": "Identity Map", - "type": "object" - }, - "ready": { - "title": "Ready", - "type": "boolean" - }, - "report_mode": { - "default": "legacy", - "enum": [ - "legacy", - "observed" - ], - "title": "Report Mode", - "type": "string" - }, - "repos": { - "items": { - "type": "string" - }, - "title": "Repos", - "type": "array" - }, - "source_provider": { - "default": "github", - "enum": [ - "github", - "gitlab" - ], - "title": "Source Provider", - "type": "string" - }, - "update_interval_minutes": { - "title": "Update Interval Minutes", - "type": "number" - } - }, - "required": [ - "github_api_url", - "repos", - "estimator_model", - "estimator_prompt", - "backfill_days", - "update_interval_minutes", - "has_estimator_key", - "identity_map", - "has_github_token", - "default_prompt", - "available_models", - "ready" - ], - "title": "ROISettingsResponse", - "type": "object" - }, - "ROISettingsUpdate": { - "additionalProperties": false, - "properties": { - "backfill_days": { - "anyOf": [ - { - "maximum": 3650.0, - "minimum": 1.0, - "type": "integer" - }, - { - "type": "null" - } - ], - "title": "Backfill Days" - }, - "estimator_key": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Estimator Key" - }, - "estimator_model": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Estimator Model" - }, - "estimator_prompt": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Estimator Prompt" - }, - "github_api_url": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Github Api Url" - }, - "github_token": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Github Token" - }, - "gitlab_api_url": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Gitlab Api Url" - }, - "gitlab_token": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Gitlab Token" - }, - "report_mode": { - "anyOf": [ - { - "enum": [ - "legacy", - "observed" - ], - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Report Mode" - }, - "repos": { - "anyOf": [ - { - "items": { - "type": "string" - }, - "type": "array" - }, - { - "type": "null" - } - ], - "title": "Repos" - }, - "source_provider": { - "anyOf": [ - { - "enum": [ - "github", - "gitlab" - ], - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Source Provider" - }, - "update_interval_minutes": { - "anyOf": [ - { - "maximum": 43200.0, - "minimum": 0.0, - "type": "number" - }, - { - "type": "null" - } - ], - "title": "Update Interval Minutes" - } - }, - "title": "ROISettingsUpdate", - "type": "object" - }, - "ROISummaryResponse": { - "properties": { - "branch_metrics": { - "$ref": "#/components/schemas/ROIBranchMetrics" - }, - "effort_basis": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Effort Basis" - }, - "end": { - "title": "End", - "type": "string" - }, - "estimator_model": { - "title": "Estimator Model", - "type": "string" - }, - "estimator_prompt": { - "title": "Estimator Prompt", - "type": "string" - }, - "id": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Id" - }, - "metrics": { - "$ref": "#/components/schemas/ROIMetricsResponse" - }, - "mode": { - "title": "Mode", - "type": "string" - }, - "people": { - "items": { - "$ref": "#/components/schemas/ROIPersonResponse" - }, - "title": "People", - "type": "array" - }, - "pulls": { - "items": { - "$ref": "#/components/schemas/ROIPullResponse" - }, - "title": "Pulls", - "type": "array" - }, - "repos": { - "items": { - "type": "string" - }, - "title": "Repos", - "type": "array" - }, - "source_provider": { - "default": "github", - "enum": [ - "github", - "gitlab" - ], - "title": "Source Provider", - "type": "string" - }, - "start": { - "title": "Start", - "type": "string" - }, - "synced_at": { - "title": "Synced At", - "type": "string" - }, - "trend": { - "items": { - "$ref": "#/components/schemas/ROITrendResponse" - }, - "title": "Trend", - "type": "array" - }, - "unlinked_branches": { - "default": [], - "items": { - "$ref": "#/components/schemas/ROIBranchSpend" - }, - "title": "Unlinked Branches", - "type": "array" - }, - "warnings": { - "items": { - "type": "string" - }, - "title": "Warnings", - "type": "array" - } - }, - "required": [ - "id", - "mode", - "start", - "end", - "synced_at", - "repos", - "estimator_model", - "estimator_prompt", - "warnings", - "effort_basis", - "metrics", - "people", - "pulls", - "trend" - ], - "title": "ROISummaryResponse", - "type": "object" - }, - "ROISyncStatus": { - "properties": { - "done": { - "title": "Done", - "type": "integer" - }, - "elapsed_seconds": { - "default": 0, - "title": "Elapsed Seconds", - "type": "integer" - }, - "error": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Error" - }, - "estimated": { - "title": "Estimated", - "type": "integer" - }, - "finished_at": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Finished At" - }, - "needs_attention": { - "title": "Needs Attention", - "type": "integer" - }, - "next_update": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Next Update" - }, - "phase": { - "enum": [ - "idle", - "spend", - "repositories", - "estimates", - "complete", - "cancelled", - "error" - ], - "title": "Phase", - "type": "string" - }, - "remaining_seconds": { - "anyOf": [ - { - "type": "integer" - }, - { - "type": "null" - } - ], - "title": "Remaining Seconds" - }, - "reused": { - "title": "Reused", - "type": "integer" - }, - "running": { - "title": "Running", - "type": "boolean" - }, - "stage": { - "title": "Stage", - "type": "string" - }, - "started_at": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Started At" - }, - "total": { - "title": "Total", - "type": "integer" - } - }, - "required": [ - "running", - "phase", - "stage", - "done", - "total", - "estimated", - "reused", - "needs_attention", - "error" - ], - "title": "ROISyncStatus", - "type": "object" - }, - "ROITrendResponse": { - "properties": { - "date": { - "title": "Date", - "type": "string" - }, - "hours": { - "title": "Hours", - "type": "number" - }, - "prs": { - "title": "Prs", - "type": "integer" - }, - "spend": { - "title": "Spend", - "type": "number" - } - }, - "required": [ - "date", - "spend", - "hours", - "prs" - ], - "title": "ROITrendResponse", - "type": "object" - }, - "ValidationError": { - "properties": { - "ctx": { - "title": "Context", - "type": "object" - }, - "input": { - "title": "Input" - }, - "loc": { - "items": { - "anyOf": [ - { - "type": "string" - }, - { - "type": "integer" - } - ] - }, - "title": "Location", - "type": "array" - }, - "msg": { - "title": "Message", - "type": "string" - }, - "type": { - "title": "Error Type", - "type": "string" - } - }, - "required": [ - "loc", - "msg", - "type" - ], - "title": "ValidationError", - "type": "object" - } - } - }, - "paths": { - "/roi-calculator/connections/test": { - "post": { - "operationId": "test_roi_calculator_connections_roi_calculator_connections_test_post", - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROISettingsResponse" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Test Roi Calculator Connections", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/identity-map": { - "put": { - "operationId": "update_roi_calculator_identity_map_roi_calculator_identity_map_put", - "requestBody": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROIIdentityMapUpdate" - } - } - }, - "required": true - }, - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROIIdentityMapResponse" - } - } - }, - "description": "Successful Response" - }, - "422": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/HTTPValidationError" - } - } - }, - "description": "Validation Error" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Update Roi Calculator Identity Map", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/observed/apps": { - "get": { - "operationId": "observed_apps_roi_calculator_observed_apps_get", - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ObservedApps" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Observed Apps", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/observed/identities": { - "get": { - "operationId": "get_observed_identities_roi_calculator_observed_identities_get", - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ObservedIdentities" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Get Observed Identities", - "tags": [ - "roi_calculator" - ] - }, - "put": { - "operationId": "save_observed_identities_roi_calculator_observed_identities_put", - "requestBody": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ObservedIdentityUpdate" - } - } - }, - "required": true - }, - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ObservedReportResponse" - } - } - }, - "description": "Successful Response" - }, - "422": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/HTTPValidationError" - } - } - }, - "description": "Validation Error" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Save Observed Identities", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/observed/oauth/{provider}/start": { - "post": { - "operationId": "start_observed_authorization_roi_calculator_observed_oauth__provider__start_post", - "parameters": [ - { - "in": "path", - "name": "provider", - "required": true, - "schema": { - "enum": [ - "github", - "gitlab" - ], - "title": "Provider", - "type": "string" - } - }, - { - "in": "query", - "name": "install", - "required": false, - "schema": { - "default": false, - "title": "Install", - "type": "boolean" - } - } - ], - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ObservedAuthorization" - } - } - }, - "description": "Successful Response" - }, - "422": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/HTTPValidationError" - } - } - }, - "description": "Validation Error" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Start Observed Authorization", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/observed/report": { - "get": { - "operationId": "get_observed_report_roi_calculator_observed_report_get", - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ObservedReportResponse" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Get Observed Report", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/observed/repositories": { - "get": { - "operationId": "observed_repositories_roi_calculator_observed_repositories_get", - "parameters": [ - { - "in": "query", - "name": "connection", - "required": false, - "schema": { - "anyOf": [ - { - "maxLength": 100, - "type": "string" - }, - { - "type": "null" - } - ], - "title": "Connection" - } - }, - { - "in": "query", - "name": "query", - "required": false, - "schema": { - "default": "", - "maxLength": 200, - "title": "Query", - "type": "string" - } - }, - { - "in": "query", - "name": "page", - "required": false, - "schema": { - "default": 1, - "maximum": 1000, - "minimum": 1, - "title": "Page", - "type": "integer" - } - } - ], - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROIRepositoriesResponse" - } - } - }, - "description": "Successful Response" - }, - "422": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/HTTPValidationError" - } - } - }, - "description": "Validation Error" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Observed Repositories", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/observed/settings": { - "get": { - "operationId": "get_observed_settings_roi_calculator_observed_settings_get", - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ObservedSettings" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Get Observed Settings", - "tags": [ - "roi_calculator" - ] - }, - "put": { - "operationId": "save_observed_settings_roi_calculator_observed_settings_put", - "requestBody": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ObservedSettingsUpdate" - } - } - }, - "required": true - }, - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ObservedSettings" - } - } - }, - "description": "Successful Response" - }, - "422": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/HTTPValidationError" - } - } - }, - "description": "Validation Error" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Save Observed Settings", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/observed/sync": { - "delete": { - "operationId": "cancel_observed_sync_roi_calculator_observed_sync_delete", - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROISyncStatus" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Cancel Observed Sync", - "tags": [ - "roi_calculator" - ] - }, - "get": { - "operationId": "get_observed_sync_roi_calculator_observed_sync_get", - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROISyncStatus" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Get Observed Sync", - "tags": [ - "roi_calculator" - ] - }, - "post": { - "operationId": "start_observed_sync_roi_calculator_observed_sync_post", - "parameters": [ - { - "in": "query", - "name": "days", - "required": false, - "schema": { - "anyOf": [ - { - "maximum": 366, - "minimum": 1, - "type": "integer" - }, - { - "type": "null" - } - ], - "title": "Days" - } - } - ], - "responses": { - "202": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROISyncStatus" - } - } - }, - "description": "Successful Response" - }, - "422": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/HTTPValidationError" - } - } - }, - "description": "Validation Error" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Start Observed Sync", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/report": { - "get": { - "operationId": "get_roi_calculator_report_roi_calculator_report_get", - "parameters": [ - { - "in": "query", - "name": "mode", - "required": false, - "schema": { - "default": "live", - "enum": [ - "live", - "demo" - ], - "title": "Mode", - "type": "string" - } - } - ], - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROIReportResponse" - } - } - }, - "description": "Successful Response" - }, - "422": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/HTTPValidationError" - } - } - }, - "description": "Validation Error" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Get Roi Calculator Report", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/repositories": { - "get": { - "operationId": "get_roi_calculator_repositories_roi_calculator_repositories_get", - "parameters": [ - { - "in": "query", - "name": "query", - "required": false, - "schema": { - "default": "", - "maxLength": 200, - "title": "Query", - "type": "string" - } - }, - { - "in": "query", - "name": "page", - "required": false, - "schema": { - "default": 1, - "maximum": 1000, - "minimum": 1, - "title": "Page", - "type": "integer" - } - } - ], - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROIRepositoriesResponse" - } - } - }, - "description": "Successful Response" - }, - "422": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/HTTPValidationError" - } - } - }, - "description": "Validation Error" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Get Roi Calculator Repositories", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/settings": { - "get": { - "operationId": "get_roi_calculator_settings_roi_calculator_settings_get", - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROISettingsResponse" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Get Roi Calculator Settings", - "tags": [ - "roi_calculator" - ] - }, - "put": { - "operationId": "update_roi_calculator_settings_roi_calculator_settings_put", - "requestBody": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROISettingsUpdate" - } - } - }, - "required": true - }, - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROISettingsResponse" - } - } - }, - "description": "Successful Response" - }, - "422": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/HTTPValidationError" - } - } - }, - "description": "Validation Error" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Update Roi Calculator Settings", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/setup/reset": { - "post": { - "operationId": "reset_roi_calculator_setup_roi_calculator_setup_reset_post", - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROISettingsResponse" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Reset Roi Calculator Setup", - "tags": [ - "roi_calculator" - ] - } - }, - "/roi-calculator/sync": { - "delete": { - "operationId": "cancel_roi_calculator_sync_roi_calculator_sync_delete", - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROISyncStatus" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Cancel Roi Calculator Sync", - "tags": [ - "roi_calculator" - ] - }, - "get": { - "operationId": "get_roi_calculator_sync_status_roi_calculator_sync_get", - "responses": { - "200": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROISyncStatus" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Get Roi Calculator Sync Status", - "tags": [ - "roi_calculator" - ] - }, - "post": { - "operationId": "start_roi_calculator_sync_roi_calculator_sync_post", - "responses": { - "202": { - "content": { - "application/json": { - "schema": { - "$ref": "#/components/schemas/ROISyncStatus" - } - } - }, - "description": "Successful Response" - } - }, - "security": [ - { - "APIKeyHeader": [] - } - ], - "summary": "Start Roi Calculator Sync", - "tags": [ - "roi_calculator" - ] - } - } - } - }, "scim": { "components": { "schemas": { diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index c88c6f2570a..dca48babc81 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -19,7 +19,7 @@ from urllib.parse import urlparse from fastapi import APIRouter, Depends, HTTPException, Request, Response from fastapi.responses import JSONResponse, StreamingResponse -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError import litellm from litellm._logging import verbose_proxy_logger @@ -78,6 +78,9 @@ _PASCAL_TO_WIRE: Final[Mapping[str, str]] = { } +_DECODED_JSON: Final = TypeAdapter(object) + + def _sse_event(payload: object) -> str: """Frame a JSON-RPC object as a single A2A SSE event (``data: \\n\\n``).""" return f"data: {json.dumps(payload)}\n\n" @@ -91,7 +94,7 @@ def _to_jsonrpc_object(chunk: object) -> object: """ if isinstance(chunk, (str, bytes, bytearray)): try: - return json.loads(chunk) + return _DECODED_JSON.validate_python(json.loads(chunk)) except (json.JSONDecodeError, UnicodeDecodeError): return chunk if hasattr(chunk, "model_dump"): diff --git a/litellm/proxy/client/chat.py b/litellm/proxy/client/chat.py index a330d057490..cf48bee6958 100644 --- a/litellm/proxy/client/chat.py +++ b/litellm/proxy/client/chat.py @@ -3,9 +3,14 @@ from collections.abc import Iterator from typing import Any, Final import requests +from pydantic import ConfigDict, TypeAdapter from .exceptions import UnauthorizedError +_SSE_LINE: Final[TypeAdapter[bytes | bytearray]] = TypeAdapter( + bytes | bytearray, config=ConfigDict(arbitrary_types_allowed=True, strict=True, hide_input_in_errors=True) +) + class ChatClient: def __init__(self, base_url: str, api_key: str | None = None, timeout: int = 600): @@ -172,7 +177,7 @@ class ChatClient: # Parse SSE stream for line in response.iter_lines(): if line: - line = line.decode("utf-8") + line = _SSE_LINE.validate_python(line).decode("utf-8") if line.startswith("data: "): data_str = line[6:] # Remove 'data: ' prefix if data_str.strip() == "[DONE]": diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 98af32fa7aa..684449acd0c 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -8,6 +8,7 @@ from urllib.parse import urlencode import click import requests +from pydantic import ConfigDict, TypeAdapter from rich.console import Console from rich.table import Table from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never @@ -121,6 +122,8 @@ class CliAuthResult(TypedDict): _TeamMapping: Final = TypeVar("_TeamMapping", bound=Mapping[str, object]) +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + KEYRING_INSTALL_HINT: Final = "pip install 'litellm[cli]'" KEYRING_ENABLE_HINT: Final = "keyring --enable (or unset PYTHON_KEYRING_BACKEND)" @@ -485,7 +488,7 @@ def prompt_team_selection_fallback( def _response_error_detail(response: requests.Response) -> str | None: try: - body: Final[dict[str, object] | list[object] | str | int | float | bool | None] = response.json() + body: Final = _JSON_OBJECT.validate_python(response.json()) except ValueError: return None detail: Final = body.get("detail") if isinstance(body, dict) else None diff --git a/litellm/proxy/client/cli/commands/debug.py b/litellm/proxy/client/cli/commands/debug.py index 4e3914143d8..2aca60ca810 100644 --- a/litellm/proxy/client/cli/commands/debug.py +++ b/litellm/proxy/client/cli/commands/debug.py @@ -115,6 +115,7 @@ class RequestResponsePayload(BaseModel): _SESSION_PAGE: Final = TypeAdapter(SessionLogsPage) _PAYLOAD: Final[TypeAdapter[RequestResponsePayload | None]] = TypeAdapter(RequestResponsePayload | None) _JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +_BACKTICK_RUNS: Final = TypeAdapter(tuple[str, ...]) _SESSION_PAGE_SIZE: Final = 100 _TRANSPORT_BODY_CHARS: Final = 500 @@ -192,7 +193,7 @@ def _fmt_json(value: JsonValue, max_chars: int) -> str: def _fenced(text: str, info: str = "") -> tuple[str, str, str]: - longest_run: Final = max((len(run) for run in re.findall(r"`+", text)), default=0) + longest_run: Final = max((len(run) for run in _BACKTICK_RUNS.validate_python(re.findall(r"`+", text))), default=0) fence: Final = "`" * max(3, longest_run + 1) return (f"{fence}{info}", text, fence) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index c85169f0ba5..687df9b0348 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -272,6 +272,18 @@ def _withheld_provider_output(response: object) -> bool: return getattr(response, "has_buffered_provider_output", False) is True +async def close_guarded_stream(stream: object) -> None: + if not isinstance(stream, AsyncGenerator): + return + with anyio.CancelScope(shield=True): + try: + await stream.aclose() + except Exception as e: # noqa: BLE001 # a failing callback cleanup must not skip the refund and finalizer + verbose_proxy_logger.warning( + "Closing the guarded stream after a client disconnect raised %s", type(e).__name__ + ) + + def resolve_litellm_call_id(client_call_id: str | None) -> str: if client_call_id is not None and 0 < len(client_call_id) <= MAX_LITELLM_CALL_ID_LENGTH: return client_call_id @@ -3909,13 +3921,14 @@ class ProxyBaseLLMRequestProcessing: client_disconnected = False delivered_chunk = False recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the yielded bytes + guarded_stream: Final[AsyncGenerator[object, None]] = proxy_logging_obj.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=response, + request_data=request_data, + ) try: str_so_far = "" - async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=response, - request_data=request_data, - ): + async for chunk in guarded_stream: # ``.format(chunk)`` was previously evaluated for every chunk # regardless of log level; gate it behind the level check. if debug_enabled: @@ -3971,6 +3984,7 @@ class ProxyBaseLLMRequestProcessing: # Starlette closes on disconnect, so the nested iterator hook (which # only sees GeneratorExit on GC) cannot own the refund. client_disconnected = not stream_completed + await close_guarded_stream(guarded_stream) if not delivered_chunk and not _withheld_provider_output(response): from litellm.proxy.spend_tracking.budget_reservation import ( release_budget_reservation_on_cancel, diff --git a/litellm/proxy/common_utils/fips.py b/litellm/proxy/common_utils/fips.py new file mode 100644 index 00000000000..07bd4d2462c --- /dev/null +++ b/litellm/proxy/common_utils/fips.py @@ -0,0 +1,150 @@ +import hashlib +import os +from collections.abc import Callable +from dataclasses import dataclass +from typing import Final + +from typing_extensions import assert_never + +from litellm.secret_managers.main import str_to_bool + +FIPS_MODE_ENV_VAR: Final = "LITELLM_FIPS_MODE" +SSL_VERIFY_ENV_VAR: Final = "SSL_VERIFY" +SSL_VERIFY_SETTING: Final = "litellm_settings.ssl_verify" +REFUSAL_PREFIX: Final = "LiteLLM proxy refused to start" + +_TRUE_VALUES: Final = frozenset({"true", "1", "yes", "on"}) +_FALSE_VALUES: Final = frozenset({"false", "0", "no", "off", ""}) + + +@dataclass(frozen=True, slots=True) +class FipsModeOff: + pass + + +@dataclass(frozen=True, slots=True) +class FipsModeOn: + pass + + +@dataclass(frozen=True, slots=True) +class MalformedFipsMode: + value: str + + +FipsModeSetting = FipsModeOff | FipsModeOn | MalformedFipsMode + + +@dataclass(frozen=True, slots=True) +class ProviderDoesNotEnforceFips: + pass + + +@dataclass(frozen=True, slots=True) +class TlsVerificationDisabled: + sources: tuple[str, ...] + + +FipsBootRefusal = MalformedFipsMode | ProviderDoesNotEnforceFips | TlsVerificationDisabled +FipsBootVerdict = FipsModeOff | FipsModeOn | FipsBootRefusal + + +class FipsModeError(Exception): + pass + + +def parse_fips_mode(raw: str | None) -> FipsModeSetting: + if raw is None: + return FipsModeOff() + normalized: Final = raw.strip().lower() + if normalized in _TRUE_VALUES: + return FipsModeOn() + if normalized in _FALSE_VALUES: + return FipsModeOff() + return MalformedFipsMode(value=raw) + + +def is_fips_mode(environ: Callable[[str], str | None] = os.environ.get) -> bool: + return isinstance(parse_fips_mode(environ(FIPS_MODE_ENV_VAR)), FipsModeOn) + + +def openssl_enforces_fips() -> bool: + """MD5 is not an approved digest, so an enforcing FIPS provider refuses it even when asked for security use.""" + try: + hashlib.md5(b"", usedforsecurity=True) + except ValueError: + return True + return False + + +def fips_boot_verdict( + *, + raw_fips_mode: str | None, + provider_enforces_fips: Callable[[], bool], + ssl_verify_environment: str | None, + ssl_verify_setting: object, +) -> FipsBootVerdict: + setting: Final = parse_fips_mode(raw_fips_mode) + match setting: + case FipsModeOff() | MalformedFipsMode(): + return setting + case FipsModeOn(): + pass + case _: + assert_never(setting) + disabled: Final = tuple( + source + for source, off in ( + (SSL_VERIFY_ENV_VAR, _is_off(ssl_verify_environment)), + (SSL_VERIFY_SETTING, _is_off(ssl_verify_setting)), + ) + if off + ) + if disabled: + return TlsVerificationDisabled(sources=disabled) + if not provider_enforces_fips(): + return ProviderDoesNotEnforceFips() + return setting + + +def enforce_fips_boot_verdict(verdict: FipsBootVerdict, announce: Callable[[str], object]) -> None: + match verdict: + case FipsModeOff() | FipsModeOn(): + return + case MalformedFipsMode() | ProviderDoesNotEnforceFips() | TlsVerificationDisabled(): + message: Final = render_refusal(verdict) + announce(f"\n{message}\n\n") + raise FipsModeError(message) + case _: + assert_never(verdict) + + +def render_refusal(refusal: FipsBootRefusal) -> str: + match refusal: + case MalformedFipsMode(): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR}={refusal.value} is not a boolean.\n" + f"Set {FIPS_MODE_ENV_VAR} to true or false, or unset it." + ) + case ProviderDoesNotEnforceFips(): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR} is on but this Python does not enforce FIPS.\n" + "Its OpenSSL still allows non-approved algorithms (MD5 succeeded), so passwords and keys would be\n" + "protected with algorithms the FIPS 140-3 policy forbids. Run the proxy from a FIPS image whose\n" + f"OpenSSL FIPS provider is enabled, or unset {FIPS_MODE_ENV_VAR} on a non-FIPS runtime." + ) + case TlsVerificationDisabled(): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR} is on but TLS certificate verification is disabled by " + f"{' and '.join(refusal.sources)}.\nFIPS deployments must verify upstream certificates, so remove the " + "override or point ssl_verify at a CA bundle instead." + ) + return assert_never(refusal) + + +def _is_off(value: object) -> bool: + if isinstance(value, bool): + return value is False + if isinstance(value, str): + return str_to_bool(value) is False + return False diff --git a/litellm/proxy/db/model_insights_tasks.py b/litellm/proxy/db/model_insights_tasks.py index 865965dcf75..e117afb881e 100644 --- a/litellm/proxy/db/model_insights_tasks.py +++ b/litellm/proxy/db/model_insights_tasks.py @@ -1,14 +1,20 @@ import json +from collections.abc import Mapping from functools import lru_cache from pathlib import Path from typing import Final +from pydantic import ConfigDict, TypeAdapter + from litellm.types.model_insights import ModelInsightTask _TASKS_FILE: Final = Path(__file__).resolve().parent.parent / "model_insights_tasks.json" +_TASK_ENTRIES: Final = TypeAdapter( + Mapping[str, Mapping[str, object]], config=ConfigDict(strict=True, hide_input_in_errors=True) +) @lru_cache(maxsize=1) def load_model_insight_tasks() -> dict[str, ModelInsightTask]: - raw: Final = json.loads(_TASKS_FILE.read_text()) - return {name: ModelInsightTask(task_type=name, **entry) for name, entry in raw.items()} + raw: Final = _TASK_ENTRIES.validate_python(json.loads(_TASKS_FILE.read_text())) + return {name: ModelInsightTask.model_validate(dict(task_type=name, **entry)) for name, entry in raw.items()} diff --git a/litellm/proxy/fine_tuning_endpoints/endpoints.py b/litellm/proxy/fine_tuning_endpoints/endpoints.py index a13ad00713d..e09f1ec8ba8 100644 --- a/litellm/proxy/fine_tuning_endpoints/endpoints.py +++ b/litellm/proxy/fine_tuning_endpoints/endpoints.py @@ -426,7 +426,7 @@ async def list_fine_tuning_jobs( route_type=CallTypes.alist_fine_tuning_jobs.value, ) - response: Any | None = None + response: object = None if target_model_names and isinstance(target_model_names, str): target_model_names_list: Final = target_model_names.split(",") if len(target_model_names_list) != 1: diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 620b24df95d..f900d14bdc1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -10,6 +10,7 @@ import sys sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path import asyncio +import contextlib import copy import json import re @@ -2747,14 +2748,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): UnifiedLLMGuardrails, ) - async for streamed_chunk in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=response, - request_data=request_data, - guardrail_to_apply=self, - buffer_until_moderated_default=False, - ): - yield streamed_chunk + async with contextlib.aclosing( + UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=response, + request_data=request_data, + guardrail_to_apply=self, + buffer_until_moderated_default=False, + ) + ) as guarded: + async for streamed_chunk in guarded: + yield streamed_chunk return # Responses-API events are neither chat-completions chunks nor raw diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py new file mode 100644 index 00000000000..44b19f82218 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py @@ -0,0 +1,33 @@ +from typing import TYPE_CHECKING, Final + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .llm_shield_proxy import LLMShieldProxyGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail") -> LLMShieldProxyGuardrail: + import litellm + + _llm_shield_guardrail_callback: Final = LLMShieldProxyGuardrail( + api_key=litellm_params.api_key, + api_base=litellm_params.api_base, + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + + litellm.logging_callback_manager.add_litellm_callback(_llm_shield_guardrail_callback) + return _llm_shield_guardrail_callback + + +guardrail_initializer_registry: Final = { + SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: initialize_guardrail, +} + + +guardrail_class_registry: Final = { + SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: LLMShieldProxyGuardrail, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml new file mode 100644 index 00000000000..4f732773a75 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/example_config.yaml @@ -0,0 +1,14 @@ +model_list: + - model_name: gpt-4o + litellm_params: + model: openai/gpt-4o + api_key: os.environ/OPENAI_API_KEY + +guardrails: + - guardrail_name: "llm_shield_proxy" + litellm_params: + guardrail: llm_shield_proxy + mode: ["pre_call", "post_call"] + default_on: true + api_base: "http://localhost:8000" + api_key: os.environ/LLM_SHIELD_PROXY_API_KEY diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py new file mode 100644 index 00000000000..fa23a110a74 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -0,0 +1,744 @@ +import copy +import functools +import os +import uuid +from collections.abc import AsyncGenerator, Callable, Mapping, Sequence +from typing import ( + TYPE_CHECKING, + Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__ + ClassVar, + Final, + Literal, + Optional, +) + +import httpx + +from litellm._logging import verbose_proxy_logger +from litellm.exceptions import GuardrailRaisedException +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import GenericGuardrailAPIInputs, TextChoices + +if TYPE_CHECKING: + from litellm.caching.caching import DualCache + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.types.proxy.guardrails.guardrail_hooks.llm_shield_proxy import ( + LLMShieldProxyGuardrailConfigModel, + ) + from litellm.types.utils import CallTypes, LLMResponseTypes +from .payload import ( + JsonBody, + MutableRequest, + RequestTooDeep, + Slot, + SlotSink, + as_array, + as_object, + choice_index, + collect_json_leaves, + collect_response_item, + detached, + read_field, + read_list, + rehydrate_slots, + write_field, +) +from .request_walk import ( + locate_request_texts, +) +from .stream_restorers import ( + AnthropicSSERestorer, + CarryKey, + CarryWindows, + ResponsesStreamRestorer, + carry_sort_key, + continuation_delta, + responses_event_type, +) + +GUARDRAIL_NAME: Final = "llm_shield_proxy" + +_DEFAULT_API_BASE: Final = "http://localhost:8000" +_REDACT_PATH: Final = "/v1/guard/redact" +_REHYDRATE_PATH: Final = "/v1/guard/rehydrate" +_REHYDRATE_STREAM_PATH: Final = "/v1/guard/rehydrate/stream" + +_SESSION_METADATA_KEY: Final = "llm_shield_session_id" + +_DEPLOYMENT_RESTORE_KEY: Final = "llm_shield_restore_at_deployment" + +_VAULT_PREFIX: Final = f"litellm-{uuid.uuid4().hex}" + +_DEFAULT_TIMEOUT_SECONDS: Final = 10.0 + + +class LLMShieldProxyGuardrail(CustomGuardrail): + """Redacts PII before it leaves the proxy and restores it in the response. + + Unlike a masking guardrail, the substitution is reversible. Outbound text is + replaced with placeholders held in a session vault inside the user's own LLM + Shield deployment; the model's reply is then restored so the end user sees the + original values while the provider never received them. + + Streaming is restored incrementally rather than by buffering the response. LLM + Shield holds back only the trailing characters that could still turn out to be + part of a placeholder, so tokens are forwarded as they arrive and a placeholder + split across two chunks is never emitted in fragments. + """ + + use_native_lifecycle_hooks: ClassVar[bool] = True + + def __init__( + self, + guardrail_name: str = GUARDRAIL_NAME, + api_base: str | None = None, + api_key: str | None = None, + **kwargs: Any, # noqa: LIT008 # kwargs-ok: forwarded verbatim to CustomGuardrail.__init__ + ) -> None: + self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) + env_base: Final = os.environ.get("LLM_SHIELD_PROXY_API_BASE") + self.api_base: Final = (api_base or env_base or _DEFAULT_API_BASE).rstrip("/") + self.api_key: Final = api_key or os.environ.get("LLM_SHIELD_PROXY_API_KEY") + super().__init__(guardrail_name=guardrail_name, **kwargs) + + @classmethod + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: parent's signature. + return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] + + @staticmethod + def get_config_model() -> type["LLMShieldProxyGuardrailConfigModel"]: + from litellm.types.proxy.guardrails.guardrail_hooks.llm_shield_proxy import ( + LLMShieldProxyGuardrailConfigModel, + ) + + return LLMShieldProxyGuardrailConfigModel + + async def async_pre_call_deployment_hook( + self, + kwargs: MutableRequest, + call_type: "CallTypes | None", + ) -> MutableRequest | None: + """Redacts a model-level guardrail's request, and keeps it out of the response cache. + + Outside the proxy this hook is the only redaction step, and the deployment post-call + hook the only restoration step. LiteLLM builds the cache key after this hook, from + the redacted request, and a cache hit returns before the post-call hook runs. So a + cached reply would either reach the caller unrestored or, stored after restoration, + hand this caller's values to the next caller whose redacted request matches. The + request is therefore neither read from nor written to the cache. Inside the proxy + this hook does not redact -- the proxy's pre-call hook already ran -- and caching + is left alone, because the proxy restores after the cache write. + + A streamed request is refused once redacted. No hook restores an SDK stream, and the + stream's cache writer reads the request from before this hook, so it would also be + cached despite the bypass. + """ + before: Final = self._minted_session_id(kwargs) + _ = await super().async_pre_call_deployment_hook(kwargs, call_type) + session_id: Final = self._minted_session_id(kwargs) + if session_id is None or session_id == before: + return kwargs + if kwargs.get("stream") is True: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=( + "LLM Shield Proxy cannot restore a streamed reply for a model-level guardrail " + "outside the LiteLLM proxy; send the request through the proxy or without stream=True." + ), + ) + metadata: Final = as_object(kwargs.get("litellm_metadata")) + if metadata is not None: + metadata[_DEPLOYMENT_RESTORE_KEY] = session_id + cache_controls: Final = as_object(kwargs.get("cache")) + kwargs["cache"] = {**(cache_controls or {}), "no-cache": True, "no-store": True} + return kwargs + + async def async_post_call_success_deployment_hook( + self, + request_data: MutableRequest, + response: "LLMResponseTypes", + call_type: "CallTypes | None", + ) -> "LLMResponseTypes | None": + """Restores the reply here only when the deployment pre-call hook redacted it. + + LiteLLM caches what this hook returns. Inside the proxy the request was redacted by + the proxy's pre-call hook and the proxy's post-call hook restores the reply after + the cache write, so restoring here as well would cache this caller's plaintext under + a key built from the redacted request. Outside the proxy nothing restores later, and + the pre-call deployment hook has already kept that request out of the cache. + """ + metadata: Final = as_object(request_data.get("litellm_metadata")) + marker: Final = metadata.get(_DEPLOYMENT_RESTORE_KEY) if metadata is not None else None + session_id: Final = self._minted_session_id(request_data) + if session_id is None or marker != session_id: + return None + return await super().async_post_call_success_deployment_hook(request_data, response, call_type) + + def _headers(self, session_id: str) -> JsonBody: + headers: Final[JsonBody] = { + "Content-Type": "application/json", + "X-Session-ID": session_id, + } + if self.api_key: + headers["Authorization"] = f"Bearer {self.api_key}" + return headers + + async def _call_shield(self, path: str, session_id: str, payload: JsonBody) -> Mapping[str, object]: + """Posts to LLM Shield Proxy, failing closed on any transport or status error. + + A redaction guardrail that fails open sends the very data it exists to + protect to a third-party provider, so an unreachable or erroring shield + blocks the request instead of passing it through. + """ + try: + response: Final = await self.async_handler.post( + f"{self.api_base}{path}", + headers=self._headers(session_id), + json=payload, + timeout=_DEFAULT_TIMEOUT_SECONDS, + ) + response.raise_for_status() + return response.json() + except httpx.HTTPStatusError as exc: + verbose_proxy_logger.exception("LLM Shield Proxy returned %s for %s", exc.response.status_code, path) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"LLM Shield Proxy returned {exc.response.status_code}; blocking the request.", + ) from exc + except Exception as exc: + verbose_proxy_logger.exception("LLM Shield Proxy call to %s failed", path) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message="LLM Shield Proxy is unreachable; blocking the request.", + ) from exc + + async def _redact(self, texts: Sequence[str], session_id: str) -> Sequence[str]: + payload: Final[JsonBody] = {"texts": list(texts)} + body: Final = await self._call_shield(_REDACT_PATH, session_id, payload) + return self._same_length_or_raise(body.get("texts"), texts, "redact") + + async def _rehydrate(self, texts: Sequence[str], session_id: str) -> Sequence[str]: + payload: Final[JsonBody] = {"texts": list(texts)} + body: Final = await self._call_shield(_REHYDRATE_PATH, session_id, payload) + return self._same_length_or_raise(body.get("texts"), texts, "rehydrate") + + def _same_length_or_raise(self, returned: object, sent: Sequence[str], operation: str) -> Sequence[str]: + """Guards the positional mapping the callers rely on to write results back.""" + entries: Final = as_array(returned) + texts: Final = tuple(entry for entry in entries or () if isinstance(entry, str)) + if entries is None or len(entries) != len(sent) or len(texts) != len(entries): + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"LLM Shield Proxy {operation} returned an unexpected payload; blocking the request.", + ) + return texts + + @staticmethod + def _mint_session_id(data: MutableRequest) -> str: + """Mints a vault id for this request, overwriting anything already there. + + Redaction and restoration both happen inside one request/response pair, so + a fresh id per request is all that is needed, and it is what keeps one + caller from reaching another caller's vault. + """ + session_id: Final = f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" + metadata: Final = data.setdefault("litellm_metadata", {}) + if isinstance(metadata, dict): + metadata[_SESSION_METADATA_KEY] = session_id + return session_id + + @staticmethod + def _minted_session_id(data: MutableRequest) -> str | None: + """The vault id this process minted for `data`, or None if it has none. + + Read only from `litellm_metadata`, the proxy-private store `_mint_session_id` writes + to. A caller can populate `metadata`; they cannot populate this. + """ + metadata: Final = as_object(data.get("litellm_metadata")) + existing: Final = metadata.get(_SESSION_METADATA_KEY) if metadata is not None else None + return existing if isinstance(existing, str) and existing.startswith(_VAULT_PREFIX) else None + + @staticmethod + def _session_id(data: MutableRequest) -> str: + """Reads back the vault id minted while redacting this request. + + Falls back to an unused id rather than to anything the caller supplied: a + reply that cannot be restored is a visible placeholder, while trusting a + caller-supplied id would hand them someone else's plaintext. + """ + existing: Final = LLMShieldProxyGuardrail._minted_session_id(data) + return existing if existing is not None else f"{_VAULT_PREFIX}-{uuid.uuid4().hex}" + + @staticmethod + def _locate_request_texts(data: MutableRequest) -> tuple[Sequence[Slot], Sequence[Slot]]: + return locate_request_texts(data) + + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: "DualCache", + data: MutableRequest, + call_type: str, + ) -> MutableRequest | None: + """Replaces PII anywhere in the outbound request with vault placeholders.""" + if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True: + return data + + try: + slots, privileged = self._locate_request_texts(data) + except RequestTooDeep as exc: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"Request {exc} nests deeper than LLM Shield Proxy inspects; blocking the request.", + ) from exc + if not slots and not privileged: + return data + + session_id: Final = self._mint_session_id(data) + if privileged: + await self._redact_into(privileged, f"{_VAULT_PREFIX}-{uuid.uuid4().hex}") + if slots: + await self._redact_into(slots, session_id) + return data + + async def _redact_into(self, slots: Sequence[Slot], session_id: str) -> None: + """Redacts every span in `slots` under one vault and writes the result back.""" + redacted: Final = await self._redact(tuple(text for text, _ in slots), session_id) + for (_, write), replacement in zip(slots, redacted): + write(replacement) + + async def async_post_call_success_hook( + self, + data: MutableRequest, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + """Restores the original values in a copy of a non-streaming response. + + The copy is what keeps plaintext out of the response cache. LiteLLM caches the + reply it received from the provider, and on some paths -- a native Anthropic dict, + an in-memory cache -- it stores the object itself rather than a serialised + snapshot. Restoring that object in place would cache this caller's values under a + key built from the redacted request, which another caller's identical-looking + request then hits. Left untouched, the cached reply holds placeholders, and a hit + is restored against the new caller's own vault. + """ + if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True: + return response + response = detached(response) # rebind-ok: everything below restores the copy. + + if self._is_anthropic_message_response(response): + return await self._restore_anthropic_response(response, data) + + response_slots: Final = self._responses_api_slots(response) + if response_slots: + return await self._restore_responses_api_response(response, response_slots, data) + + choices: Final = getattr(response, "choices", None) + if not choices: + return response + + pending: Final[SlotSink] = [] + for choice in choices: + message = getattr(choice, "message", None) + if message is None: + text = read_field(choice, "text") + if isinstance(text, str) and text: + pending.append((text, functools.partial(write_field, choice, "text"))) + continue + content = getattr(message, "content", None) + if isinstance(content, str) and content: + pending.append((content, functools.partial(setattr, message, "content"))) + for tool_call in getattr(message, "tool_calls", None) or (): + function = getattr(tool_call, "function", None) + arguments = getattr(function, "arguments", None) if function is not None else None + if isinstance(arguments, str) and arguments: + pending.append((arguments, functools.partial(setattr, function, "arguments"))) + legacy = getattr(message, "function_call", None) + legacy_arguments = getattr(legacy, "arguments", None) if legacy is not None else None + if isinstance(legacy_arguments, str) and legacy_arguments: + pending.append((legacy_arguments, functools.partial(setattr, legacy, "arguments"))) + if not pending: + return response + + restored: Final = await self._rehydrate(tuple(text for text, _ in pending), self._session_id(data)) + for (_, write), replacement in zip(pending, restored): + write(replacement) + return response + + @staticmethod + def _is_anthropic_message_response(response: object) -> bool: + """Anthropic's native /v1/messages reply arrives as a plain dict.""" + body: Final = as_object(response) + return body is not None and body.get("type") == "message" and isinstance(body.get("content"), list) + + async def _restore_anthropic_response(self, response: MutableRequest, data: MutableRequest) -> MutableRequest: + """Restores text blocks and tool inputs in an Anthropic native message reply. + + This shape has no `choices`, so without its own branch the reply would go + back to the caller still carrying placeholders. + + A `tool_use` block's payload is `input`, an arbitrary JSON object rather than a + string, and the request path redacts its string leaves -- so the reply's leaves + have to come back or the application invokes the tool with placeholders. + """ + slots: Final[SlotSink] = [] + for entry in read_list(response, "content"): + block = as_object(entry) + if block is None: + continue + kind = block.get("type") + text = block.get("text") + if kind == "text" and isinstance(text, str) and text: + slots.append((text, functools.partial(block.__setitem__, "text"))) + elif kind == "tool_use" and as_object(block.get("input")) is not None: + collect_json_leaves(block.get("input"), slots) + if not slots: + return response + + restored: Final = await self._rehydrate(tuple(text for text, _ in slots), self._session_id(data)) + for (_, write), replacement in zip(slots, restored): + write(replacement) + return response + + @staticmethod + def _responses_api_slots(response: object) -> Sequence[Slot]: + """Restorable spans in a Responses API reply. + + That shape carries `output` items rather than `choices`, so it needs its own + walk; without one the reply goes back to the caller still holding + placeholders even though the request was redacted correctly. Items and blocks + come through as dicts or as objects depending on how far the reply has been + deserialised, so both are handled. + + The item-level fields mirror `collect_responses_fields`, which walks the same + fields on the request side -- a function_call item holds `arguments`, a + function_call_output holds `output` -- so the two directions stay symmetric. + """ + slots: Final[SlotSink] = [] + for item in getattr(response, "output", None) or (): + collect_response_item(item, slots) + return tuple(slots) + + async def _restore_responses_api_response(self, response: Any, slots: Sequence[Slot], data: MutableRequest) -> Any: + """Puts the original values back into a Responses API reply.""" + await rehydrate_slots(slots, functools.partial(self._rehydrate, session_id=self._session_id(data))) + return response + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + request_data: MutableRequest, + ) -> AsyncGenerator[Any, None]: + """Restores original values incrementally, without buffering the stream. + + Each choice -- and each tool call within a choice -- is its own token stream, so + the sliding window is tracked per (choice index, tool call) pair. One shared + window would splice the characters held back for one stream onto another. The + windows are locals of this generator, so they are scoped to a single stream and + cannot leak between concurrent requests. + + The two native stream shapes have no `choices` and are restored by their own + walkers, with the same per-stream windows: Anthropic `/v1/messages` arrives as raw + SSE frames, and the Responses API as typed events. + """ + if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True: + async for chunk in response: + yield chunk + return + + session_id: Final = self._session_id(request_data) + step: Final = functools.partial(self._stream_step, session_id=session_id) + rehydrate: Final = functools.partial(self._rehydrate, session_id=session_id) + sse: Final = AnthropicSSERestorer(step) + events: Final = ResponsesStreamRestorer(step, rehydrate) + carries: Final[CarryWindows] = {} + last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. + + async for chunk in response: + if isinstance(chunk, (bytes, str)): + for frames in await sse.feed(chunk): + yield frames + continue + if responses_event_type(chunk) is not None: + for event in await events.restore(detached(chunk)): + yield event + continue + restored_chunk = detached(chunk) + last_chunk = restored_chunk + for choice in getattr(restored_chunk, "choices", None) or (): + await self._restore_choice(choice, carries, session_id) + yield restored_chunk + + for frames in await sse.finish(): + yield frames + for event in await events.finish(): + yield event + if last_chunk is not None and any(carries.values()): + async for trailing in self._flush_trailing(last_chunk, carries, session_id): + yield trailing + + async def _restore_choice(self, choice: object, carries: CarryWindows, session_id: str) -> None: + """Restores one choice's delta, advancing that choice's own windows. + + Content and each tool call are separate token streams, so each gets its own + window: `(choice_index, None)` for content, `(choice_index, tool_call_index)` for + one tool call's accumulating `arguments`. A shared window would splice the text + held back for one stream onto another. + """ + delta: Final = getattr(choice, "delta", None) + index: Final = choice_index(choice) + is_final: Final = bool(getattr(choice, "finish_reason", None)) + if isinstance(choice, TextChoices): + await self._restore_text_window(choice, (index, None), carries, session_id, is_final) + return + if delta is None: + return + + await self._restore_content_window(delta, (index, None), carries, session_id, is_final) + + for tool_call in getattr(delta, "tool_calls", None) or (): + await self._restore_tool_call_window(tool_call, index, carries, session_id) + + if is_final: + await self._flush_finished_choice(delta, index, carries, session_id) + + async def _restore_text_window( + self, + choice: object, + key: CarryKey, + carries: CarryWindows, + session_id: str, + is_final: bool, + ) -> None: + """Restores a Completions stream choice's `text` through its window.""" + carry: Final = carries.get(key, "") + text: Final = read_field(choice, "text") + if not isinstance(text, str) or not text: + if not (is_final and carry): + return + emitted, remaining = await self._stream_step(text if isinstance(text, str) else "", carry, is_final, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + if emitted or text: + write_field(choice, "text", emitted) + + async def _restore_content_window( + self, + delta: Any, + key: CarryKey, + carries: CarryWindows, + session_id: str, + is_final: bool, + ) -> None: + """Restores one delta's content through its own window.""" + carry: Final = carries.get(key, "") + text: Final = getattr(delta, "content", None) + + if not isinstance(text, str) or not text: + if is_final and carry: + flushed, remaining = await self._stream_step("", carry, True, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + if flushed: + delta.content = flushed + return + + emitted, remaining = await self._stream_step(text, carry, is_final, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + delta.content = emitted + + async def _restore_tool_call_window( + self, + tool_call: Any, + choice_index: int, + carries: CarryWindows, + session_id: str, + ) -> None: + """Restores one streamed tool call's argument fragment. + + A tool call's `arguments` is a JSON document delivered as fragments that clients + concatenate per tool-call index, so each index gets a window of its own rather + than sharing the content stream's. + """ + tool_index: Final = read_field(tool_call, "index") + if not isinstance(tool_index, int): + return + function: Final = read_field(tool_call, "function") + if function is None: + return + arguments: Final = read_field(function, "arguments") + if not isinstance(arguments, str) or not arguments: + return + + key: Final = (choice_index, tool_index) + emitted, remaining = await self._stream_step(arguments, carries.get(key, ""), False, session_id) + carries[key] = remaining # rebind-ok: this tool call's window advances. + write_field(function, "arguments", emitted) + + async def _flush_finished_choice( + self, + delta: Any, + choice_index: int, + carries: CarryWindows, + session_id: str, + ) -> None: + """Emits everything this finishing choice still holds, into this chunk. + + A client parses a tool call's `arguments` when the chunk carrying the + finish_reason arrives, so a flush delivered afterwards is too late -- the client + has already tried to parse truncated JSON. Content lands back on `content`; held + tool-call text is appended as an index-only continuation entry, which is the shape + clients concatenate by index, so no id or name is needed. Appending is correct + even when this chunk already carried a fragment for that tool call. + """ + continuations: Final[list[dict[str, object]]] = [] # mutable-ok: built into this chunk's delta. + for key in sorted((held for held in carries if held[0] == choice_index), key=carry_sort_key): + carry = carries[key] + if not carry: + continue + _, tool_index = key + text, remaining = await self._stream_step("", carry, True, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + if not text: + continue + if tool_index is None: + delta.content = text + else: + continuations.extend(continuation_delta(tool_index, text)) + if continuations: + existing: Final = tuple(getattr(delta, "tool_calls", None) or ()) + delta.tool_calls = [*existing, *continuations] + + async def _flush_trailing( + self, last_chunk: Any, carries: CarryWindows, session_id: str + ) -> AsyncGenerator[Any, None]: + """Empties every window still holding text, one chunk per window. + + This is the net for a stream that ended with no finish_reason at all; a stream + that ended with one is flushed into its own terminal chunk by + `_flush_finished_choice`, because that is the moment a client parses tool + arguments. + + Driven by the windows rather than by the last chunk's choices. A choice that + finished earlier is not present in the terminal chunk, and flushing only what + that chunk carries would drop its held text and truncate its answer. + """ + for key in sorted(carries, key=carry_sort_key): + carry = carries[key] + if not carry: + continue + choice_index, tool_index = key + text, remaining = await self._stream_step("", carry, True, session_id) + carries[key] = remaining # rebind-ok: this stream's window advances. + if not text: + continue + chunk = self._chunk_for_choice(last_chunk, choice_index) + if chunk is None: + continue + choice: object = chunk.choices[0] + delta = read_field(choice, "delta") + if isinstance(choice, TextChoices): + choice.text = text + elif tool_index is None: + write_field(delta, "content", text) + else: + write_field(delta, "content", None) + write_field(delta, "tool_calls", continuation_delta(tool_index, text)) + yield chunk + + @staticmethod + def _chunk_for_choice(last_chunk: Any, index: int) -> Any: + """A single-choice copy of the last chunk, carrying only `index`. + + Emitting one choice per chunk keeps a flush from reading as content on a + choice it does not belong to. + """ + chunk: Final = last_chunk.model_copy(deep=True) + raw_choices: Final = getattr(chunk, "choices", None) + if not raw_choices: + return None + choices: Final[tuple[object, ...]] = tuple(raw_choices) + position: Final = next((at for at, choice in enumerate(choices) if choice_index(choice) == index), 0) + kept: Final = raw_choices[position] + if getattr(kept, "delta", None) is None and not isinstance(kept, TextChoices): + return None + kept.index = index + kept.finish_reason = None + chunk.choices = [kept] + if hasattr(chunk, "usage"): + del chunk.usage + return chunk + + async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]: + """Returns ``(text safe to emit now, window still being held)``.""" + body: Final = await self._call_shield( + _REHYDRATE_STREAM_PATH, + session_id, + {"text": text, "carry": carry, "final": final}, + ) + emitted: Final = body.get("text") + remaining: Final = body.get("carry") + if not isinstance(emitted, str) or not isinstance(remaining, str): + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message="LLM Shield Proxy stream rehydration returned an unexpected payload.", + ) + return emitted, remaining + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: MutableRequest, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"] = None, + ) -> GenericGuardrailAPIInputs: + """Unified entry point: what the UI's Test guardrail button and the translation + handlers call. + + `tool_calls` is handled on the response side only. LiteLLM populates the field here, + and on a reply it holds the model's tool arguments -- the same text the native hook + restores, and restoring one but not the other would leave the placeholder on + whichever path ran. The request side is left to the native pre-call hook, because + redacting it here as well would redact it twice. + """ + text_list: Final = tuple(inputs.get("texts") or ()) + tool_calls: Final = tuple(inputs.get("tool_calls") or ()) if input_type == "response" else () + if not text_list and not tool_calls: + return inputs + + restored_calls: Final[list[object]] = [copy.deepcopy(call) for call in tool_calls] # mutable-ok: a new list. + spans: Final[list[str]] = list(text_list) # mutable-ok: ordered batch, frozen before the call. + writers: Final[list[Callable[[str], None]]] = [] # mutable-ok: one per span appended below. + for call in restored_calls: + function = read_field(call, "function") + arguments = read_field(function, "arguments") if function is not None else None + if isinstance(arguments, str) and arguments: + spans.append(arguments) + writers.append(functools.partial(write_field, function, "arguments")) + + replaced: Final = ( + await self._redact(tuple(spans), self._mint_session_id(request_data)) + if input_type == "request" + else await self._rehydrate(tuple(spans), self._session_id(request_data)) + ) + restored_values: Final[list[str]] = list(replaced) # mutable-ok: sliced into the texts list. + + for write, replacement in zip(writers, restored_values[len(text_list) :]): + write(replacement) + merged: Final[JsonBody] = {**inputs} + if text_list: + merged["texts"] = restored_values[: len(text_list)] + if restored_calls: + merged["tool_calls"] = restored_calls + return merged diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/payload.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/payload.py new file mode 100644 index 00000000000..644ac76efb8 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/payload.py @@ -0,0 +1,190 @@ +import copy +import functools +from collections.abc import Awaitable, Callable, Sequence +from typing import ( + Final, + TypeAlias, +) + +MutableRequest: TypeAlias = dict[str, object] + +JsonBody: TypeAlias = dict[str, object] + +MAX_JSON_DEPTH: Final = 64 + +Slot: TypeAlias = tuple[str, Callable[[str], None]] + +StreamStep: TypeAlias = Callable[[str, str, bool], Awaitable[tuple[str, str]]] + +Rehydrate: TypeAlias = Callable[[Sequence[str]], Awaitable[Sequence[str]]] + +SlotSink: TypeAlias = list[Slot] + +MutableSeq: TypeAlias = list[object] + + +def as_object(value: object) -> MutableRequest | None: + """`value` as a JSON object, or None. + + `isinstance(value, dict)` alone leaves the keys and values unknown to the type + checker. A JSON object's keys are strings, so the type is stated once, here. + """ + return value if isinstance(value, dict) else None + + +def as_array(value: object) -> MutableSeq | None: + """`value` as a JSON array, or None. See `as_object`.""" + return value if isinstance(value, list) else None + + +def detached(value: object) -> object: + """A deep copy of a reply or chunk, for restoring without touching LiteLLM's own object. + + LiteLLM keeps the object it handed the hooks to fill its response cache and its + logs, so writing restored plaintext into that object would put it there too. + """ + return copy.deepcopy(value) + + +def is_container(value: object) -> bool: + """Whether `value` is a JSON object or array, without narrowing it to unknown types.""" + return isinstance(value, (dict, list)) + + +def collect(container: MutableRequest, key: str, slots: SlotSink) -> None: + """Records the string at `key`, along with the write that replaces it.""" + value: Final = container.get(key) + if isinstance(value, str) and value: + slots.append((value, functools.partial(container.__setitem__, key))) + + +def collect_entry(entries: MutableSeq, index: int, slots: SlotSink) -> None: + """Records a string held directly in a list, rather than under a key.""" + value: Final = entries[index] + if isinstance(value, str) and value: + slots.append((value, functools.partial(entries.__setitem__, index))) + + +class RequestTooDeep(Exception): + """A request nests text past a walk's bound. + + Skipping the rest would forward it unredacted while the guardrail reports as + enabled, so the pre-call hook refuses the request instead. + """ + + +def collect_text_parts(container: MutableRequest, key: str, slots: SlotSink) -> None: + """Collects the `text` of every part in the list held at `key`.""" + for entry in as_array(container.get(key)) or (): + part = as_object(entry) + if part is not None: + collect(part, "text", slots) + + +def choice_index(choice: object) -> int: + """Streaming choices are matched across chunks by their index.""" + index: Final = getattr(choice, "index", 0) + return index if isinstance(index, int) else 0 + + +def read_field(holder: object, name: str) -> object: + """Reads one field from a dict or from an object. + + LiteLLM's replies arrive as Pydantic models on some paths and as plain dicts on + others, depending how far they have been deserialised, so every response walk here + has to handle both shapes. + """ + fields: Final = as_object(holder) + if fields is not None: + return fields.get(name) + return getattr(holder, name, None) + + +def read_list(holder: object, name: str) -> Sequence[object]: + """Reads a list field from a dict or an object; anything else reads as empty. + + The entries are the reply's own objects, so writing through them edits the reply. + """ + value: Final = read_field(holder, name) + if isinstance(value, tuple): + return value + return tuple(as_array(value) or ()) + + +def write_field(holder: object, name: str, value: object) -> None: + """Writes one string field back into a dict or an object. Pairs with read_field.""" + if isinstance(holder, dict): + holder[name] = value + else: + setattr(holder, name, value) + + +def collect_json_leaves(node: object, slots: SlotSink, *, strict: bool = False) -> None: + """Collects every string leaf of a JSON-ish structure, with a write-back per leaf. + + An Anthropic `tool_use` block carries `input`, an arbitrary JSON object rather than a + string, so a value worth restoring can sit at any depth. Bounded by `MAX_JSON_DEPTH`: + the shape is caller or model controlled, and the bound is what stops a crafted one from + becoming an unbounded descent. Walked with an explicit stack rather than recursively, + so a deeply nested value cannot spend stack frames proportional to attacker-chosen + depth. + + `strict` is for the request side, where a leaf left behind would reach the provider + unredacted: past the bound it raises `RequestTooDeep`. On the reply side a leaf past + the bound just keeps its placeholder, which leaks nothing, so it is skipped. + """ + pending: Final[list[tuple[object, int]]] = [(node, 0)] # mutable-ok: local walk stack. + while pending: + current, current_depth = pending.pop() + if current_depth > MAX_JSON_DEPTH: + if strict and is_container(current) and current: + raise RequestTooDeep("json") + continue + current_object = as_object(current) + if current_object is not None: + for key in tuple(current_object): + value = current_object[key] + if isinstance(value, str) and value: + slots.append((value, functools.partial(current_object.__setitem__, key))) + else: + pending.append((value, current_depth + 1)) + continue + entries = as_array(current) + if entries is not None: + for index, value in enumerate(entries): + if isinstance(value, str) and value: + slots.append((value, functools.partial(entries.__setitem__, index))) + else: + pending.append((value, current_depth + 1)) + + +def collect_response_item(item: object, slots: SlotSink) -> None: + """Restorable spans in one Responses API output item, dict or object. + + Mirrors `collect_responses_fields` on the request side -- a function_call or + mcp_call item holds `arguments`, their outputs `output`, a reasoning item `summary` + parts -- so the two directions stay symmetric. A custom tool call carries `input` and + a code interpreter call `code`, both model-written. + """ + for block in read_list(item, "content"): + for field in ("text", "refusal"): + text = read_field(block, field) + if isinstance(text, str) and text: + slots.append((text, lambda new, b=block, f=field: write_field(b, f, new))) + for part in read_list(item, "summary"): + text = read_field(part, "text") + if isinstance(text, str) and text: + slots.append((text, lambda new, p=part: write_field(p, "text", new))) + for field in ("arguments", "output", "input", "code"): + value = read_field(item, field) + if isinstance(value, str) and value: + slots.append((value, lambda new, i=item, f=field: write_field(i, f, new))) + + +async def rehydrate_slots(slots: Sequence[Slot], rehydrate: Rehydrate) -> None: + """Restores every span in `slots` in one batch and writes each result back.""" + if not slots: + return + restored: Final = await rehydrate(tuple(text for text, _ in slots)) + for (_, write), replacement in zip(slots, restored): + write(replacement) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/request_walk.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/request_walk.py new file mode 100644 index 00000000000..c9c0d2ccc47 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/request_walk.py @@ -0,0 +1,363 @@ +from collections.abc import Sequence +from typing import ( + Final, +) + +from .payload import ( + MAX_JSON_DEPTH, + MutableRequest, + RequestTooDeep, + Slot, + SlotSink, + as_array, + as_object, + collect, + collect_entry, + collect_json_leaves, + collect_text_parts, + is_container, + read_list, +) + +PRIVILEGED_ROLES: Final = frozenset({"system", "developer"}) + +MAX_CONTENT_DEPTH: Final = 8 + +SCHEMA_STRUCTURAL_KEYWORDS: Final = frozenset( + ( + "type", + "format", + "pattern", + "required", + "dependentRequired", + "propertyOrdering", + "discriminator", + "contentEncoding", + "contentMediaType", + "$ref", + "$id", + "$schema", + "$anchor", + "$dynamicRef", + "$dynamicAnchor", + "$recursiveRef", + "$recursiveAnchor", + "$vocabulary", + ) +) + +SCHEMA_VALUE_KEYWORDS: Final = frozenset(("examples", "default")) + +SCHEMA_LITERAL_KEYWORDS: Final = frozenset(("enum", "const")) + +SCHEMA_MAP_KEYWORDS: Final = frozenset( + ("properties", "patternProperties", "$defs", "definitions", "dependentSchemas", "dependencies") +) + + +def collect_prompt(data: MutableRequest, slots: SlotSink) -> None: + """The Completions API sends its text in `prompt`, and its tail in `suffix`.""" + collect(data, "suffix", slots) + prompt: Final = data.get("prompt") + if isinstance(prompt, str): + collect(data, "prompt", slots) + return + prompt_object: Final = as_object(prompt) + if prompt_object is not None: + variables: Final = as_object(prompt_object.get("variables")) + if variables is not None: + for name in tuple(variables): + collect(variables, name, slots) + typed = as_object(variables[name]) + if typed is not None: + collect(typed, "text", slots) + return + entries: Final = as_array(prompt) + if entries is None: + return + for index in range(len(entries)): + collect_entry(entries, index, slots) + + +def collect_content(container: MutableRequest, slots: SlotSink) -> None: + """Collects `content`, a string or a list of typed parts. + + An Anthropic tool_result nests its own content, so this has to descend. It walks + with an explicit stack and a depth bound rather than by recursion: the nesting is + caller controlled, and an unbounded descent is a JSON bomb. Content nested past the + bound raises `RequestTooDeep` rather than being skipped. + """ + pending: Final[list[tuple[MutableRequest, int]]] = [(container, 0)] # mutable-ok: local queue, never escapes. + cursor = 0 # rebind-ok: advances through the queue. + while cursor < len(pending): + node, depth = pending[cursor] + cursor += 1 + content = node.get("content") + if isinstance(content, str): + collect(node, "content", slots) + continue + if depth >= MAX_CONTENT_DEPTH and content: + raise RequestTooDeep("content") + for item in as_array(content) or (): + part = as_object(item) + if part is None: + continue + collect(part, "text", slots) + if part.get("type") == "tool_use": + collect_json_leaves(part.get("input"), slots, strict=True) + source = as_object(part.get("source")) if part.get("type") == "document" else None + if source is not None: + collect(part, "title", slots) + collect(part, "context", slots) + if source.get("type") == "text": + collect(source, "data", slots) + elif source.get("type") == "content": + pending.append((source, depth + 1)) + if "content" in part: + pending.append((part, depth + 1)) + + +def collect_participant_name(message: MutableRequest, slots: SlotSink) -> None: + """Redacts `name` where it identifies a person, never where it names a function. + + On a user or assistant turn `name` is the participant, which is personal data. + On a tool or function turn the same field carries the function's name and has + to reach the provider unchanged, or the call no longer routes. + """ + if message.get("role") in ("tool", "function"): + return + collect(message, "name", slots) + + +def collect_tool_arguments(message: MutableRequest, slots: SlotSink) -> None: + """Tool arguments carry the values a user asked the model to act on.""" + for tool_call in read_list(message, "tool_calls"): + tool_call_object = as_object(tool_call) + function = as_object(tool_call_object.get("function")) if tool_call_object is not None else None + if function is not None: + collect(function, "arguments", slots) + legacy: Final = as_object(message.get("function_call")) + if legacy is not None: + collect(legacy, "arguments", slots) + + +def collect_system(data: MutableRequest, slots: SlotSink) -> None: + """Anthropic's /v1/messages carries its system prompt at the top level.""" + system: Final = data.get("system") + if isinstance(system, str): + collect(data, "system", slots) + return + collect_text_parts(data, "system", slots) + + +def collect_responses_fields(data: MutableRequest, slots: SlotSink, privileged: SlotSink) -> None: + """The Responses API sends text outside `messages`, in `instructions` and `input`. + + `instructions` is written by the application, not by the caller, so it is + collected into the privileged sink; `input` is the caller's own text, except for + system and developer items in it, which go to the privileged sink like their Chat + counterparts. + """ + collect(data, "instructions", privileged) + request_input: Final = data.get("input") + if isinstance(request_input, str): + collect(data, "input", slots) + return + entries: Final = as_array(request_input) + if entries is None: + return + for index, entry in enumerate(entries): + if isinstance(entry, str): + collect_entry(entries, index, slots) + continue + item = as_object(entry) + if item is None: + continue + collect_content(item, privileged if item.get("role") in PRIVILEGED_ROLES else slots) + collect(item, "arguments", slots) + collect(item, "output", slots) + collect_text_parts(item, "output", slots) + collect(item, "input", slots) + collect(item, "code", slots) + collect_text_parts(item, "summary", slots) + + +def collect_tool_definitions(data: MutableRequest, slots: SlotSink, privileged: SlotSink) -> None: + """Tool definitions are application-authored free text bound for the provider. + + A tool's description and the free text in its parameter schema are where callers put + examples and customer context, so they carry PII as often as a prompt does. They are + collected into the privileged sink, like a system prompt: redacted outbound, and never + restorable from the reply. `enum` and `const` values are the exception, and go to the + caller's vault -- see `SCHEMA_LITERAL_KEYWORDS`. Names and types are left as sent. + + Covers Chat `tools[].function`, the legacy `functions[]`, and the flat tool shape the + Responses API and Anthropic share, whose schema is `parameters` or `input_schema`. + """ + for key in ("tools", "functions"): + for entry in as_array(data.get(key)) or (): + tool = as_object(entry) + if tool is None: + continue + function = as_object(tool.get("function")) + for holder in (tool, function) if function is not None else (tool,): + collect(holder, "description", privileged) + collect_schema_text(holder.get("parameters"), slots, privileged) + collect_schema_text(holder.get("input_schema"), slots, privileged) + + +def collect_schema_text(schema: object, slots: SlotSink, privileged: SlotSink) -> None: + """Collects the text in a JSON Schema, at any depth. + + Scan by default: every string is collected except under the keywords in + `SCHEMA_STRUCTURAL_KEYWORDS`, whose values must go out verbatim. A list of keywords + *to* collect would leak every one it forgot -- draft-07 `dependencies`, a `$comment`, + a vendor `x-` extension -- which is how this walk started out. Free text goes to the + privileged sink; `enum` / `const` literals go to the caller's, so the model's use of + them is restored. + + Structure matters in two places. Under `properties` and the other name -> subschema + maps, keys are property names rather than keywords, so a property called `type` is a + subschema to walk, not a keyword to skip. And `examples` / `default` hold JSON values, + so all their strings are collected whatever the keys around them are called. Nested + past `MAX_JSON_DEPTH`, the request is refused. + """ + pending: Final[list[tuple[object, int]]] = [(schema, 0)] # mutable-ok: local walk stack. + while pending: + node, depth = pending.pop() + if depth > MAX_JSON_DEPTH: + if is_container(node) and node: + raise RequestTooDeep("schema") + continue + entries = as_array(node) + if entries is not None: + for index, item in enumerate(entries): + collect_entry(entries, index, privileged) + if is_container(item): + pending.append((item, depth + 1)) + continue + schema_object = as_object(node) + if schema_object is None: + continue + for keyword, value in tuple(schema_object.items()): + if keyword in SCHEMA_STRUCTURAL_KEYWORDS: + continue + subschemas = as_object(value) if keyword in SCHEMA_MAP_KEYWORDS else None + if keyword in SCHEMA_LITERAL_KEYWORDS: + collect(schema_object, keyword, slots) + collect_json_leaves(value, slots, strict=True) + elif keyword in SCHEMA_VALUE_KEYWORDS: + collect(schema_object, keyword, privileged) + collect_json_leaves(value, privileged, strict=True) + elif subschemas is not None: + pending.extend((child, depth + 1) for child in subschemas.values()) + elif isinstance(value, str): + collect(schema_object, keyword, privileged) + elif is_container(value): + pending.append((value, depth + 1)) + + +def collect_output_contracts(data: MutableRequest, slots: SlotSink, privileged: SlotSink) -> None: + """Text the caller sends to shape the reply rather than to prompt it. + + A predicted output (`prediction.content`) is the caller's own draft of the answer, so + it goes with their text: the model largely repeats it, and it has to come back. A + structured-output schema -- Chat `response_format.json_schema`, Responses + `text.format` -- is application-authored like a tool schema, so its free text goes + to the privileged sink, and its names and types stay as sent. + """ + prediction: Final = as_object(data.get("prediction")) + if prediction is not None: + collect(prediction, "content", slots) + collect_text_parts(prediction, "content", slots) + response_format: Final = as_object(data.get("response_format")) + text_options: Final = as_object(data.get("text")) + for declared in ( + response_format.get("json_schema") if response_format is not None else None, + text_options.get("format") if text_options is not None else None, + ): + wrapper = as_object(declared) + if wrapper is not None: + collect(wrapper, "description", privileged) + collect_schema_text(wrapper.get("schema"), slots, privileged) + + +def collect_user_locations(data: MutableRequest, privileged: SlotSink) -> None: + """Web search forwards the user's approximate location, whose `city` and `region` + are free text and can hold a street address. + + Chat carries it in `web_search_options.user_location.approximate`; the Responses + and Anthropic web-search tools carry it flat on the tool's `user_location`. Nothing + restores it from a reply, hence the privileged sink. + """ + options: Final = data.get("web_search_options") + tools: Final = as_array(data.get("tools")) or () + for declared in (options, *tools): + holder = as_object(declared) + location = as_object(holder.get("user_location")) if holder is not None else None + if location is None: + continue + approximate = as_object(location.get("approximate")) + for container in (location, approximate) if approximate is not None else (location,): + collect(container, "city", privileged) + collect(container, "region", privileged) + + +def collect_end_user_ids(data: MutableRequest, privileged: SlotSink) -> None: + """`user` and `safety_identifier` are forwarded to the provider and often hold an email. + + Only detected PII is replaced, so an opaque id reaches the provider unchanged. LiteLLM's + own end-user spend tracking reads the id resolved at authentication, before this hook + runs, so rewriting the field here does not move spend. Nothing restores these from a + reply, hence the privileged sink. + """ + collect(data, "user", privileged) + collect(data, "safety_identifier", privileged) + + +def locate_request_texts( + data: MutableRequest, +) -> tuple[Sequence[Slot], Sequence[Slot]]: + """Finds every redactable span, split by whether the caller can see it. + + Anything missed here reaches the provider in the clear while the guardrail + still reports as enabled, so the walk covers every request shape that + carries text. + + The split exists because the response is restored against one vault only. + Server-authored spans -- system and developer turns, Anthropic's top-level + `system`, the Responses API `instructions`, tool and output schemas -- go into a + vault nothing is ever restored against, so a caller who gets the model to + echo one of their placeholders back receives the placeholder, not the value + behind it. End-user identifiers go there too: nothing in a reply needs them. + + Tool *results* stay on the caller's side deliberately. The model reads them in + order to answer, so it can already repeat anything in them; restoring the + placeholder gives the caller the answer they would have had without this + guardrail, and an agent that reads a file and quotes an address from it needs + that address back. + + `extra_body` is walked the same way as the request itself. LiteLLM merges it over + the transformed request just before sending, so a field there -- `input`, + `messages`, `system` -- replaces the redacted one on the wire. + """ + slots: Final[SlotSink] = [] + privileged: Final[SlotSink] = [] + for payload in (data, as_object(data.get("extra_body"))): + if payload is None: + continue + for entry in read_list(payload, "messages"): + message = as_object(entry) + if message is not None: + sink = privileged if message.get("role") in PRIVILEGED_ROLES else slots + collect_content(message, sink) + collect_participant_name(message, sink) + collect_tool_arguments(message, sink) + collect_responses_fields(payload, slots, privileged) + collect_prompt(payload, slots) + collect_system(payload, privileged) + collect_tool_definitions(payload, slots, privileged) + collect_output_contracts(payload, slots, privileged) + collect_user_locations(payload, privileged) + collect_end_user_ids(payload, privileged) + return tuple(slots), tuple(privileged) diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/stream_restorers.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/stream_restorers.py new file mode 100644 index 00000000000..4ef8bbdbb3e --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/stream_restorers.py @@ -0,0 +1,345 @@ +import copy +import functools +import itertools +import json +import re +from enum import Enum +from types import MappingProxyType +from typing import ( + Final, + TypeAlias, +) + +from .payload import ( + JsonBody, + MutableRequest, + Rehydrate, + SlotSink, + StreamStep, + as_object, + collect_response_item, + read_field, + read_list, + rehydrate_slots, + write_field, +) + +ANTHROPIC_DELTA_FIELDS: Final = MappingProxyType({"text_delta": "text", "input_json_delta": "partial_json"}) + +SSE_EVENT_BOUNDARY: Final = re.compile(rb"(\r?\n\r?\n)") + +SSE_OPENINGS: Final = (b"event:", b"data:", b"id:", b"retry:", b":") + +RESPONSES_BINARY_DELTAS: Final = frozenset(("response.audio.delta",)) + +RESPONSES_STRUCTURAL_FIELDS: Final = frozenset( + ("type", "id", "item_id", "call_id", "name", "server_label", "status", "obfuscation") +) + +RESPONSES_TERMINAL_EVENTS: Final = frozenset(("response.completed", "response.incomplete")) + +CarryKey: TypeAlias = tuple[int, int | None] +CarryWindows: TypeAlias = dict[CarryKey, str] + +ResponsesStreamKey: TypeAlias = tuple[str, object, object, object] + + +def carry_sort_key(key: CarryKey) -> tuple[int, int]: + """Orders streaming windows without ever comparing None to an int. + + `sorted()` over the raw keys raises as soon as one choice holds both a content window + and a tool-call window, because `None < 0` is not orderable. Content sorts first, then + tool calls by their index. + """ + choice_index, tool_index = key + return (choice_index, -1 if tool_index is None else tool_index) + + +def continuation_delta(tool_index: int, text: str) -> list[dict[str, object]]: + """A `tool_calls` delta carrying `text` as an index-only continuation. + + Clients concatenate tool-call fragments by index, so no id or name is needed. + """ + return [{"index": tool_index, "function": {"arguments": text}}] + + +def opens_like_sse(head: bytes) -> bool | None: + """Whether a raw stream is SSE, judged by its opening bytes; None while undecidable. + + An SSE stream opens with a field name or a `:` comment. Anything else -- a JSON array + streamed in pieces, say -- has no event boundaries to wait for. A chunk that ends + partway through a field name decides nothing yet, so that case waits for more. + """ + opening: Final = head.lstrip() + if not opening: + return None + if opening.startswith(SSE_OPENINGS): + return True + if any(field.startswith(opening) for field in SSE_OPENINGS): + return None + return False + + +def responses_event_type(chunk: object) -> str | None: + """The event type of a Responses API stream event, or None for any other chunk. + + The type arrives as a plain string on dicts and as a str-valued Enum on LiteLLM's + event models. The Enum is unwrapped because it does not hash like its value, so it + would miss every lookup in the event tables above. + """ + if isinstance(chunk, (bytes, str)): + return None + kind: Final = read_field(chunk, "type") + value: Final = kind.value if isinstance(kind, Enum) else kind + return value if isinstance(value, str) and value.startswith("response.") else None + + +class AnthropicSSERestorer: + """Restores an Anthropic `/v1/messages` stream, which reaches the hook as raw SSE. + + Each content block is its own token stream with its own window, keyed by the block's + `index`: `text_delta` carries prose and `input_json_delta` a tool call's arguments. + When a block stops, whatever its window still holds is emitted as one more delta for + that block, just ahead of the `content_block_stop` frame, so the client has the whole + block before it is told the block is complete. + + Frames are processed whole. A network chunk can end in the middle of an event, so the + unfinished tail is kept until the rest arrives; that delays one partial event, never + a completed one. A frame that is not an Anthropic event -- another endpoint's SSE, or + anything that fails to parse -- is passed through byte for byte, and a raw stream that + does not open like SSE at all is passed through chunk by chunk, never buffered. + """ + + def __init__(self, step: StreamStep) -> None: + self._step: Final = step + self._carries: Final[dict[int, str]] = {} # mutable-ok: per-block windows advanced in place. + self._delta_types: Final[dict[int, str]] = {} # mutable-ok: each block's delta type, for its flush. + self._pending = b"" + self._as_text = False + self._is_sse: bool | None = None + + async def feed(self, chunk: bytes | str) -> tuple[bytes | str, ...]: + """Restores every event this chunk completes; holds back an unfinished tail.""" + if isinstance(chunk, str): + self._as_text = True + if self._is_sse is False: + return (chunk,) + raw: Final = chunk.encode("utf-8") if isinstance(chunk, str) else chunk + buffered: Final = self._pending + raw + if self._is_sse is None: + self._is_sse = opens_like_sse(buffered) + if self._is_sse is None: + self._pending = buffered + return () + if not self._is_sse: + self._pending = b"" + return self._emit(buffered) + boundaries: Final = tuple(SSE_EVENT_BOUNDARY.finditer(buffered)) + if not boundaries: + self._pending = buffered + return () + cut: Final = boundaries[-1].end() + self._pending = buffered[cut:] + parts: Final = SSE_EVENT_BOUNDARY.split(buffered[:cut]) + restored: Final = tuple( + [await self._restore_event(parts[index]) + parts[index + 1] for index in range(0, len(parts) - 1, 2)] + ) + return self._emit(b"".join(restored)) + + async def finish(self) -> tuple[bytes | str, ...]: + """Emits an unterminated final event and any window a block never closed.""" + held: Final = self._pending + self._pending = b"" + if not self._is_sse: + return self._emit(held) + tail: Final = await self._restore_event(held) if held.strip() else held + flushed: Final = await self._flush_all() + separator: Final = b"\n\n" if tail.strip() and flushed else b"" + return self._emit(tail + separator + flushed) + + def _emit(self, frames: bytes) -> tuple[bytes | str, ...]: + if not frames: + return () + return (frames.decode("utf-8") if self._as_text else frames,) + + async def _restore_event(self, block: bytes) -> bytes: + """Rewrites one SSE event, or returns it untouched if it carries nothing to restore.""" + try: + lines: Final = block.decode("utf-8").split("\n") + except UnicodeDecodeError: + return block + data_lines: Final = tuple(index for index, line in enumerate(lines) if line.startswith("data:")) + if len(data_lines) != 1: + return block + line: Final = lines[data_lines[0]] + try: + parsed: Final[object] = json.loads(line[len("data:") :]) + except ValueError: + return block + event: Final = as_object(parsed) + if event is None: + return block + kind: Final = event.get("type") + index: Final = event.get("index") + if kind == "content_block_stop" and isinstance(index, int): + return await self._flush(index) + block + if kind == "message_stop": + return await self._flush_all() + block + if kind != "content_block_delta" or not await self._restore_delta(event): + return block + ending: Final = "\r" if line.endswith("\r") else "" + rewritten: Final = ( + *lines[: data_lines[0]], + f"data: {json.dumps(event, ensure_ascii=False)}{ending}", + *lines[data_lines[0] + 1 :], + ) + return "\n".join(rewritten).encode("utf-8") + + async def _restore_delta(self, event: MutableRequest) -> bool: + """Advances one block's window through this delta. False if it holds no text.""" + index: Final = event.get("index") + delta: Final = event.get("delta") + if not isinstance(index, int) or not isinstance(delta, dict): + return False + delta_type: Final = delta.get("type") + if not isinstance(delta_type, str): + return False + field: Final = ANTHROPIC_DELTA_FIELDS.get(delta_type) + text: Final = delta.get(field) if field is not None else None + if field is None or not isinstance(text, str) or not text: + return False + emitted, remaining = await self._step(text, self._carries.get(index, ""), False) + self._carries[index] = remaining + self._delta_types[index] = delta_type + delta[field] = emitted + return True + + async def _flush(self, index: int) -> bytes: + """One synthetic delta frame carrying whatever `index`'s window still holds.""" + carry: Final = self._carries.pop(index, "") + delta_type: Final = self._delta_types.pop(index, None) + field: Final = ANTHROPIC_DELTA_FIELDS.get(delta_type) if isinstance(delta_type, str) else None + if not carry or field is None: + return b"" + text, _ = await self._step("", carry, True) + if not text: + return b"" + event: Final[JsonBody] = { + "type": "content_block_delta", + "index": index, + "delta": {"type": delta_type, field: text}, + } + return f"event: content_block_delta\ndata: {json.dumps(event, ensure_ascii=False)}\n\n".encode() + + async def _flush_all(self) -> bytes: + flushed: Final = tuple([await self._flush(index) for index in tuple(self._carries)]) + return b"".join(flushed) + + +class ResponsesStreamRestorer: + """Restores a Responses API event stream. + + The event families are matched by shape rather than listed, so a text stream the + API adds later is restored by default instead of leaking a placeholder: + + - Any `*.delta` event whose `delta` is a string is a token stream (output_text, + refusal, function-call and MCP arguments, reasoning summaries, ...). Each gets its + own window, keyed by the family, the item id and the part index. + - Any `*.done` event closes the stream of the same family. Whatever its window still + holds goes out first, as a copy of that stream's last delta event -- so it carries + the stream's own ids, and repeats that event's `sequence_number`. Then every text + field on the done event is restored in full: its string fields other than + identifiers, plus any `part` or `item` it repeats. + - `response.completed` / `response.incomplete` repeat the whole reply, and are + restored the same way the non-streaming reply is. + """ + + def __init__(self, step: StreamStep, rehydrate: Rehydrate) -> None: + self._step: Final = step + self._rehydrate: Final = rehydrate + self._carries: Final[dict[ResponsesStreamKey, str]] = {} # mutable-ok: per-stream windows advanced in place. + self._last_deltas: Final[dict[ResponsesStreamKey, object]] = {} # mutable-ok: newest delta per stream. + + async def restore(self, event: object) -> tuple[object, ...]: + """The events to emit in place of `event`: any flush, then the event itself.""" + kind: Final = responses_event_type(event) + if kind is None: + return (event,) + if kind.endswith(".delta") and kind not in RESPONSES_BINARY_DELTAS: + await self._restore_delta(event, kind) + return (event,) + slots: Final[SlotSink] = [] + flushed: Final = await self._flush(responses_stream_key(event, kind)) if kind.endswith(".done") else () + if kind.endswith(".done"): + collect_event_text(event, slots) + part: Final = read_field(event, "part") + if part is not None: + collect_response_item({"content": [part]}, slots) + collect_response_item(read_field(event, "item"), slots) + elif kind in RESPONSES_TERMINAL_EVENTS: + for item in read_list(read_field(event, "response"), "output"): + collect_response_item(item, slots) + await rehydrate_slots(slots, self._rehydrate) + return (*flushed, event) + + async def finish(self) -> tuple[object, ...]: + """Flushes every stream the provider never closed, e.g. a truncated reply.""" + flushed: Final = tuple([await self._flush(key) for key in tuple(self._carries)]) + return tuple(itertools.chain.from_iterable(flushed)) + + async def _restore_delta(self, event: object, kind: str) -> None: + text: Final = read_field(event, "delta") + if not isinstance(text, str) or not text: + return + key: Final = responses_stream_key(event, kind) + emitted, remaining = await self._step(text, self._carries.get(key, ""), False) + self._carries[key] = remaining + self._last_deltas[key] = event + write_field(event, "delta", emitted) + + async def _flush(self, key: ResponsesStreamKey) -> tuple[object, ...]: + carry: Final = self._carries.pop(key, "") + template: Final = self._last_deltas.pop(key, None) + if not carry or template is None: + return () + text, _ = await self._step("", carry, True) + if not text: + return () + flush: Final = copy.deepcopy(template) + write_field(flush, "delta", text) + return (flush,) + + +def responses_stream_key(event: object, kind: str) -> ResponsesStreamKey: + """Identifies the delta stream an event belongs to, the same for its delta and done. + + The family is the event type without its `.delta` / `.done` suffix, so an output_text + stream and a refusal stream on the same part never share a window. + """ + family: Final = kind.rsplit(".", 1)[0] + part_index: Final = read_field(event, "content_index") + summary_index: Final = read_field(event, "summary_index") + return ( + family, + read_field(event, "item_id"), + read_field(event, "output_index"), + part_index if part_index is not None else summary_index, + ) + + +def collect_event_text(event: object, slots: SlotSink) -> None: + """Collects every top-level text field of a Responses API event, dict or model. + + Scan by default, with identifiers excluded, rather than a list of known fields: the + `.done` event of each stream family names its text differently (`text`, `refusal`, + `arguments`, ...), and a family added upstream would otherwise leak a placeholder. + """ + attributes: Final[object] = getattr(event, "__dict__", None) + fields: Final = as_object(event) or as_object(attributes) + if fields is None: + return + for name, value in tuple(fields.items()): + if name in RESPONSES_STRUCTURAL_FIELDS or name.endswith("_id"): + continue + if isinstance(value, str) and value: + slots.append((value, functools.partial(write_field, event, name))) diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 37c1829def4..0c4ea6b29b5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -6,11 +6,14 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint 3. Implements a way to call /applyGuardrail endpoint for `/chat/completions` + `/v1/messages` requests on async_post_call_streaming_iterator_hook """ +import asyncio +import contextlib import copy import json from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Callable, Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final, Protocol +from typing import TYPE_CHECKING, Any, Final, Protocol, TypeAlias, cast +import anyio from fastapi import HTTPException from litellm._logging import verbose_proxy_logger @@ -19,6 +22,7 @@ from litellm.cost_calculator import _infer_call_type from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route +from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings from litellm.proxy._types import UserAPIKeyAuth from litellm.types.guardrails import GuardrailEventHooks @@ -28,6 +32,7 @@ from litellm.types.utils import ( CallTypesLiteral, Delta, ModelResponseStream, + StandardLoggingGuardrailInformation, StreamingChoices, ) @@ -45,6 +50,8 @@ A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message) GUARDRAIL_NAME: Final = "unified_llm_guardrails" +_RequestData: TypeAlias = dict[str, object] + class _EndpointTranslation(Protocol): @property @@ -59,6 +66,9 @@ class _EndpointTranslation(Protocol): @property def get_streaming_scan_key(self) -> "Callable[[Sequence[object]], StreamingScanKey | None]": ... + @property + def released_stream_as_ended(self) -> "Callable[[Sequence[object]], tuple[object, ...]]": ... + @property def build_block_sse_chunks(self) -> "Callable[..., Sequence[bytes] | None]": ... @@ -109,6 +119,18 @@ def _held_choices(held_chars_per_choice: Mapping[int, int]) -> frozenset[int]: return frozenset(idx for idx, held in held_chars_per_choice.items() if held > 0) +def _recorded_guardrail_information(request_data: _RequestData) -> tuple[StandardLoggingGuardrailInformation, ...]: + _metadata_key, metadata_bucket = get_or_create_metadata_bucket(request_data) + entries: Final = metadata_bucket.get("standard_logging_guardrail_information") + if not isinstance(entries, list): + return () + return tuple( + cast( # cast-ok: only the guardrail logging helpers write this metadata key + "list[StandardLoggingGuardrailInformation]", entries + ) + ) + + def _is_redundant_scan(scan_key: "StreamingScanKey | None", last_scan_key: "StreamingScanKey | None") -> bool: if scan_key is None: return False @@ -601,6 +623,7 @@ class UnifiedLLMGuardrails(CustomLogger): finish_reason_per_choice: dict[int, str | None], held_chars_per_choice: dict[int, int], is_final: bool, + terminated: asyncio.Event, ) -> AsyncGenerator[object, None]: """Run one guardrail processing round and emit the resulting diff chunk. @@ -632,6 +655,7 @@ class UnifiedLLMGuardrails(CustomLogger): is_final=is_final, ) except ModifyResponseException as e: + terminated.set() if e.original_response is None: e.original_response = responses_so_far async for block_chunk in self.handle_streaming_block( @@ -643,6 +667,7 @@ class UnifiedLLMGuardrails(CustomLogger): yield block_chunk raise _StreamTerminated() except HTTPException as e: + terminated.set() async for error_item in self.emit_streaming_http_error( e, call_type, @@ -664,7 +689,7 @@ class UnifiedLLMGuardrails(CustomLogger): *, guardrail_to_apply: CustomGuardrail, response: AsyncIterable[object], - request_data: dict, + request_data: _RequestData, user_api_key_dict: UserAPIKeyAuth, call_type: str, sampling_rate: int, @@ -687,6 +712,7 @@ class UnifiedLLMGuardrails(CustomLogger): held_chars_per_choice: Final[dict[int, int]] = {} chunk_counter = 0 last_chunk: object | None = None + terminated: Final = asyncio.Event() def _round(reference_chunk: object, is_final: bool) -> AsyncGenerator[object, None]: return self._emit_transform_round( @@ -702,10 +728,13 @@ class UnifiedLLMGuardrails(CustomLogger): finish_reason_per_choice=finish_reason_per_choice, held_chars_per_choice=held_chars_per_choice, is_final=is_final, + terminated=terminated, ) saw_tool_calls = False saw_text_content = False + tool_calls_released = False # rebind-ok: set once a raw tool call reaches the client unscanned + end_of_stream_inspection_started = False # rebind-ok: set once the end-of-stream inspection owns the verdict try: async for item in response: @@ -742,6 +771,7 @@ class UnifiedLLMGuardrails(CustomLogger): held_choices=_held_choices(held_chars_per_choice), ) responses_yielded.append(tool_only) + tool_calls_released = True yield tool_only continue @@ -781,6 +811,7 @@ class UnifiedLLMGuardrails(CustomLogger): # ``stream_transform_underflow`` 400 from mismatched prefixes. A shallow # list copy wouldn't help โ€” the mutation is on the chunk objects # themselves โ€” so we deepcopy. + end_of_stream_inspection_started = True if saw_tool_calls: async for out in self._inspect_full_response_for_block( endpoint_translation=endpoint_translation, @@ -801,6 +832,37 @@ class UnifiedLLMGuardrails(CustomLogger): yield out except _StreamTerminated: return + except (GeneratorExit, asyncio.CancelledError): + await self._scan_uninspected_tool_calls_after_disconnect( + uninspected=tool_calls_released and not end_of_stream_inspection_started and not terminated.is_set(), + endpoint_translation=endpoint_translation, + responses_released=responses_yielded, + guardrail_to_apply=guardrail_to_apply, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + raise + + @staticmethod + async def _scan_uninspected_tool_calls_after_disconnect( + *, + uninspected: bool, + endpoint_translation: _EndpointTranslation, + responses_released: Sequence[object], + guardrail_to_apply: CustomGuardrail, + user_api_key_dict: UserAPIKeyAuth, + request_data: _RequestData, + ) -> None: + if not uninspected: + return + await UnifiedLLMGuardrails._scan_released_stream_after_disconnect( + endpoint_translation=endpoint_translation, + responses_released=responses_released, + last_scan_key=None, + guardrail_to_apply=guardrail_to_apply, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) async def _emit_stream_tail( self, @@ -841,14 +903,15 @@ class UnifiedLLMGuardrails(CustomLogger): from litellm.integrations.custom_guardrail import ModifyResponseException try: - await endpoint_translation.process_output_streaming_response( - responses_so_far=responses_so_far, - guardrail_to_apply=guardrail_to_apply, - litellm_logging_obj=request_data.get("litellm_logging_obj"), - user_api_key_dict=user_api_key_dict, - request_data=request_data, - stream_transform_sink=None, - ) + with anyio.CancelScope(shield=bool(responses_yielded)): + await endpoint_translation.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + stream_transform_sink=None, + ) except ModifyResponseException as e: if e.original_response is None: e.original_response = responses_so_far @@ -968,6 +1031,45 @@ class UnifiedLLMGuardrails(CustomLogger): config_value: Final = config.get(name, attribute_value) if isinstance(config, dict) else attribute_value return self.optional_params.get(name, config_value) + @staticmethod + async def _scan_released_stream_after_disconnect( + *, + endpoint_translation: _EndpointTranslation, + responses_released: Sequence[object], + last_scan_key: "StreamingScanKey | None", + guardrail_to_apply: CustomGuardrail, + user_api_key_dict: UserAPIKeyAuth, + request_data: _RequestData, + ) -> None: + scanned: Final = endpoint_translation.released_stream_as_ended(copy.deepcopy(tuple(responses_released))) + if _is_redundant_scan(endpoint_translation.get_streaming_scan_key(scanned), last_scan_key): + return + recorded_before: Final = len(_recorded_guardrail_information(request_data)) + with anyio.CancelScope(shield=True): + try: + await endpoint_translation.process_output_streaming_response( + responses_so_far=scanned, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + except Exception as e: # noqa: BLE001 # the client is gone, so the verdict can only be recorded + verbose_proxy_logger.warning( + "UnifiedLLMGuardrails: %s scanned a stream the client disconnected from and raised %s", + guardrail_to_apply.guardrail_name, + type(e).__name__, + ) + recorded_during_scan: Final = _recorded_guardrail_information(request_data)[recorded_before:] + if any(entry.get("guardrail_status") != "success" for entry in recorded_during_scan): + return + guardrail_to_apply.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=e, + request_data=request_data, + guardrail_status="guardrail_failed_to_respond", + event_type=GuardrailEventHooks.post_call, + ) + async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, @@ -995,6 +1097,7 @@ class UnifiedLLMGuardrails(CustomLogger): if guardrail_to_apply is None: guardrail_to_apply = request_data.pop("guardrail_to_apply", None) + typed_request_data: Final[_RequestData] = request_data def _streaming_flag(name: str, default: object) -> Any: return self.resolve_streaming_flag(guardrail_to_apply, name, default) @@ -1061,17 +1164,20 @@ class UnifiedLLMGuardrails(CustomLogger): mappings=mappings, ) if transform_call_type is not None: - async for transformed_item in self._run_incremental_transform_stream( - guardrail_to_apply=guardrail_to_apply, - response=response, - request_data=request_data, - user_api_key_dict=user_api_key_dict, - call_type=transform_call_type, - sampling_rate=sampling_rate, - end_of_stream_only=end_of_stream_only, - mappings=mappings, - ): - yield transformed_item + async with contextlib.aclosing( + self._run_incremental_transform_stream( + guardrail_to_apply=guardrail_to_apply, + response=response, + request_data=typed_request_data, + user_api_key_dict=user_api_key_dict, + call_type=transform_call_type, + sampling_rate=sampling_rate, + end_of_stream_only=end_of_stream_only, + mappings=mappings, + ) + ) as transformed: + async for transformed_item in transformed: + yield transformed_item return verbose_proxy_logger.warning( "UnifiedLLMGuardrails: streaming_transform_mode=incremental_diff is only supported " @@ -1093,217 +1199,240 @@ class UnifiedLLMGuardrails(CustomLogger): chunks_yielded = False last_scan_key: StreamingScanKey | None = None # rebind-ok: replaced after every scan round tool_calls_in_flight = False # rebind-ok: tracks the latest scan key's unscanned tool calls + verdict_settled = False # rebind-ok: set once the end-of-stream scan or a block owns the verdict - async for item in response: - chunk_counter += 1 - responses_so_far.append(item) + try: + async for item in response: + chunk_counter += 1 + responses_so_far.append(item) - # Infer call type from first chunk if not already done - if call_type is None and user_api_key_dict.request_route is not None: - call_types = get_call_types_for_route(user_api_key_dict.request_route) - if call_types is not None: - call_type = call_types[0].value + # Infer call type from first chunk if not already done + if call_type is None and user_api_key_dict.request_route is not None: + call_types = get_call_types_for_route(user_api_key_dict.request_route) + if call_types is not None: + call_type = call_types[0].value - if call_type is None: - call_type = _infer_call_type(call_type=None, completion_response=item) + if call_type is None: + call_type = _infer_call_type(call_type=None, completion_response=item) - # If call type not supported, just pass through all chunks - if call_type is None or CallTypes(call_type) not in mappings: - yield item - async for remaining_item in response: - yield remaining_item - return + # If call type not supported, just pass through all chunks + if call_type is None or CallTypes(call_type) not in mappings: + yield item + async for remaining_item in response: + yield remaining_item + return - # If end_of_stream_only mode, yield chunks without processing. - # When buffering, withhold them instead -- they are released (or - # replaced by the block message) only after end-of-stream - # moderation runs below. - if end_of_stream_only: - if not buffer_until_moderated: - endpoint_translation = mappings[CallTypes(call_type)]() - stream_has_ended = hasattr( - endpoint_translation, "_check_streaming_has_ended" - ) and endpoint_translation._check_streaming_has_ended(responses_so_far) - if pending_end_of_stream_items or stream_has_ended: - pending_end_of_stream_items.append(item) - else: - chunks_yielded = True - responses_yielded.append(item) - yield item - else: - withheld_items.append(item) - continue - - # Process chunk based on sampling rate - if buffer_until_moderated: - withheld_items.append(item) - if chunk_counter % sampling_rate == 0: - endpoint_translation = mappings[CallTypes(call_type)]() - scan_key = endpoint_translation.get_streaming_scan_key(responses_so_far) - if scan_key is not None: - tool_calls_in_flight = scan_key.tool_calls_in_flight - hold_window = buffer_until_moderated and (scan_key is None or tool_calls_in_flight) - if _is_redundant_scan(scan_key, last_scan_key): - verbose_proxy_logger.debug( - "Skipping streaming chunk %s for guardrail %s: nothing new to scan since the last round", - chunk_counter, - guardrail_to_apply.guardrail_name, - ) - if buffer_until_moderated: - if hold_window: - continue - for withheld_item in withheld_items: + # If end_of_stream_only mode, yield chunks without processing. + # When buffering, withhold them instead -- they are released (or + # replaced by the block message) only after end-of-stream + # moderation runs below. + if end_of_stream_only: + if not buffer_until_moderated: + endpoint_translation = mappings[CallTypes(call_type)]() + stream_has_ended = hasattr( + endpoint_translation, "_check_streaming_has_ended" + ) and endpoint_translation._check_streaming_has_ended(responses_so_far) + if pending_end_of_stream_items or stream_has_ended: + pending_end_of_stream_items.append(item) + else: chunks_yielded = True - responses_yielded.append(withheld_item) - yield withheld_item - withheld_items.clear() + responses_yielded.append(item) + yield item else: - chunks_yielded = True - responses_yielded.append(item) - yield item + withheld_items.append(item) continue + # Process chunk based on sampling rate + if buffer_until_moderated: + withheld_items.append(item) + if chunk_counter % sampling_rate == 0: + endpoint_translation = mappings[CallTypes(call_type)]() + scan_key = endpoint_translation.get_streaming_scan_key(responses_so_far) + if scan_key is not None: + tool_calls_in_flight = scan_key.tool_calls_in_flight + hold_window = buffer_until_moderated and (scan_key is None or tool_calls_in_flight) + if _is_redundant_scan(scan_key, last_scan_key): + verbose_proxy_logger.debug( + "Skipping streaming chunk %s for guardrail %s: nothing new to scan since the last round", + chunk_counter, + guardrail_to_apply.guardrail_name, + ) + if buffer_until_moderated: + if hold_window: + continue + for withheld_item in withheld_items: + chunks_yielded = True + responses_yielded.append(withheld_item) + yield withheld_item + withheld_items.clear() + else: + chunks_yielded = True + responses_yielded.append(item) + yield item + continue + + verbose_proxy_logger.debug( + "Processing streaming chunk %s (sampling_rate=%s) with guardrail %s", + chunk_counter, + sampling_rate, + guardrail_to_apply.guardrail_name, + ) + + original_items = ( + tuple(copy.deepcopy(withheld_items)) if buffer_until_moderated else (copy.deepcopy(item),) + ) + + try: + await endpoint_translation.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + except ModifyResponseException as e: + verdict_settled = True + if e.original_response is None: + e.original_response = responses_so_far + # Guardrail blocked the response mid-stream. Emit a clean + # terminating SSE sequence delivering the block message + # instead of letting the exception propagate into a bare + # `data: {"error": ...}` blob (which truncates the stream). + # Chunks have already been forwarded here, so the block + # continues the in-progress message (stream_started=True). + # The current chunk was appended to responses_so_far but not + # yet yielded, so exclude it: the continuation must reflect + # only what the client has actually received. + async for block_chunk in self.handle_streaming_block( + e, + endpoint_translation, + stream_started=chunks_yielded, + responses_so_far=responses_yielded, + ): + yield block_chunk + return + except HTTPException as e: + verdict_settled = True + # Response already started (we already yielded chunks); cannot send 400. + async for error_item in self.emit_streaming_http_error( + e, + call_type, + responses_so_far, + request_data, + endpoint_translation=endpoint_translation, + stream_started=chunks_yielded, + responses_yielded=responses_yielded, + ): + yield error_item + return + if scan_key is not None: + last_scan_key = scan_key + if hold_window: + verbose_proxy_logger.debug( + "Holding %s buffered chunks for guardrail %s: this round could not scan the whole window", + len(withheld_items), + guardrail_to_apply.guardrail_name, + ) + withheld_items[:] = original_items + continue + for original_item in original_items: + chunks_yielded = True + responses_yielded.append(original_item) + yield original_item + withheld_items.clear() + else: + if not buffer_until_moderated: + chunks_yielded = True + responses_yielded.append(item) + yield item + + # Stream has ended - do final processing with all collected chunks + if call_type is not None and CallTypes(call_type) in mappings: verbose_proxy_logger.debug( - "Processing streaming chunk %s (sampling_rate=%s) with guardrail %s", - chunk_counter, - sampling_rate, + "Processing final streaming response with all %s chunks for guardrail %s", + len(responses_so_far), guardrail_to_apply.guardrail_name, ) - original_items = ( - tuple(copy.deepcopy(withheld_items)) if buffer_until_moderated else (copy.deepcopy(item),) + endpoint_translation = mappings[CallTypes(call_type)]() + + buffered_items: Final = ( + tuple(copy.deepcopy(withheld_items)) + if buffer_until_moderated and release_on_scan and not end_of_stream_only + else tuple(withheld_items) + if buffer_until_moderated + else None ) + end_scan_key: Final = endpoint_translation.get_streaming_scan_key(responses_so_far) + verdict_settled = True + if _is_redundant_scan(end_scan_key, last_scan_key): + verbose_proxy_logger.debug( + "Skipping end-of-stream scan for guardrail %s: the last sampled round already scanned it all", + guardrail_to_apply.guardrail_name, + ) + for buffered_item in buffered_items or (): + yield buffered_item + for pending_item in pending_end_of_stream_items: + responses_yielded.append(pending_item) + yield pending_item + return try: - await endpoint_translation.process_output_streaming_response( - responses_so_far=responses_so_far, - guardrail_to_apply=guardrail_to_apply, - litellm_logging_obj=request_data.get("litellm_logging_obj"), - user_api_key_dict=user_api_key_dict, - request_data=request_data, - ) + with anyio.CancelScope(shield=chunks_yielded): + await endpoint_translation.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + # Moderation passed: release the withheld original chunks. + if buffered_items is not None: + for buffered_item in buffered_items: + yield buffered_item + for pending_item in pending_end_of_stream_items: + responses_yielded.append(pending_item) + yield pending_item except ModifyResponseException as e: if e.original_response is None: e.original_response = responses_so_far - # Guardrail blocked the response mid-stream. Emit a clean - # terminating SSE sequence delivering the block message - # instead of letting the exception propagate into a bare - # `data: {"error": ...}` blob (which truncates the stream). - # Chunks have already been forwarded here, so the block - # continues the in-progress message (stream_started=True). - # The current chunk was appended to responses_so_far but not - # yet yielded, so exclude it: the continuation must reflect - # only what the client has actually received. + # Block detected during end-of-stream processing. Emit a clean + # terminating SSE sequence with the block message rather than + # propagating into a bare error blob that truncates the stream. + # The withheld original chunks are never released. async for block_chunk in self.handle_streaming_block( e, endpoint_translation, - stream_started=chunks_yielded, + stream_started=bool(responses_yielded), responses_so_far=responses_yielded, ): yield block_chunk return except HTTPException as e: - # Response already started (we already yielded chunks); cannot send 400. async for error_item in self.emit_streaming_http_error( e, call_type, responses_so_far, request_data, endpoint_translation=endpoint_translation, - stream_started=chunks_yielded, + stream_started=bool(responses_yielded), responses_yielded=responses_yielded, ): yield error_item - return - if scan_key is not None: - last_scan_key = scan_key - if hold_window: - verbose_proxy_logger.debug( - "Holding %s buffered chunks for guardrail %s: this round could not scan the whole window", - len(withheld_items), - guardrail_to_apply.guardrail_name, - ) - withheld_items[:] = original_items - continue - for original_item in original_items: - chunks_yielded = True - responses_yielded.append(original_item) - yield original_item - withheld_items.clear() - else: - if not buffer_until_moderated: - chunks_yielded = True - responses_yielded.append(item) - yield item - - # Stream has ended - do final processing with all collected chunks - if call_type is not None and CallTypes(call_type) in mappings: - verbose_proxy_logger.debug( - "Processing final streaming response with all %s chunks for guardrail %s", - len(responses_so_far), - guardrail_to_apply.guardrail_name, - ) - - endpoint_translation = mappings[CallTypes(call_type)]() - - buffered_items: Final = ( - tuple(copy.deepcopy(withheld_items)) - if buffer_until_moderated and release_on_scan and not end_of_stream_only - else tuple(withheld_items) - if buffer_until_moderated - else None - ) - end_scan_key: Final = endpoint_translation.get_streaming_scan_key(responses_so_far) - if _is_redundant_scan(end_scan_key, last_scan_key): - verbose_proxy_logger.debug( - "Skipping end-of-stream scan for guardrail %s: the last sampled round already scanned it all", - guardrail_to_apply.guardrail_name, - ) - for buffered_item in buffered_items or (): - yield buffered_item - for pending_item in pending_end_of_stream_items: - responses_yielded.append(pending_item) - yield pending_item - return - - try: - await endpoint_translation.process_output_streaming_response( - responses_so_far=responses_so_far, + except (GeneratorExit, asyncio.CancelledError): + translation_class: Final = None if call_type is None else mappings.get(CallTypes(call_type)) + if ( + chunks_yielded + and not verdict_settled + and translation_class is not None + and isinstance(guardrail_to_apply, CustomGuardrail) + ): + await self._scan_released_stream_after_disconnect( + endpoint_translation=translation_class(), + responses_released=responses_yielded, + last_scan_key=last_scan_key, guardrail_to_apply=guardrail_to_apply, - litellm_logging_obj=request_data.get("litellm_logging_obj"), user_api_key_dict=user_api_key_dict, - request_data=request_data, + request_data=typed_request_data, ) - # Moderation passed: release the withheld original chunks. - if buffered_items is not None: - for buffered_item in buffered_items: - yield buffered_item - for pending_item in pending_end_of_stream_items: - responses_yielded.append(pending_item) - yield pending_item - except ModifyResponseException as e: - if e.original_response is None: - e.original_response = responses_so_far - # Block detected during end-of-stream processing. Emit a clean - # terminating SSE sequence with the block message rather than - # propagating into a bare error blob that truncates the stream. - # The withheld original chunks are never released. - async for block_chunk in self.handle_streaming_block( - e, - endpoint_translation, - stream_started=bool(responses_yielded), - responses_so_far=responses_yielded, - ): - yield block_chunk - return - except HTTPException as e: - async for error_item in self.emit_streaming_http_error( - e, - call_type, - responses_so_far, - request_data, - endpoint_translation=endpoint_translation, - stream_started=bool(responses_yielded), - responses_yielded=responses_yielded, - ): - yield error_item + raise diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 00f0ef1756f..2a0b528a50b 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -8,7 +8,7 @@ import time import traceback from collections.abc import Iterable, Mapping from datetime import datetime, timedelta, timezone -from typing import Any, Final, Literal, TypedDict, cast +from typing import Final, Literal, TypedDict import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, Response, status @@ -1431,7 +1431,7 @@ async def shared_health_check_status_endpoint( ) -def _read_license_data() -> dict[str, Any] | None: +def _read_license_data() -> EnterpriseLicenseData | None: from litellm.proxy.proxy_server import _license_check, premium_user_data license_data: EnterpriseLicenseData | None = premium_user_data or _license_check.airgapped_license_data @@ -1453,10 +1453,10 @@ def _read_license_data() -> dict[str, Any] | None: if license_data is None: return None - return cast(dict[str, Any], license_data) + return license_data -def _read_allowed_features(license_data: dict[str, Any]) -> list: +def _read_allowed_features(license_data: Mapping[str, object]) -> list: raw_allowed_features: Final = license_data.get("allowed_features") if isinstance(raw_allowed_features, list): return list(raw_allowed_features) @@ -1707,7 +1707,7 @@ def _show_env_credential_login_warning() -> bool: async def _get_health_readiness_details( response: Response | None = None, -) -> dict[str, Any]: +) -> dict[str, object]: """ Detailed health payload for authenticated diagnostics. """ @@ -1726,7 +1726,7 @@ async def _get_health_readiness_details( success_callback_names = litellm.success_callback # check Cache - cache_type: Any = None + cache_type: object = None if litellm.cache is not None: from litellm.caching.caching import RedisSemanticCache @@ -1735,7 +1735,7 @@ async def _get_health_readiness_details( if isinstance(litellm.cache.cache, RedisSemanticCache): # ping the cache # TODO: @ishaan-jaff - we should probably not ping the cache on every /health/readiness check - index_info: Any + index_info: object try: index_info = await litellm.cache.cache._index_info() except Exception as e: diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 33fb069afbd..5db0fce7d92 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -54,6 +54,8 @@ from litellm.repositories.autorouter_session_repository import AutoRouterSession from litellm.repositories.base_repository import SupportsModelDump from litellm.repositories.daily_activity_sql import build_where_clause from litellm.repositories.team_repository import TeamRepository +from litellm.repositories.user_repository import UserRepository +from litellm.repositories.verification_token_repository import VerificationTokenRepository from litellm.router_strategy.complexity_router import ComplexityRouter from litellm.router_utils.auto_router_model_naming import ( StrategyRouterDependencyRole, @@ -187,15 +189,15 @@ def _team_table(prisma_client: "PrismaClient") -> _TeamTable: def _verification_tokens(prisma_client: "PrismaClient") -> _VerificationTokenTable: - return prisma_client.db.litellm_verificationtoken + return VerificationTokenRepository(prisma_client).table def _team_rows(prisma_client: "PrismaClient") -> _TeamRowsTable: - return prisma_client.db.litellm_teamtable + return TeamRepository(prisma_client).table def _user_rows(prisma_client: "PrismaClient") -> _UserRowsTable: - return prisma_client.db.litellm_usertable + return UserRepository(prisma_client).table def _shadow_eval_jobs(prisma_client: "PrismaClient") -> _ShadowEvalJobTable: @@ -571,13 +573,16 @@ async def preview_auto_router_routing( llm_router=llm_router, ) - complexity_router: Final = ComplexityRouter( - model_name=resolved.router_name, - litellm_router_instance=llm_router, - complexity_router_config=resolved.complexity_router_config.model_dump(exclude_none=True), - default_model=resolved.default_model, - derive_savings_baseline=False, - ) + try: + complexity_router: Final = ComplexityRouter( + model_name=resolved.router_name, + litellm_router_instance=llm_router, + complexity_router_config=resolved.complexity_router_config.model_dump(exclude_none=True), + default_model=resolved.complexity_router_config.resolve_default_model(resolved.default_model), + derive_savings_baseline=False, + ) + except ValueError as e: + raise HTTPException(status_code=400, detail={"error": f"Could not route this prompt: {e}"}) from e request_kwargs: Final = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata( data=request_data, diff --git a/litellm/proxy/management_endpoints/callback_management_endpoints.py b/litellm/proxy/management_endpoints/callback_management_endpoints.py index 4f9f46af08f..d910b475b57 100644 --- a/litellm/proxy/management_endpoints/callback_management_endpoints.py +++ b/litellm/proxy/management_endpoints/callback_management_endpoints.py @@ -7,12 +7,15 @@ import os from typing import Final from fastapi import APIRouter, Depends +from pydantic import TypeAdapter from litellm.litellm_core_utils.logging_callback_manager import CallbacksByType from litellm.proxy.auth.user_api_key_auth import user_api_key_auth router: Final = APIRouter() +_DECODED_JSON: Final = TypeAdapter(object) + @router.get( "/callbacks/list", @@ -51,6 +54,6 @@ async def get_callback_configs(): ) with open(config_path, "r") as f: - configs: Final = json.load(f) + configs: Final = _DECODED_JSON.validate_python(json.load(f)) return configs diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 323b9e434a8..037b9082cd4 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -355,9 +355,12 @@ _KEY_METADATA_REQUEST_FIELDS: Final = frozenset( ) +_DECODED_JSON: Final = TypeAdapter(object) + + def _decode_json_string_column(column: str, value: object) -> object: if column in _KEY_UPDATE_JSON_STRING_COLUMNS and isinstance(value, str): - return json.loads(value) + return _DECODED_JSON.validate_python(json.loads(value)) return value diff --git a/litellm/proxy/management_endpoints/management_v1/budgets.py b/litellm/proxy/management_endpoints/management_v1/budgets.py index 106dbfaf7b7..f1760b857c5 100644 --- a/litellm/proxy/management_endpoints/management_v1/budgets.py +++ b/litellm/proxy/management_endpoints/management_v1/budgets.py @@ -85,8 +85,7 @@ class PrismaBudgetListExecutor: async def count(self, where: tuple[Predicate, ...]) -> int: clauses, params = where_sql(where) sql: Final = f"SELECT COUNT(*) AS count FROM {BUDGET_TABLE}" + (f" WHERE {clauses}" if clauses else "") - rows: Final = await self.prisma_client.db.query_raw(sql, *params) - counted: Final = _ROW_COUNTS.validate_python(rows) + counted: Final = _ROW_COUNTS.validate_python(await self.prisma_client.db.query_raw(sql, *params)) return counted[0].count if counted else 0 async def find_many(self, plan: QueryPlan) -> Sequence[BudgetListItem]: @@ -97,8 +96,7 @@ class PrismaBudgetListExecutor: + f" ORDER BY {order_by_sql(plan.order)}" + f" LIMIT ${len(params) + 1} OFFSET ${len(params) + 2}" ) - rows: Final = await self.prisma_client.db.query_raw(sql, *params, plan.take, plan.skip) - return _BUDGET_ROWS.validate_python(rows) + return _BUDGET_ROWS.validate_python(await self.prisma_client.db.query_raw(sql, *params, plan.take, plan.skip)) def _serialize(row: BudgetListItem) -> BudgetListItem: diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index e1504172d9e..e5bea71ac7c 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -157,13 +157,16 @@ if MCP_AVAILABLE: get_user_env_vars, get_user_env_vars_bulk, get_user_oauth_credential, + is_resubmitted_oauth_client, list_server_user_credentials, list_user_oauth_credentials, mcp_oauth_token_identity, merge_user_env_vars, + oauth_credentials_for_upstream_edit, purge_user_oauth_credentials_for_server, reject_mcp_server, set_mcp_server_pinned_tools, + stale_mcp_auth_fields, store_user_credential, store_user_oauth_credential, update_mcp_server, @@ -530,6 +533,8 @@ if MCP_AVAILABLE: except Exception as e: verbose_proxy_logger.debug("Failed to write temporary MCP server to Redis cache: %s", e) + _CACHED_VALUE: Final = TypeAdapter(object) + @with_service_target(MCP_SERVERS_TARGET) async def _get_temporary_mcp_server_from_redis( server_id: str, @@ -547,8 +552,8 @@ if MCP_AVAILABLE: return None try: - cached_server: Final = await cache_backend.async_get_cache( - key=f"{TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX}:{server_id}" + cached_server: Final = _CACHED_VALUE.validate_python( + await cache_backend.async_get_cache(key=f"{TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX}:{server_id}") ) except Exception as e: verbose_proxy_logger.debug("Failed reading temporary MCP server from Redis cache: %s", e) @@ -870,6 +875,9 @@ if MCP_AVAILABLE: ("authentication_token", "auth_value"), ("client_id", "client_id"), ("client_secret", "client_secret"), + ("token_endpoint_auth_method", "token_endpoint_auth_method"), + ("dcr_issuer", "dcr_issuer"), + ("dcr_server_url", "dcr_server_url"), ("scopes", "scopes"), ("aws_access_key_id", "aws_access_key_id"), ("aws_secret_access_key", "aws_secret_access_key"), @@ -890,30 +898,73 @@ if MCP_AVAILABLE: value for key, value in as_dict.items() if key not in MCP_ADMIN_CONFIG_CREDENTIAL_KEYS and key != "scopes" ) + def _oauth_session_changes_upstream( + payload: NewMCPServerRequest, existing: MCPServer | LiteLLM_MCPServerTable + ) -> bool: + return existing.auth_type == MCPAuth.oauth2 and ( + payload.url != existing.url + or ("issuer" in payload.model_fields_set and (payload.issuer or None) != existing.issuer) + or (payload.auth_type is not None and payload.auth_type != existing.auth_type) + ) + def _inherit_credentials_from_existing_server( payload: NewMCPServerRequest, ) -> NewMCPServerRequest: - if not payload.server_id or _has_non_admin_config_credentials(payload.credentials): + if not payload.server_id: return payload existing_server: Final = global_mcp_server_manager.get_mcp_server_by_id(payload.server_id) if existing_server is None: return payload - - inherited_credentials: dict[str, object] = { + upstream_changed: Final = _oauth_session_changes_upstream(payload, existing_server) + issuer_changed: Final = ( + "issuer" in payload.model_fields_set and (payload.issuer or None) != existing_server.issuer + ) + cleared: Final = ( + stale_mcp_auth_fields( + payload.model_dump(exclude_unset=True), + lambda field: ( + getattr(existing_server, f"configured_{field}", None) or getattr(existing_server, field, None) + ), + ) + if upstream_changed + else {} + ) + resolved_payload: Final = ( + payload.model_copy(update={**cleared, "oauth2_flow": payload.oauth2_flow}) if upstream_changed else payload + ) + supplied: Final = dict(resolved_payload.credentials or {}) + existing_credentials: Final[dict[str, object]] = { credential_key: value for server_attr, credential_key in _INHERITED_CREDENTIAL_FIELDS if (value := getattr(existing_server, server_attr, None)) } + resubmitted_client: Final = is_resubmitted_oauth_client(supplied, existing_credentials) + if _has_non_admin_config_credentials(resolved_payload.credentials) and not ( + upstream_changed and resubmitted_client + ): + return resolved_payload + + bound_credentials: Final = ( + oauth_credentials_for_upstream_edit( + {**existing_credentials, **supplied}, + existing_server.issuer, + existing_server.url, + issuer_changed=issuer_changed + or (resolved_payload.auth_type or existing_server.auth_type) != existing_server.auth_type, + ) + if upstream_changed + else existing_credentials + ) # The gate above guarantees anything still supplied is admin config, which the admin just # typed, so it wins over the stored value. - inherited_credentials = {**inherited_credentials, **dict(payload.credentials or {})} + inherited_credentials: Final = bound_credentials if upstream_changed else {**bound_credentials, **supplied} if not inherited_credentials: - return payload + return resolved_payload.model_copy(update={"credentials": {}}) if upstream_changed else resolved_payload try: - return payload.model_copy(update={"credentials": inherited_credentials}) + return resolved_payload.model_copy(update={"credentials": inherited_credentials}) except AttributeError: pass @@ -937,15 +988,20 @@ if MCP_AVAILABLE: supplied: Final = payload.server_id if not supplied: return str(uuid.uuid4()) - if global_mcp_server_manager.get_mcp_server_by_id(supplied) is not None: - return supplied + registered: Final = global_mcp_server_manager.get_mcp_server_by_id(supplied) + if registered is not None: + return str(uuid.uuid4()) if _oauth_session_changes_upstream(payload, registered) else supplied prisma_client: Final = _get_prisma_client_or_none() if prisma_client is None: return supplied # A draft is another session's row, not a saved server, so re-supplying an id this # endpoint previously handed back must not let a later session adopt its configuration. existing: Final = await get_mcp_server(prisma_client, supplied) - if existing is None or existing.approval_status == MCPApprovalStatus.draft: + if ( + existing is None + or existing.approval_status == MCPApprovalStatus.draft + or _oauth_session_changes_upstream(payload, existing) + ): return str(uuid.uuid4()) return supplied diff --git a/litellm/proxy/management_endpoints/model_insights_endpoints.py b/litellm/proxy/management_endpoints/model_insights_endpoints.py index b787aac2d9f..53e8dfcc71e 100644 --- a/litellm/proxy/management_endpoints/model_insights_endpoints.py +++ b/litellm/proxy/management_endpoints/model_insights_endpoints.py @@ -115,7 +115,7 @@ def _deployment_filter(rows: list[_GroupedModel]) -> list[dict[str, str]]: def _daily_metric(row: _GroupedDaily) -> ModelInsightDailyMetric: - return ModelInsightDailyMetric(date=row.date, **_metric(row).model_dump()) + return ModelInsightDailyMetric.model_validate({**_metric(row).model_dump(), "date": row.date}) def _daily_total(row: _GroupedDate) -> ModelInsightDailyTotal: @@ -146,12 +146,14 @@ def _summarize_tasks(rows: list[_GroupedTask], metric: ModelInsightsMetric) -> l } grand: Final = sum(totals.values()) return [ - ModelInsightTaskSummary( - **(catalog.get(task) or _UNCATEGORIZED_TASK).model_copy(update={"task_type": task}).model_dump(), - value=value, - share=value / grand * 100 if grand else 0.0, - leader=leaders[task].model_group, - provider=leaders[task].custom_llm_provider, + ModelInsightTaskSummary.model_validate( + { + **(catalog.get(task) or _UNCATEGORIZED_TASK).model_copy(update={"task_type": task}).model_dump(), + "value": value, + "share": value / grand * 100 if grand else 0.0, + "leader": leaders[task].model_group, + "provider": leaders[task].custom_llm_provider, + } ) for task, value in sorted(totals.items(), key=lambda item: item[1], reverse=True) ] diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 11a0075ffd0..253013a2ebf 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -2547,8 +2547,8 @@ async def add_new_model( enforced=bool(general_settings.get(ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING, False)), ) - clean_model_info: Final = ModelInfo( - **without_server_derived_pricing(model_params.model_info.model_dump(exclude_none=True)) + clean_model_info: Final = ModelInfo.model_validate( + dict(without_server_derived_pricing(model_params.model_info.model_dump(exclude_none=True))) ) model_params.model_info = ( # rebind-ok: downstream team-model handling mutates this same object clean_model_info.model_copy(update=MappingProxyType({"member_auto_router": True})) @@ -3194,6 +3194,9 @@ def _deduplicate_litellm_router_models(models: list[dict]) -> list[dict]: return unique_models +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + + def model_info_as_mapping(model_info: object) -> Mapping[str, object] | None: """A DB row's model_info column arrives as a dict or as its JSON string depending on the query path, and every consumer needs the mapping. Single owner of that parse: @@ -3204,10 +3207,9 @@ def model_info_as_mapping(model_info: object) -> Mapping[str, object] | None: if not isinstance(model_info, str): return None try: - parsed: Final = json.loads(model_info) + return _JSON_OBJECT.validate_python(json.loads(model_info)) except (TypeError, ValueError): return None - return parsed if isinstance(parsed, Mapping) else None def _expects_liveness_on_this_pod(model_info: object) -> bool: diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index c8ae7af41db..26abe35b49a 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -27,7 +27,7 @@ from typing import ( import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, status -from pydantic import TypeAdapter +from pydantic import ConfigDict, TypeAdapter from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger @@ -306,7 +306,7 @@ async def _verify_org_access( ) -_STR_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object]) +_STR_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) _BUDGET_SETTABLE_FIELDS: Final = frozenset(LiteLLM_BudgetTable.model_fields.keys()) - {"budget_id"} _ORG_COLUMN_FIELDS: Final = frozenset({"organization_alias", "models"}) _ORG_METADATA_FIELDS: Final = tuple( @@ -720,7 +720,7 @@ async def update_organization( ) # Transform UI payload to expected format - raw_data: Final[dict[str, object]] = await request.json() + raw_data: Final = _STR_OBJECT_DICT_ADAPTER.validate_python(await request.json()) raw_data_with_flat_budget_fields: Final = handle_nested_budget_structure_in_organization_update_request(raw_data) # Create validated data model @@ -767,7 +767,9 @@ async def update_organization( # Merge metadata from existing organization with updated metadata if updated_organization_row_json.get("metadata") is not None: existing_metadata: Final = existing_organization_row.metadata or {} - updated_metadata: Final[dict[str, object]] = updated_organization_row_json.get("metadata", {}) + updated_metadata: Final = _STR_OBJECT_DICT_ADAPTER.validate_python( + updated_organization_row_json.get("metadata", {}) + ) merged_metadata: Final[Mapping[str, object]] = _update_dictionary( existing_dict=cast( # cast-ok: prisma de-serializes a Json column to the plain python dict it stores "dict[str, object]", existing_metadata @@ -786,12 +788,16 @@ async def update_organization( ) budget_fields: Final = { - k: v for k, v in data.model_dump().items() if k in _BUDGET_SETTABLE_FIELDS and k in data.model_fields_set + k: v + for k, v in _STR_OBJECT_DICT_ADAPTER.validate_python(data.model_dump()).items() + if k in _BUDGET_SETTABLE_FIELDS and k in data.model_fields_set } if budget_fields and existing_organization_row.budget_id: await update_budget( - budget_obj=BudgetNewRequest(budget_id=existing_organization_row.budget_id, **budget_fields), + budget_obj=BudgetNewRequest.model_validate( + {"budget_id": existing_organization_row.budget_id, **budget_fields} + ), user_api_key_dict=user_api_key_dict, ) diff --git a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py index 69356922ea1..c3c499eeec1 100644 --- a/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/policy_endpoints/endpoints.py @@ -12,12 +12,12 @@ All /policy management endpoints import copy import json import os -from collections.abc import AsyncGenerator, AsyncIterator +from collections.abc import AsyncGenerator, AsyncIterator, Mapping from typing import TYPE_CHECKING, Final, Literal, cast from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import Response, StreamingResponse -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field from typing_extensions import TypedDict import litellm @@ -449,6 +449,23 @@ async def validate_policy( return result +class _LoadedPolicy(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True, hide_input_in_errors=True) + + inherit: str | None = None + scope: PolicyScopeResponse = Field(default_factory=PolicyScopeResponse) + guardrails: PolicyGuardrailsResponse = Field(default_factory=PolicyGuardrailsResponse) + resolved_guardrails: tuple[str, ...] = () + inheritance_chain: tuple[str, ...] = () + + +class _LoadedPolicies(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True, hide_input_in_errors=True) + + policies: Mapping[str, _LoadedPolicy] = Field(default_factory=dict) + total_count: int = 0 + + @router.get( "/policy/list", tags=["policy management"], @@ -472,19 +489,19 @@ async def list_policies( """ from litellm.proxy.policy_engine.init_policies import get_policies_summary - summary: Final = get_policies_summary() + summary: Final = _LoadedPolicies.model_validate(get_policies_summary()) return PolicyListResponse( policies={ name: PolicySummaryItem( - inherit=data.get("inherit"), - scope=PolicyScopeResponse(**data.get("scope", {})), - guardrails=PolicyGuardrailsResponse(**data.get("guardrails", {})), - resolved_guardrails=data.get("resolved_guardrails", []), - inheritance_chain=data.get("inheritance_chain", []), + inherit=policy.inherit, + scope=policy.scope, + guardrails=policy.guardrails, + resolved_guardrails=list(policy.resolved_guardrails), + inheritance_chain=list(policy.inheritance_chain), ) - for name, data in summary.get("policies", {}).items() + for name, policy in summary.policies.items() }, - total_count=summary.get("total_count", 0), + total_count=summary.total_count, ) diff --git a/litellm/proxy/management_endpoints/router_settings_endpoints.py b/litellm/proxy/management_endpoints/router_settings_endpoints.py index a557e3a6082..d8ea96e6a46 100644 --- a/litellm/proxy/management_endpoints/router_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/router_settings_endpoints.py @@ -116,7 +116,7 @@ async def get_router_settings( config: Final = await proxy_config.get_config() router_settings_from_config: Final = config.get("router_settings", {}) - current_values: Final[dict[str, Any]] = {} + current_values: Final[dict[str, object]] = {} if llm_router is not None: # Router exposes routing groups as private `_routing_groups`; the # generic `hasattr` loop below would miss them. diff --git a/litellm/proxy/management_endpoints/tag_management_endpoints.py b/litellm/proxy/management_endpoints/tag_management_endpoints.py index b1094684389..9a1cf701142 100644 --- a/litellm/proxy/management_endpoints/tag_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tag_management_endpoints.py @@ -17,6 +17,7 @@ from datetime import datetime from typing import TYPE_CHECKING, Final, Protocol, TypedDict, overload from fastapi import APIRouter, Depends, HTTPException, Query +from pydantic import TypeAdapter from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth, user_api_key_has_admin_view @@ -58,6 +59,8 @@ if TYPE_CHECKING: router: Final = APIRouter() +_DECODED_JSON: Final = TypeAdapter(object) + class _TagRecord(Protocol): tag_name: str @@ -523,7 +526,7 @@ async def info_tag( model_info: object = {} if tag_record.model_info: if isinstance(tag_record.model_info, str): - model_info = json.loads(tag_record.model_info) + model_info = _DECODED_JSON.validate_python(json.loads(tag_record.model_info)) else: model_info = tag_record.model_info @@ -646,7 +649,7 @@ async def list_tags( model_info: object = {} if tag_record.model_info: if isinstance(tag_record.model_info, str): - model_info = json.loads(tag_record.model_info) + model_info = _DECODED_JSON.validate_python(json.loads(tag_record.model_info)) else: model_info = tag_record.model_info diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index fe976c861e5..e452afccef9 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -1721,7 +1721,7 @@ async def new_team( complete_team_data_dict = complete_team_data.model_dump(exclude_none=True) # Serialize router_settings to JSON (matching key creation pattern) - router_settings_value: Final = getattr(data, "router_settings", None) + router_settings_value: Final = data.router_settings router_settings_json: Final = ( safe_dumps(router_settings_value) if router_settings_value is not None else safe_dumps({}) ) diff --git a/litellm/proxy/openai_files_endpoints/file_content_streaming_handler.py b/litellm/proxy/openai_files_endpoints/file_content_streaming_handler.py index 2381a5cc2db..3d408accba3 100644 --- a/litellm/proxy/openai_files_endpoints/file_content_streaming_handler.py +++ b/litellm/proxy/openai_files_endpoints/file_content_streaming_handler.py @@ -18,11 +18,11 @@ class FileContentStreamingHandler: *, custom_llm_provider: str, file_id: str, - data: dict[str, Any], + data: dict[str, object], should_route: bool, original_file_id: str | None, - credentials: dict[str, Any] | None, - ) -> tuple[str, str, dict[str, Any]]: + credentials: dict[str, object] | None, + ) -> tuple[str, str, dict[str, object]]: """ Resolve the provider, file ID, and request payload to use for streaming. diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index f5e62da5962..dfec77a6c16 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -1652,7 +1652,8 @@ async def delete_file( def _as_file_list_page(response: object) -> object: if not isinstance(response, list): return response - return FileListPage(**build_list_page(_LISTED_FILES_ADAPTER.validate_python(response))) + page: Final = build_list_page(_LISTED_FILES_ADAPTER.validate_python(response)) + return FileListPage.model_validate(page) @router.get( diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py index 7cec3bac207..dfb731972ac 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py @@ -5,6 +5,7 @@ from datetime import datetime from typing import TYPE_CHECKING, Any, Final, cast import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._logging import verbose_proxy_logger @@ -48,6 +49,8 @@ else: PassThroughEndpointLogging = Any EndpointType = Any +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class AnthropicPassthroughLoggingHandler: @staticmethod @@ -343,11 +346,9 @@ class AnthropicPassthroughLoggingHandler: if not line.startswith("data:"): continue try: - data = json.loads(line[len("data:") :].strip()) + data = _JSON_OBJECT.validate_python(json.loads(line[len("data:") :].strip())) except (json.JSONDecodeError, ValueError): continue - if not isinstance(data, dict): - continue etype = data.get("type") if etype == "message_delta": return False diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 24c1b865db6..be40e225cde 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, cast from urllib.parse import urlparse import httpx -from pydantic import TypeAdapter +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._logging import verbose_proxy_logger @@ -55,6 +55,7 @@ else: _VERTEX_INTERACTIONS_PATH: Final = re.compile(r"/projects/[^/]+/locations/[^/]+/interactions/?$") _INTERACTIONS_RESPONSE_BODY: Final = TypeAdapter(dict[str, object]) +_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(hide_input_in_errors=True)) def _interactions_model( @@ -99,7 +100,7 @@ class VertexPassthroughLoggingHandler: litellm_model_response: Final = ModelResponse( model=model, usage=InteractionsUsageObjectTransformation.transform_interactions_usage_object( - cast(Mapping[str, Any], usage_object) + _JSON_OBJECT.validate_python(usage_object) ), ) logging_obj.custom_llm_provider = custom_llm_provider @@ -349,7 +350,7 @@ class VertexPassthroughLoggingHandler: model: Final = VertexPassthroughLoggingHandler.extract_model_from_url(url_route) - _json_response: Final[dict[str, object]] = httpx_response.json() + _json_response: Final = _JSON_OBJECT.validate_python(httpx_response.json()) litellm_prediction_response: ModelResponse | EmbeddingResponse | ImageResponse = ModelResponse() if VertexPassthroughLoggingHandler._is_audio_predict_response( diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index f55eb4f7863..efe7a5002e0 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -421,7 +421,7 @@ class PolicyRegistry: if policy_request.condition is not None: data["condition"] = json.dumps(policy_request.condition.model_dump()) if policy_request.pipeline is not None: - validated_pipeline: Final = GuardrailPipeline(**policy_request.pipeline) + validated_pipeline: Final = GuardrailPipeline.model_validate(policy_request.pipeline) data["pipeline"] = json.dumps(validated_pipeline.model_dump()) created_policy: Final = await _policy_table(prisma_client).create(data=data) @@ -496,7 +496,7 @@ class PolicyRegistry: if policy_request.condition is not None: update_data["condition"] = json.dumps(policy_request.condition.model_dump()) if policy_request.pipeline is not None: - validated_pipeline: Final = GuardrailPipeline(**policy_request.pipeline) + validated_pipeline: Final = GuardrailPipeline.model_validate(policy_request.pipeline) update_data["pipeline"] = json.dumps(validated_pipeline.model_dump()) updated_policy: Final = await _policy_table(prisma_client).update( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f40e541dd86..ac116255082 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -399,6 +399,7 @@ from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, _is_azure_model_router_request, _should_return_raw_model_name, + close_guarded_stream, create_response, log_llm_api_exception, open_sse_before_first_byte, @@ -422,6 +423,14 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( encrypt_value_helper, ) from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_body_call_id, with_call_id +from litellm.proxy.common_utils.fips import ( + FIPS_MODE_ENV_VAR, + SSL_VERIFY_ENV_VAR, + enforce_fips_boot_verdict, + fips_boot_verdict, + is_fips_mode, + openssl_enforces_fips, +) from litellm.proxy.common_utils.healthy_model_filter import ( get_hidden_unhealthy_model_names, is_healthy_only_listing_default, @@ -1329,6 +1338,16 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState if isinstance(worker_config, dict): await initialize_from_worker_config(worker_config) + enforce_fips_boot_verdict( + fips_boot_verdict( + raw_fips_mode=os.getenv(FIPS_MODE_ENV_VAR), + provider_enforces_fips=openssl_enforces_fips, + ssl_verify_environment=os.getenv(SSL_VERIFY_ENV_VAR), + ssl_verify_setting=litellm.ssl_verify, + ), + announce=announce_on_stderr_at_exit, + ) + enforce_master_key_boot_verdict( await with_stored_secrets_counted( master_key_boot_verdict( @@ -1368,10 +1387,21 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState try: result: Final = await migrate_passwords_to_scrypt_async(prisma_client) verbose_proxy_logger.info("Password migration: %s", result) + except ValueError as e: + verbose_proxy_logger.error( + "Password migration failed, so plaintext passwords stay unhashed in the database: %s. " + "This is what an OpenSSL FIPS provider reports when the hashing algorithm is not approved.", + e, + ) + if is_fips_mode(): + raise except Exception as e: verbose_proxy_logger.warning("Password migration skipped: %s", e) - asyncio.create_task(_run_pw_migration()) + if is_fips_mode(): + await _run_pw_migration() + else: + asyncio.create_task(_run_pw_migration()) async def _run_agent_grant_id_migration() -> None: from litellm.proxy.agent_endpoints.agent_registry import ( @@ -1574,11 +1604,6 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState if not model_info_scheduler.running: model_info_scheduler.start() - if scheduler is not None and prisma_client is not None: - from litellm.proxy.management_endpoints.roi_calculator_endpoints import register_scheduled_sync - - register_scheduled_sync(scheduler) - tracing_settings: Final = cast( # cast-ok: Pydantic validates the legacy untyped settings value dict[str, object] | None, TypeAdapter(dict[str, object] | None).validate_python(general_settings.get("tracing")), @@ -9882,6 +9907,17 @@ async def async_data_generator( stream_completed = False client_disconnected = False error_state: Final = ResponsesStreamErrorState() if responses_stream_errors else None + needs_iterator_wrap: Final = proxy_logging_obj.needs_iterator_wrap() + stream_iterator: Final[AsyncIterator[object]] = ( + proxy_logging_obj.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=response, + request_data=request_data, + ) + if needs_iterator_wrap + else response + ) + stream_source: AsyncIterator[object] | None = None # rebind-ok: bound once the keepalive policy resolves try: error_message: str | None = None requested_model_from_client: Final = _get_client_requested_model_for_streaming(request_data=request_data) @@ -9906,21 +9942,11 @@ async def async_data_generator( # per-chunk hook. Coalescing them into a single flag forced wasted # ``get_response_string`` work per chunk on every deployment that # happened to ship a streaming-iterator override (the default). - needs_iterator_wrap: Final = proxy_logging_obj.needs_iterator_wrap() needs_per_chunk_hook: Final = proxy_logging_obj.needs_per_chunk_streaming_hook() is_raw_sse_stream: Final = bool(request_data.get("_litellm_raw_sse_stream")) strip_stream_usage: Final = bool(request_data.get("_litellm_strip_stream_usage")) raw_sse_buffer = "" - if needs_iterator_wrap: - stream_iterator = proxy_logging_obj.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=response, - request_data=request_data, - ) - else: - stream_iterator = response - # A stream can start on a deployment with keepalive off and fall back # mid-stream to one that enables it: only skip wrapping altogether when # there's no router to ever fall back through AND the resolved interval @@ -9929,7 +9955,7 @@ async def async_data_generator( # happens to start with it off. resolve_keepalive_seconds: Final = _make_keepalive_resolver(request_data) initial_keepalive_seconds: Final = resolve_keepalive_seconds(response) - stream_source: Final = ( + stream_source = ( _iter_with_keepalive( stream_iterator.__aiter__(), resolve_keepalive_seconds, @@ -10070,6 +10096,9 @@ async def async_data_generator( # (a nested iterator hook would only see GeneratorExit on GC). if not stream_completed: client_disconnected = True + for guarded_layer in (stream_source, stream_iterator): + if guarded_layer is not response: + await close_guarded_stream(guarded_layer) raise except Exception as e: verbose_proxy_logger.exception("litellm.proxy.proxy_server.async_data_generator(): Exception occured - %s", e) diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index 7814903e975..da9a3033187 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -496,10 +496,11 @@ _AUTOROUTER_PRESETS_ADAPTER: Final = TypeAdapter(dict[str, AutoRouterPresetRecor def _load_bundled_autorouter_presets() -> Mapping[str, AutoRouterPresetRecord]: - raw: Final = json.loads( - files("litellm.proxy.public_endpoints").joinpath("autorouter_presets.json").read_text(encoding="utf-8") + return _AUTOROUTER_PRESETS_ADAPTER.validate_python( + json.loads( + files("litellm.proxy.public_endpoints").joinpath("autorouter_presets.json").read_text(encoding="utf-8") + ) ) - return _AUTOROUTER_PRESETS_ADAPTER.validate_python(raw) async def _fetch_remote_autorouter_presets(url: str) -> Mapping[str, AutoRouterPresetRecord]: diff --git a/litellm/proxy/roi_calculator/README.md b/litellm/proxy/roi_calculator/README.md index 82dd5463398..441f8f98c1c 100644 --- a/litellm/proxy/roi_calculator/README.md +++ b/litellm/proxy/roi_calculator/README.md @@ -1,54 +1,5 @@ -# ROI Calculator +# ROI Calculator (retired) -The dashboard compares merged pull or merge requests, elapsed time from opening to merge, new bug and regression issues, and recorded gateway spend over 7, 28, or 90 complete UTC days. Compare against the immediately preceding period of the same length or the same-length period last year +The integrated ROI Calculator is retired. Its dashboard page, API endpoints, OAuth callbacks, and scheduled refresh are no longer available -The calculator combines repository activity with spend recorded by the gateway. Spend per merged change is a person's recorded gateway spend during the period divided by their matched merged changes. To track a branch's AI cost, send repository and branch tags with each request - -## Connect repositories - -Use **Preview sample report** beside the title to explore the dashboard before connecting repositories. Sample periods, engineer details, quality signals, and branch spend work without changing your connections or live report. **Exit demo** returns to your report or setup - -Open `/ui/roi-calculator/`, choose GitHub or GitLab, then connect with an app or access token. Select several repositories and start the sync. Use **Add connection** to keep both providers connected. Each provider and API host retains its credentials, repositories, and identity mappings, and the report combines their activity while counting each personโ€™s gateway spend once. Public repositories also accept an empty token, subject to the provider's anonymous API limits - -For GitHub tokens, grant read access to metadata, pull requests and issues. GitLab tokens require `read_api`. Self-hosted instances use their API URL, for example `https://git.example.com/api/v4` - -## Configure app authorization - -Register a GitHub App with read-only repository permissions for metadata, pull requests and issues. Enable expiring user access tokens and leave authorization during installation disabled, since the gateway starts authorization after installation. Generate a private key in the app settings to allow installation and store it securely. The gateway uses a generated client secret for authorization and does not need the private key - -Register a confidential GitLab OAuth application with `read_api` and `read_user` scopes - -Set `PROXY_BASE_URL` to the gateway's public URL. The callback URLs are `/roi-calculator/observed/oauth/github/callback` and `/roi-calculator/observed/oauth/gitlab/callback` - -Set the GitHub App setup URL to `/roi-calculator/observed/oauth/github/installed`, enable **Redirect on update**, and set `LITELLM_ROI_GITHUB_APP_SLUG` to its URL slug. The first connection then starts with repository installation and continues to user authorization - -Set `LITELLM_ROI_GITHUB_CLIENT_ID` and `LITELLM_ROI_GITHUB_CLIENT_SECRET` for GitHub, or `LITELLM_ROI_GITLAB_CLIENT_ID` and `LITELLM_ROI_GITLAB_CLIENT_SECRET` for GitLab. For a self-hosted provider, set `LITELLM_ROI_GITHUB_URL` or `LITELLM_ROI_GITLAB_URL` to its base URL without the API suffix - -The gateway encrypts access and refresh tokens using its configured encryption key. Authorization uses PKCE and an expiring, single-use state tied to an HTTP-only browser cookie. Refreshes are coordinated across gateway workers - -## Link people - -Use **Link accounts** to associate several current or historical usernames with one internal email. Each connection has a separate username field, so a GitHub username never matches a GitLab user implicitly. Saving immediately recalculates the report without fetching repositories again. Public profile emails match automatically when they resolve unambiguously to an internal user - -**Matched people only** is on by default for people, merged changes, and branch lists. Turn it off to include outside contributors and their branches. Matching depends on the linked internal account, even when no spend was recorded. This switch filters the lists; summary metrics and quality signals still cover all selected repositories - -Agent-authored changes count for a person only when the supported agent metadata explicitly names a requester. Repository issue counts and revert titles are quality signals, not an individual defect score - -Bug and regression counts combine repositories with issue tracking enabled. They remain unavailable when none of the selected repositories has issue tracking enabled - -## Sync behavior - -The default refresh interval is daily and applies to every connection in the workspace. Adding or editing a connection preserves it unless `update_interval_minutes` is supplied. The observed settings API accepts `update_interval_minutes: 0` for manual updates. A cancelled or failed sync preserves the last complete report - -Existing settings retain `report_mode: legacy` and their scheduled reports until an administrator saves a connection, authorizes an app, or starts an observed sync. Reading the new dashboard alone does not change the mode. The legacy settings API can explicitly select `report_mode: legacy` again - -GitHub collection splits large searches into smaller date ranges to avoid its search-result limit. Both providers validate pagination and reject incomplete responses instead of publishing partial counts - - -## Branch request tags - -Send `repo:github.com/owner/repo` or `repo:gitlab.com/group/project` together with `branch:feature/name` in `metadata.tags`, top-level `tags`, or the comma-separated `x-litellm-tags` header. The tags must identify the source repository and branch, including forks - -The report sums recorded requests inside its UTC dates. A branch cost is assigned to a merged change only when that source branch matches one change in the period. Reused branches stay visible in Branch spend without duplicating costs across changes. No retained tagged requests means unknown cost; a recorded zero remains zero - -An empty repository produces a successful report with zero merged changes and no merge duration or spend-per-change ratio +The backend implementation and stored configuration are retained. Reverting the retirement change restores the dashboard and registration points. The standalone [litellm-roi-calculator](https://github.com/BerriAI/litellm-roi-calculator) project is separate and remains available diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 2dd1c9419f1..f9235a33b42 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -11,6 +11,7 @@ from types import MappingProxyType from typing import Final, NoReturn, SupportsFloat, SupportsIndex, SupportsInt, cast from fastapi import HTTPException, status +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._internal_context import with_service_target @@ -71,6 +72,9 @@ _COUNTER_ENTITY_TYPES: Final[Mapping[str, str]] = { "Project": Litellm_EntityType.PROJECT.value, } +_CACHED_VALUE: Final = TypeAdapter(object) +_WINDOW_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) + class _CounterReservationUnavailable(Exception): def __init__( @@ -762,8 +766,8 @@ async def _get_team_member_budget_counter( else: default_budget_id: Final = (team_object.metadata or {}).get("team_member_budget_id") if isinstance(default_budget_id, str): - default_budget: Final = await user_api_key_cache.async_get_cache( - key=f"team_member_default_budget:{default_budget_id}", + default_budget: Final = _CACHED_VALUE.validate_python( + await user_api_key_cache.async_get_cache(key=f"team_member_default_budget:{default_budget_id}") ) default_cap: Final = _to_float(_get_value(default_budget, "max_budget")) if default_cap is not None and default_cap > 0: @@ -799,8 +803,8 @@ async def _get_org_budget_counter( if org_id is None: return None - org_table: Final = await user_api_key_cache.async_get_cache( - key=f"org_id:{org_id}:with_budget", + org_table: Final = _CACHED_VALUE.validate_python( + await user_api_key_cache.async_get_cache(key=f"org_id:{org_id}:with_budget") ) if org_table is None: return None @@ -833,7 +837,9 @@ async def _get_project_budget_counter( return None source_cache_key: Final = project_cache_key(valid_token.project_id) - project_object: Final = await user_api_key_cache.async_get_cache(key=source_cache_key) + project_object: Final = _CACHED_VALUE.validate_python( + await user_api_key_cache.async_get_cache(key=source_cache_key) + ) if project_object is None: return None @@ -901,10 +907,9 @@ def _coerce_window(window: object) -> Mapping[str, object]: return window if isinstance(window, str): try: - parsed: Final[object] = json.loads(window) + return _WINDOW_OBJECT.validate_python(json.loads(window)) except Exception: return {} - return parsed if isinstance(parsed, Mapping) else {} model_dump: Final = getattr(window, "model_dump", None) if not callable(model_dump): return {} diff --git a/litellm/proxy/spend_tracking/spend_log_error_logger.py b/litellm/proxy/spend_tracking/spend_log_error_logger.py index 03037ba854f..28cd535127f 100644 --- a/litellm/proxy/spend_tracking/spend_log_error_logger.py +++ b/litellm/proxy/spend_tracking/spend_log_error_logger.py @@ -25,7 +25,7 @@ troubleshoot. The UI suppression follows the same gate. import logging import os -from typing import Any, Final +from typing import Final from litellm._logging import verbose_proxy_logger from litellm.secret_managers.main import str_to_bool @@ -58,7 +58,7 @@ def should_suppress_spend_log_tracebacks() -> bool: def spend_log_error( message: str, - *args: Any, + *args: object, exc: BaseException | None = None, ) -> None: """Log a spend-tracking error, with the traceback gated on the env var. diff --git a/litellm/proxy/tracing_endpoints.py b/litellm/proxy/tracing_endpoints.py index 0d3723c45b8..6ec0c6da3e7 100644 --- a/litellm/proxy/tracing_endpoints.py +++ b/litellm/proxy/tracing_endpoints.py @@ -29,6 +29,13 @@ from litellm.proxy.common_utils.http_parsing_utils import is_otlp_trace_request from litellm.proxy.tracing_runtime import provide_receiver, require_receiver from litellm.rust_bridge.trace.errors import TraceChanged from litellm.rust_bridge.trace.generated.models import TraceQueryHelp +from litellm.rust_bridge.trace.generated.requests import ( + TraceDetailRequest, + TraceErrorPageRequest, + TraceListRequest, + TraceQueryRequest, + TraceSpanRequest, +) from litellm.rust_bridge.trace.generated.types import ( AllQueryScope, OwnedQueryScope, @@ -49,6 +56,10 @@ router = APIRouter(tags=["agent tracing"]) MS_PER_DAY: Final = 24 * 60 * 60 * 1000 +async def current_time_ms() -> int: + return int(time.time() * 1000) + + @dataclass(frozen=True, slots=True) class TraceAccessContext: receiver: TraceReceiver | None @@ -183,28 +194,21 @@ def read_failure(error: TraceChanged | ValueError | OverflowError | RuntimeError @router.get("/v1/traces", response_model=TracePage) async def list_agent_traces( context: Annotated[TraceAccessContext, Depends(provide_trace_access)], - start_ms: Annotated[int | None, Query(description="Window start, unix ms. Default: 24h ago")] = None, - end_ms: Annotated[int | None, Query(description="Window end, unix ms. Default: now")] = None, - cursor: Annotated[str | None, Query(max_length=512)] = None, + now_ms: Annotated[int, Depends(current_time_ms)], + request: Annotated[TraceListRequest, Query()], ) -> TracePage: - now_ms: Final = int(time.time() * 1000) try: tracing, scope = context.reader() return await tracing.list_traces( scope=scope, - start_ms=start_ms if start_ms is not None else now_ms - MS_PER_DAY, - end_ms=end_ms if end_ms is not None else now_ms, - cursor=cursor, + start_ms=request.start_ms if request.start_ms is not None else now_ms - MS_PER_DAY, + end_ms=request.end_ms if request.end_ms is not None else now_ms, + cursor=request.cursor, ) except (TraceChanged, ValueError, OverflowError, RuntimeError) as error: raise read_failure(error) from error -class TraceQueryRequest(BaseModel): - model_config = ConfigDict(frozen=True, extra="forbid") - sql: str - - @dataclass(frozen=True, slots=True) class TraceQueryAccess: storage: ClickHouseStorage @@ -272,13 +276,11 @@ async def help_agent_trace_queries( async def get_agent_trace( trace_id: str, context: Annotated[TraceAccessContext, Depends(provide_trace_access)], - trace_ref: Annotated[str, Query()] = "", - cursor: Annotated[str | None, Query(max_length=512)] = None, - page_size: Annotated[int | None, Query(ge=1, le=500)] = None, + request: Annotated[TraceDetailRequest, Query()], ) -> Trace: tracing, scope = context.reader() try: - trace: Final = await tracing.get_trace(trace_id, scope, trace_ref, cursor, page_size) + trace: Final = await tracing.get_trace(trace_id, scope, request.trace_ref, request.cursor, request.page_size) except (TraceChanged, ValueError, OverflowError, RuntimeError) as error: raise read_failure(error) from error if trace is None: @@ -291,11 +293,11 @@ async def get_agent_trace_span( trace_id: str, span_id: str, context: Annotated[TraceAccessContext, Depends(provide_trace_access)], - trace_ref: Annotated[str, Query()] = "", + request: Annotated[TraceSpanRequest, Query()], ) -> SpanDetail: tracing, scope = context.reader() try: - span: Final = await tracing.get_span(trace_id, span_id, scope, trace_ref) + span: Final = await tracing.get_span(trace_id, span_id, scope, request.trace_ref) except (TraceChanged, ValueError, OverflowError, RuntimeError) as error: raise read_failure(error) from error if span is None: @@ -308,12 +310,11 @@ async def get_agent_trace_span_error( trace_id: str, span_id: str, context: Annotated[TraceAccessContext, Depends(provide_trace_access)], - trace_ref: Annotated[str, Query()] = "", - cursor: Annotated[str | None, Query(max_length=512)] = None, + request: Annotated[TraceErrorPageRequest, Query()], ) -> SpanErrorPage: try: tracing, scope = context.reader() - page: Final = await tracing.get_span_error(trace_id, span_id, scope, trace_ref, cursor) + page: Final = await tracing.get_span_error(trace_id, span_id, scope, request.trace_ref, request.cursor) except (TraceChanged, ValueError, OverflowError, RuntimeError) as error: raise read_failure(error) from error if page is None: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 1c738aefb97..29f2f46f001 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -44,6 +44,7 @@ from typing import ( Union, cast, overload, + runtime_checkable, ) from typing_extensions import ReadOnly, TypedDict @@ -524,8 +525,13 @@ class _UpstreamStreamBoundary(Generic[_T]): raise +@runtime_checkable +class _ClosableAsyncIterator(Protocol): + def aclose(self) -> object: ... + + class _StreamIteratorHook(Protocol[_T]): - def __call__(self, *, response: AsyncIterator[_T]) -> AsyncGenerator[_T, None]: ... + def __call__(self, *, response: AsyncIterator[_T]) -> AsyncIterator[_T]: ... def _is_client_error_exception(exc: Exception) -> bool: @@ -2773,8 +2779,22 @@ class ProxyLogging: ) -> AsyncGenerator[_T, None]: upstream: Final = _UpstreamStreamBoundary(response) try: - async for chunk in hook(response=upstream): - yield chunk + guarded: Final = hook(response=upstream) + try: + async for chunk in guarded: + yield chunk + finally: + if isinstance(guarded, _ClosableAsyncIterator): + try: + closing: Final = guarded.aclose() + if inspect.isawaitable(closing): + await closing + except Exception as e: # noqa: BLE001 # a finished stream must not fail on callback cleanup + verbose_proxy_logger.warning( + "Closing the streaming iterator of %s raised %s", + getattr(callback, "guardrail_name", None) or type(callback).__name__, + type(e).__name__, + ) except Exception as e: if e is not upstream.failure: enrich_http_exception_with_guardrail_context(e, callback) @@ -3959,6 +3979,7 @@ class ProxyLogging: stream_needs_translation: Final = ProxyLogging._stream_requires_guardrail_translation(user_api_key_dict) pipeline_gated_names: Final = _pipeline_step_guardrail_names(post_call_pipelines) + guarded_layers: Final[list[AsyncGenerator[object, None]]] = [] # mutable-ok: closed on disconnect for resolved_callback, kind in caps.iterator_overrides: if isinstance(resolved_callback, CustomGuardrail): if resolved_callback.guardrail_name in pipeline_gated_names: @@ -4001,6 +4022,7 @@ class ProxyLogging: hook, request_data=request_data, ) + guarded_layers.append(current_response) pipeline_translation: Final = ( resolve_endpoint_translation(user_api_key_dict, None) if post_call_pipelines else None @@ -4013,6 +4035,7 @@ class ProxyLogging: pipelines=post_call_pipelines, translation=pipeline_translation, ) + guarded_layers.append(current_response) served_chunks: Final[list[object]] = [] # mutable-ok: accumulates while yielding to the client try: @@ -4020,6 +4043,7 @@ class ProxyLogging: served_chunks.append(chunk) yield chunk except (GeneratorExit, asyncio.CancelledError): + await ProxyLogging._close_guarded_layers(guarded_layers) ProxyLogging._record_served_stream_output(request_data, served_chunks) raise except Exception as e: @@ -4100,6 +4124,16 @@ class ProxyLogging: for buffered_item in buffered: yield buffered_item + @staticmethod + async def _close_guarded_layers(layers: Sequence[AsyncGenerator[object, None]]) -> None: + for layer in reversed(layers): + try: + await layer.aclose() + except Exception as e: # noqa: BLE001 # one failing callback cleanup must not skip the inner ones + verbose_proxy_logger.warning( + "Closing a streaming callback layer after a client disconnect raised %s", type(e).__name__ + ) + @staticmethod def _record_served_stream_output(request_data: Mapping[str, object], served_chunks: Sequence[object]) -> None: logging_obj: Final = request_data.get("litellm_logging_obj") diff --git a/litellm/rag/ingestion/base_ingestion.py b/litellm/rag/ingestion/base_ingestion.py index da1bc0a1feb..fb799546b37 100644 --- a/litellm/rag/ingestion/base_ingestion.py +++ b/litellm/rag/ingestion/base_ingestion.py @@ -23,6 +23,7 @@ from litellm.constants import DEFAULT_CHUNK_OVERLAP, DEFAULT_CHUNK_SIZE from litellm.litellm_core_utils.url_utils import async_safe_get from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, + header_value, httpxSpecialProvider, ) from litellm.rag.ingestion.file_parsers import extract_text_from_pdf @@ -79,7 +80,7 @@ class BaseRAGIngestion(ABC): from litellm.litellm_core_utils.credential_accessor import CredentialAccessor credential_name: Final = self.vector_store_config.get("litellm_credential_name") - if credential_name and litellm.credential_list: + if isinstance(credential_name, str) and credential_name and litellm.credential_list: credential_values: Final = CredentialAccessor.get_credential_values(credential_name) if not credential_values: return @@ -125,7 +126,8 @@ class BaseRAGIngestion(ABC): response.raise_for_status() file_content = response.content filename = file_url.split("/")[-1] or "document" - content_type = response.headers.get("content-type", "application/octet-stream") + content_type_header: Final = header_value(response.headers, "content-type") + content_type = "application/octet-stream" if content_type_header is None else content_type_header return filename, file_content, content_type, None if file_id: diff --git a/litellm/rag/ingestion/gemini_ingestion.py b/litellm/rag/ingestion/gemini_ingestion.py index 1cf5db549e4..a015f3cf614 100644 --- a/litellm/rag/ingestion/gemini_ingestion.py +++ b/litellm/rag/ingestion/gemini_ingestion.py @@ -9,6 +9,8 @@ from __future__ import annotations from typing import TYPE_CHECKING, Final, cast +from pydantic import BaseModel, ConfigDict + from litellm._logging import verbose_logger from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, @@ -23,6 +25,13 @@ if TYPE_CHECKING: from litellm.types.rag import RAGIngestOptions +class _WhiteSpaceConfig(BaseModel): + model_config = ConfigDict(extra="ignore", frozen=True, hide_input_in_errors=True) + + max_tokens_per_chunk: object = 800 + max_overlap_tokens: object = 400 + + class GeminiRAGIngestion(BaseRAGIngestion): """ Gemini-specific RAG ingestion using File Search API. @@ -234,12 +243,13 @@ class GeminiRAGIngestion(BaseRAGIngestion): # Add chunking configuration if provided chunking_strategy: Final = self.chunking_strategy if chunking_strategy and isinstance(chunking_strategy, dict): - white_space_config: Final = chunking_strategy.get("white_space_config") + white_space_config: Final[object] = chunking_strategy.get("white_space_config") if white_space_config: + white_space: Final = _WhiteSpaceConfig.model_validate(white_space_config) request_body["chunkingConfig"] = { "whiteSpaceConfig": { - "maxTokensPerChunk": white_space_config.get("max_tokens_per_chunk", 800), - "maxOverlapTokens": white_space_config.get("max_overlap_tokens", 400), + "maxTokensPerChunk": white_space.max_tokens_per_chunk, + "maxOverlapTokens": white_space.max_overlap_tokens, } } diff --git a/litellm/rag/ingestion/openai_ingestion.py b/litellm/rag/ingestion/openai_ingestion.py index 925a49d0f7a..6c7c619237d 100644 --- a/litellm/rag/ingestion/openai_ingestion.py +++ b/litellm/rag/ingestion/openai_ingestion.py @@ -7,7 +7,7 @@ so this implementation skips the embedding step and directly uploads files. from __future__ import annotations -from typing import TYPE_CHECKING, Any, Final, cast +from typing import TYPE_CHECKING, Final import litellm from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion @@ -106,7 +106,7 @@ class OpenAIRAGIngestion(BaseRAGIngestion): vector_store_id=vector_store_id, file_id=existing_file_id, custom_llm_provider="openai", - chunking_strategy=cast(dict[str, Any] | None, self.chunking_strategy), + chunking_strategy=self.chunking_strategy, api_key=api_key, api_base=api_base, ) @@ -134,7 +134,7 @@ class OpenAIRAGIngestion(BaseRAGIngestion): vector_store_id=vector_store_id, file_id=result_file_id, custom_llm_provider="openai", - chunking_strategy=cast(dict[str, Any] | None, self.chunking_strategy), + chunking_strategy=self.chunking_strategy, api_key=api_key, api_base=api_base, ) diff --git a/litellm/rag/ingestion/s3_vectors_ingestion.py b/litellm/rag/ingestion/s3_vectors_ingestion.py index e2aa5555eec..e3eafd680de 100644 --- a/litellm/rag/ingestion/s3_vectors_ingestion.py +++ b/litellm/rag/ingestion/s3_vectors_ingestion.py @@ -18,7 +18,7 @@ from __future__ import annotations import hashlib import uuid from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final, TypedDict +from typing import TYPE_CHECKING, Final, TypedDict import litellm from litellm._logging import verbose_logger @@ -194,7 +194,7 @@ class S3VectorsRAGIngestion(BaseRAGIngestion, BaseAWSLLM): url: str, data: str | None = None, headers: dict[str, str] | None = None, - ) -> Any: + ) -> httpx.Response: """ Helper to sign and execute AWS API requests using httpx + SigV4. diff --git a/litellm/rag/main.py b/litellm/rag/main.py index 1a5301f0579..5ab39bec080 100644 --- a/litellm/rag/main.py +++ b/litellm/rag/main.py @@ -18,6 +18,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final import httpx +from pydantic import ConfigDict, TypeAdapter import litellm from litellm._internal_context import is_internal_call @@ -71,6 +72,13 @@ _SEARCH_ARGS_SET_BY_PIPELINE: Final = frozenset( ) +_VECTOR_STORE_OPTIONS: Final = TypeAdapter(Mapping[object, object], config=ConfigDict(hide_input_in_errors=True)) + + +def _vector_store_provider(ingest_options: Mapping[str, object]) -> object: + return _VECTOR_STORE_OPTIONS.validate_python(ingest_options.get("vector_store", {})).get("custom_llm_provider") + + def get_ingestion_class(provider: str) -> type[BaseRAGIngestion]: """ Get the ingestion class for a given provider. @@ -137,7 +145,7 @@ async def _execute_ingest_pipeline( @client async def aingest( - ingest_options: dict[str, Any], + ingest_options: Mapping[str, object], file_data: tuple[str, bytes, str] | None = None, file: dict[str, str] | None = None, file_url: str | None = None, @@ -197,7 +205,7 @@ async def aingest( except Exception as e: raise litellm.exception_type( model=None, - custom_llm_provider=ingest_options.get("vector_store", {}).get("custom_llm_provider"), + custom_llm_provider=_vector_store_provider(ingest_options), original_exception=e, completion_kwargs=local_vars, extra_kwargs=kwargs, @@ -226,7 +234,7 @@ def _suppressed_sub_call_billing() -> Iterator[None]: async def _execute_query_pipeline( model: str, messages: list[AllMessageValues], - retrieval_config: dict[str, Any], + retrieval_config: Mapping[str, object], rerank: dict[str, Any] | None = None, stream: bool = False, vector_store_params: Mapping[str, object] | None = None, @@ -276,12 +284,16 @@ async def _execute_query_pipeline( search_provider: Final = retrieval_config.get("custom_llm_provider", "openai") try: - search_cost = sum( - vector_store_search_cost( - model=search_provider if "/" in search_provider else None, - custom_llm_provider=search_provider, - response=search_response, + search_cost = ( + sum( + vector_store_search_cost( + model=search_provider if "/" in search_provider else None, + custom_llm_provider=search_provider, + response=search_response, + ) ) + if isinstance(search_provider, str) + else 0.0 ) except Exception: # noqa: BLE001 - cost accounting must never break the query path search_cost = 0.0 @@ -354,7 +366,7 @@ async def _execute_query_pipeline( async def aquery( model: str, messages: list[AllMessageValues], - retrieval_config: dict[str, Any], + retrieval_config: Mapping[str, object], rerank: dict[str, Any] | None = None, stream: bool = False, vector_store_params: Mapping[str, object] | None = None, @@ -403,7 +415,7 @@ async def aquery( def query( model: str, messages: list[AllMessageValues], - retrieval_config: dict[str, Any], + retrieval_config: Mapping[str, object], rerank: dict[str, Any] | None = None, stream: bool = False, vector_store_params: Mapping[str, object] | None = None, @@ -450,7 +462,7 @@ def query( @client def ingest( - ingest_options: dict[str, Any], + ingest_options: Mapping[str, object], file_data: tuple[str, bytes, str] | None = None, file: dict[str, str] | None = None, file_url: str | None = None, @@ -517,7 +529,7 @@ def ingest( except Exception as e: raise litellm.exception_type( model=None, - custom_llm_provider=ingest_options.get("vector_store", {}).get("custom_llm_provider"), + custom_llm_provider=_vector_store_provider(ingest_options), original_exception=e, completion_kwargs=local_vars, extra_kwargs=kwargs, diff --git a/litellm/repositories/autorouter_session_repository.py b/litellm/repositories/autorouter_session_repository.py index 82e2c091728..640fd26ffe3 100644 --- a/litellm/repositories/autorouter_session_repository.py +++ b/litellm/repositories/autorouter_session_repository.py @@ -2,7 +2,7 @@ Repository for the auto-router per-session rollup (LiteLLM_AutoRouterSession). """ -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, Protocol from litellm.models.autorouter_session import LiteLLM_AutoRouterSession from litellm.repositories.base_repository import BaseRepository @@ -12,10 +12,21 @@ if TYPE_CHECKING: from prisma import models as prisma_models +class _AutoRouterSessionDb(Protocol): + @property + def litellm_autoroutersession(self) -> TableActions["prisma_models.LiteLLM_AutoRouterSession"]: ... + + +class _PrismaClientView(Protocol): + @property + def db(self) -> _AutoRouterSessionDb: ... + + class AutoRouterSessionRepository(BaseRepository[LiteLLM_AutoRouterSession]): @property def table(self) -> TableActions["prisma_models.LiteLLM_AutoRouterSession"]: - return self.prisma_client.db.litellm_autoroutersession + client: Final[_PrismaClientView] = self.prisma_client + return client.db.litellm_autoroutersession @property def model_class(self) -> type[LiteLLM_AutoRouterSession]: diff --git a/litellm/repositories/model_repository.py b/litellm/repositories/model_repository.py index cbacb6e8f90..6541db4c00b 100644 --- a/litellm/repositories/model_repository.py +++ b/litellm/repositories/model_repository.py @@ -4,20 +4,24 @@ Model repository for database operations on LiteLLM_ProxyModelTable. import json from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Final + +from pydantic import ConfigDict, TypeAdapter from litellm.models.model import LiteLLM_ProxyModelTable from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.base_repository import BaseRepository, DbRecord, record_to_dict from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import PrismaTableRepository if TYPE_CHECKING: from prisma import models as prisma_models +_LITELLM_PARAMS: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(strict=True, hide_input_in_errors=True)) + class _ProxyModelTableRepository(PrismaTableRepository["prisma_models.LiteLLM_ProxyModelTable"]): table_name = "litellm_proxymodeltable" @@ -60,22 +64,26 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): decrypted[key] = value return decrypted - def _to_model(self, record: Any) -> LiteLLM_ProxyModelTable | None: + def _to_model(self, record: DbRecord | None) -> LiteLLM_ProxyModelTable | None: """Convert a database record to a Model with decryption.""" if record is None: return None - data: Final = record.dict() if hasattr(record, "dict") else dict(record) + data: Final = dict(record_to_dict(record)) - if isinstance(data.get("litellm_params"), str): - data["litellm_params"] = json.loads(data["litellm_params"]) - if isinstance(data.get("model_info"), str): - data["model_info"] = json.loads(data["model_info"]) + litellm_params: Final = data.get("litellm_params") + if isinstance(litellm_params, str): + data["litellm_params"] = json.loads(litellm_params) + model_info: Final = data.get("model_info") + if isinstance(model_info, str): + data["model_info"] = json.loads(model_info) if data.get("litellm_params"): - data["litellm_params"] = self._decrypt_litellm_params(data["litellm_params"]) + data["litellm_params"] = self._decrypt_litellm_params( + _LITELLM_PARAMS.validate_python(data["litellm_params"]) + ) - return LiteLLM_ProxyModelTable(**data) + return LiteLLM_ProxyModelTable.model_validate(data) async def find_by_id(self, model_id: str, id_field: str = "model_id") -> LiteLLM_ProxyModelTable | None: return await super().find_by_id(model_id, id_field) diff --git a/litellm/repositories/object_permission_repository.py b/litellm/repositories/object_permission_repository.py index b732d2ff94c..99d6aeed88b 100644 --- a/litellm/repositories/object_permission_repository.py +++ b/litellm/repositories/object_permission_repository.py @@ -2,7 +2,7 @@ ObjectPermission repository for database operations on LiteLLM_ObjectPermissionTable. """ -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Final, Protocol from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.repositories.base_repository import BaseRepository @@ -12,6 +12,19 @@ if TYPE_CHECKING: from prisma import models as prisma_models +class _ObjectPermissionDb(Protocol): + @property + def litellm_objectpermissiontable(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]: ... + + +class _PrismaClientView(Protocol): + @property + def db(self) -> _ObjectPermissionDb: ... + + @property + def writer_db(self) -> _ObjectPermissionDb: ... + + class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): """Repository for object permission database operations.""" @@ -21,7 +34,8 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): @property def table(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]: - database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db + client: Final[_PrismaClientView] = self.prisma_client + database: Final = client.writer_db if self._use_writer else client.db return database.litellm_objectpermissiontable @property @@ -48,7 +62,7 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): skills: list[str] | None = None, ) -> LiteLLM_ObjectPermissionTable: """Create a new object permission record.""" - data: Final[dict[str, Any]] = {} + data: Final[dict[str, object]] = {} if mcp_servers is not None: data["mcp_servers"] = mcp_servers if mcp_access_groups is not None: @@ -90,7 +104,7 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): skills: list[str] | None = None, ) -> LiteLLM_ObjectPermissionTable | None: """Update an object permission record.""" - data: Final[dict[str, Any]] = {} + data: Final[dict[str, object]] = {} if mcp_servers is not None: data["mcp_servers"] = mcp_servers if mcp_access_groups is not None: diff --git a/litellm/repositories/project_repository.py b/litellm/repositories/project_repository.py index 905e813f35e..0ea5009a5b6 100644 --- a/litellm/repositories/project_repository.py +++ b/litellm/repositories/project_repository.py @@ -3,7 +3,7 @@ Project repository for database operations on LiteLLM_ProjectTable. """ from collections.abc import Mapping -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, Protocol from litellm.models.project import LiteLLM_ProjectTable from litellm.repositories.base_repository import BaseRepository @@ -13,12 +13,23 @@ if TYPE_CHECKING: from prisma import models as prisma_models +class _ProjectDb(Protocol): + @property + def litellm_projecttable(self) -> TableActions["prisma_models.LiteLLM_ProjectTable"]: ... + + +class _PrismaClientView(Protocol): + @property + def db(self) -> _ProjectDb: ... + + class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): """Repository for project database operations.""" @property def table(self) -> TableActions["prisma_models.LiteLLM_ProjectTable"]: - return self.prisma_client.db.litellm_projecttable + client: Final[_PrismaClientView] = self.prisma_client + return client.db.litellm_projecttable @property def model_class(self) -> type[LiteLLM_ProjectTable]: diff --git a/litellm/repositories/user_repository.py b/litellm/repositories/user_repository.py index b201dbc566b..5ed1574c6bc 100644 --- a/litellm/repositories/user_repository.py +++ b/litellm/repositories/user_repository.py @@ -5,7 +5,7 @@ User repository for database operations on LiteLLM_UserTable. import json from collections.abc import Mapping, Sequence from itertools import chain -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, Protocol from pydantic import TypeAdapter @@ -37,6 +37,19 @@ ORDER BY p.user_id _PLACEHOLDER_ROWS_ADAPTER: Final = TypeAdapter(tuple[SCIMPlaceholder, ...]) +class _UserDb(Protocol): + @property + def litellm_usertable(self) -> TableActions["prisma_models.LiteLLM_UserTable"]: ... + + +class _PrismaClientView(Protocol): + @property + def db(self) -> _UserDb: ... + + @property + def writer_db(self) -> _UserDb: ... + + class UserRepository(BaseRepository[LiteLLM_UserTable]): """Repository for user database operations.""" @@ -46,7 +59,8 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]): @property def table(self) -> TableActions["prisma_models.LiteLLM_UserTable"]: - database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db + client: Final[_PrismaClientView] = self.prisma_client + database: Final = client.writer_db if self._use_writer else client.db return database.litellm_usertable @property diff --git a/litellm/router.py b/litellm/router.py index 0b9f12c8da3..afc842a05a0 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -9210,23 +9210,13 @@ class Router: if limit_violation is not None: raise ValueError(limit_violation) - default_model: str | None = deployment.litellm_params.complexity_router_default_model - - # If no default model specified, try to get from config tiers. Derived from the - # validated model, not the raw dict, so normalization (e.g. fallback_tier - # whitespace) is applied by its one owner before the tiers lookup. - if default_model is None and complexity_router_config: - validated: Final = ComplexityRouterConfig.model_validate(complexity_router_config) - # Custom tier sets name their fallback tier; built-in sets default to MEDIUM or SIMPLE - derived: Final = ( - (validated.tiers.get(validated.fallback_tier) if validated.fallback_tier is not None else None) - or validated.tiers.get("MEDIUM") - or validated.tiers.get("SIMPLE") + default_model: Final = ( + ComplexityRouterConfig.model_validate(complexity_router_config).resolve_default_model( + deployment.litellm_params.complexity_router_default_model ) - if isinstance(derived, list): - default_model = derived[0] if derived else None - else: - default_model = derived + if complexity_router_config + else deployment.litellm_params.complexity_router_default_model + ) if default_model is None: raise ValueError( diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index 9aa881fc4c9..95c981502e7 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -123,6 +123,9 @@ class AdaptiveRouter: prefs = self.model_to_prefs.get(model) or _default_prefs() self._cells[(rt, model)] = initial_cell(prefs, rt) + def cell(self, request_type: RequestType, model: str) -> BanditCell: + return self._cells[(request_type, model)] + async def load_state_from_db(self, prisma_client: object) -> None: """Add each row's persisted delta to a freshly computed cold-start prior. diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index fc0865654c7..43200271f9f 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -2945,7 +2945,7 @@ class ComplexityRouter(CustomLogger): raise ValueError(f"No candidate models left for tier {tier_key} after routing-plugin filtering") return self._pick_from_tier_value(context.candidate_models, tier_key) - def _ensure_adaptive_router(self) -> Any | None: + def _ensure_adaptive_router(self) -> AdaptiveRouter | None: if not self.config.adaptive: return None if self.adaptive_router is not None: @@ -3087,7 +3087,7 @@ class ComplexityRouter(CustomLogger): pools: Final = self._tier_pools() classified_candidates: Final = _allowed(tuple(pools.get(_tier_name(classified_tier), ())), fit_filter) cold_start_candidates: Final = tuple( - model for model in classified_candidates if adaptive._cells[(request_type, model)].total_samples == 0 + model for model in classified_candidates if adaptive.cell(request_type, model).total_samples == 0 ) if cold_start_candidates: chosen_model: Final = random.choice(cold_start_candidates) @@ -3106,7 +3106,7 @@ class ComplexityRouter(CustomLogger): "candidates": [ { "model": model, - "total_samples": adaptive._cells[(request_type, model)].total_samples, + "total_samples": adaptive.cell(request_type, model).total_samples, } for model in cold_start_candidates ], @@ -3123,7 +3123,7 @@ class ComplexityRouter(CustomLogger): best_score = float("-inf") candidate_scores: Final[list[dict[str, object]]] = [] for model in self._adaptive_candidate_models(classified_tier, hard_floor, hard_ceiling, fit_filter): - cell = adaptive._cells[(request_type, model)] + cell = adaptive.cell(request_type, model) quality_sample = thompson_sample(cell) cost_score = normalized_cost(adaptive.model_to_cost.get(model, 0.0), all_costs) if self.config.adaptive_eligible == "classified_tier": diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 88907731468..8d00b45fed0 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -1955,6 +1955,20 @@ class ComplexityRouterConfig(BaseModel): def _normalize_classification_examples_field(cls, value: str | None) -> str | None: return normalize_classification_examples(value) + def resolve_default_model(self, default_model: str | None = None) -> str | None: + if default_model is not None: + return default_model + if self.default_model is not None: + return self.default_model + derived: Final = ( + (self.tiers.get(self.fallback_tier) if self.fallback_tier is not None else None) + or self.tiers.get("MEDIUM") + or self.tiers.get("SIMPLE") + ) + if isinstance(derived, list): + return derived[0] if derived else None + return derived + @property def has_custom_tiers(self) -> bool: """True when the operator replaced the built-in tier set via tier_definitions.""" diff --git a/litellm/router_utils/reasoning_effort_capability.py b/litellm/router_utils/reasoning_effort_capability.py index 1d7656f253e..200de29b664 100644 --- a/litellm/router_utils/reasoning_effort_capability.py +++ b/litellm/router_utils/reasoning_effort_capability.py @@ -34,7 +34,7 @@ from typing import Final, get_args import litellm from litellm.types.llms.openai import REASONING_EFFORT -REASONING_EFFORT_ADVERTISEMENT_ORDER: Final = get_args(REASONING_EFFORT) +REASONING_EFFORT_ADVERTISEMENT_ORDER: Final[tuple[str, ...]] = get_args(REASONING_EFFORT) _EMPTY_ENTRY: Final[Mapping[str, object]] = MappingProxyType({}) _EFFORT_FLAGS: Final = ( diff --git a/litellm/rust_bridge/trace/AGENTS.md b/litellm/rust_bridge/trace/AGENTS.md new file mode 100644 index 00000000000..5132b6ce74c --- /dev/null +++ b/litellm/rust_bridge/trace/AGENTS.md @@ -0,0 +1,7 @@ +# Trace contract boundary + +- `generated/` is output of `scripts/generate_trace_types.py` from the `litellm-traces` schemas; change the Rust type and regenerate, never edit these files +- CI runs the generator with `--check` and fails on drift +- `storage.py` validates every native result against the generated response models before returning it +- Keep the trace methods in `_native.pyi` matching `litellm-rust/crates/python-bridge/src/routes/traces.rs` +- Native trace methods take scalar arguments; moving them to the generated request types changes overflow errors, so it needs its own behavior-change PR diff --git a/litellm/rust_bridge/trace/generated/requests.py b/litellm/rust_bridge/trace/generated/requests.py new file mode 100644 index 00000000000..dd5304ea09b --- /dev/null +++ b/litellm/rust_bridge/trace/generated/requests.py @@ -0,0 +1,59 @@ +# @generated by scripts/generate_trace_types.py, do not edit + +from __future__ import annotations + +from typing import Annotated, TypeAlias + +from pydantic import BaseModel, ConfigDict, Field + + +class TraceDetailRequest(BaseModel): + model_config = ConfigDict( + frozen=True, + ) + + trace_ref: str = "" + cursor: str | None = Field(None, max_length=512) + page_size: int | None = Field(None, ge=1, le=500) + + +class TraceErrorPageRequest(BaseModel): + model_config = ConfigDict( + frozen=True, + ) + + trace_ref: str = "" + cursor: str | None = Field(None, max_length=512) + + +class TraceListRequest(BaseModel): + model_config = ConfigDict( + frozen=True, + ) + + start_ms: int | None = Field(None, description="Window start, unix ms. Default: 24h ago") + end_ms: int | None = Field(None, description="Window end, unix ms. Default: now") + cursor: str | None = Field(None, max_length=512) + + +class TraceQueryRequest(BaseModel): + model_config = ConfigDict( + extra="forbid", + frozen=True, + ) + + sql: str + + +class TraceSpanRequest(BaseModel): + model_config = ConfigDict( + frozen=True, + ) + + trace_ref: str = "" + + +TraceWireRequests: TypeAlias = Annotated[ + TraceDetailRequest | TraceErrorPageRequest | TraceListRequest | TraceQueryRequest | TraceSpanRequest, + Field(..., title="TraceWireRequests"), +] diff --git a/litellm/secret_managers/aws_secret_manager.py b/litellm/secret_managers/aws_secret_manager.py index 1aed44e9a31..b3a9fbf31c8 100644 --- a/litellm/secret_managers/aws_secret_manager.py +++ b/litellm/secret_managers/aws_secret_manager.py @@ -14,9 +14,13 @@ import os import re from typing import Any, Final +from pydantic import TypeAdapter + import litellm from litellm.proxy._types import KeyManagementSystem +_PARSED_LITERAL: Final = TypeAdapter(object) + def validate_environment(): if "AWS_REGION_NAME" not in os.environ: @@ -107,7 +111,7 @@ class AWSKeyManagementService_V2: if isinstance(secret, str): secret = secret.strip() try: - secret_value_as_bool: Final = ast.literal_eval(secret) + secret_value_as_bool: Final = _PARSED_LITERAL.validate_python(ast.literal_eval(secret)) if isinstance(secret_value_as_bool, bool): return secret_value_as_bool except Exception: diff --git a/litellm/secret_managers/main.py b/litellm/secret_managers/main.py index fbed4fb8e75..275c029b0b7 100644 --- a/litellm/secret_managers/main.py +++ b/litellm/secret_managers/main.py @@ -6,7 +6,7 @@ import traceback from typing import Final import httpx -from pydantic import BaseModel, ValidationError +from pydantic import BaseModel, TypeAdapter, ValidationError import litellm from litellm._logging import verbose_logger @@ -19,6 +19,8 @@ from litellm.secret_managers.get_azure_ad_token_provider import ( oidc_cache: Final = DualCache() +_PARSED_LITERAL: Final = TypeAdapter(object) + _OIDC_TOKEN_EXPIRY_MARGIN_SECONDS: Final = 60 @@ -348,7 +350,7 @@ def get_secret( secret = os.getenv(secret_name) try: if isinstance(secret, str): - secret_value_as_bool = ast.literal_eval(secret) + secret_value_as_bool = _PARSED_LITERAL.validate_python(ast.literal_eval(secret)) if isinstance(secret_value_as_bool, bool): return secret_value_as_bool else: diff --git a/litellm/tracing/AGENTS.md b/litellm/tracing/AGENTS.md index d71664c38eb..82c4d4ac8a1 100644 --- a/litellm/tracing/AGENTS.md +++ b/litellm/tracing/AGENTS.md @@ -4,3 +4,4 @@ - Use `litellm.rust_bridge.trace.storage.ClickHouseStorage` for ClickHouse; keep trace schema, SQL and encoding in `litellm-traces`, and generic transport in `litellm-storage-clickhouse` - Derive tenant fields from authentication and overwrite matching fields supplied by the exporter - Test confirmed writes, failures, tenant isolation and read behavior through public functions +- Trace routes in `litellm/proxy/tracing_endpoints.py` bind query parameters with `Annotated[, Query()]` and bodies with the generated model; never redeclare field constraints in `Query(...)` or a local model diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 46026c12d24..8f1140cef62 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -145,6 +145,7 @@ class SupportedGuardrailIntegrations(Enum): STRAIKER = "straiker" ALICE = "alice" AGENT_365 = "agent_365" + LLM_SHIELD_PROXY = "llm_shield_proxy" CONDUCT = "conduct" diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index ce86e79a6f0..949506898b1 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -9,7 +9,7 @@ from urllib.parse import urlsplit import httpx from pydantic import BaseModel, ConfigDict, Field -from typing_extensions import TypedDict +from typing_extensions import ReadOnly, TypedDict from litellm.types.llms.base import HiddenParams @@ -161,6 +161,9 @@ MCPTokenEndpointAuthMethod = Literal["client_secret_basic", "client_secret_post" class MCPCredentials(TypedDict, total=False): + dcr_issuer: ReadOnly[str | None] + dcr_server_url: ReadOnly[str | None] + auth_value: str | None """ Authentication value diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 2ee19b3e59a..343fb731b06 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -38,6 +38,7 @@ class MCPOAuthMetadata(BaseModel): authorization_url: str | None = None token_url: str | None = None registration_url: str | 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 trust-on-first-use as the server's ``issuer`` when none is configured, so that later rebuilds @@ -118,6 +119,9 @@ class MCPServer(BaseModel): client_secret: str | None = None issuer: str | None = None issuer_is_anchored: bool = False + authorization_response_iss_parameter_supported: bool = False + dcr_issuer: str | None = None + dcr_server_url: str | None = None scopes: list[str] | None = None authorization_url: str | None = None token_url: str | None = None diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py b/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py new file mode 100644 index 00000000000..967d7ee75c2 --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/llm_shield_proxy.py @@ -0,0 +1,24 @@ +from pydantic import Field + +from .base import GuardrailConfigModel + + +class LLMShieldProxyGuardrailConfigModel(GuardrailConfigModel): + api_key: str | None = Field( + default=None, + description=( + "The virtual key for the LLM Shield Proxy instance. If not provided, the " + "`LLM_SHIELD_PROXY_API_KEY` environment variable is checked." + ), + ) + api_base: str | None = Field( + default=None, + description=( + "The base URL of the LLM Shield Proxy instance. If not provided, the `LLM_SHIELD_PROXY_API_BASE` " + "environment variable is checked, then `http://localhost:8000`." + ), + ) + + @staticmethod + def ui_friendly_name() -> str: + return "LLM Shield Proxy" diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index e78d2aa5f6a..c0cb3796934 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -10,6 +10,8 @@ from typing import ( get_args, ) +from pydantic import ConfigDict, TypeAdapter + from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import remove_items_at_indices from litellm.repositories.table_repositories import ( @@ -30,6 +32,8 @@ if TYPE_CHECKING: else: PrismaClient = Any +_DB_ROW: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(strict=True, hide_input_in_errors=True)) + class VectorStoreIndexRegistry: def __init__(self, vector_store_indexes: list[LiteLLM_ManagedVectorStoreIndex] = []): @@ -95,7 +99,9 @@ class VectorStoreIndexRegistry: ) for vector_store in _vector_stores_from_db: _dict_vector_store = dict(vector_store) - _litellm_managed_vector_store = LiteLLM_ManagedVectorStoreIndex(**_dict_vector_store) + _litellm_managed_vector_store = LiteLLM_ManagedVectorStoreIndex.model_validate( + _DB_ROW.validate_python(_dict_vector_store) + ) vector_stores_from_db.append(_litellm_managed_vector_store) return vector_stores_from_db diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 2d0985bcb64..28ad23fd372 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -4663,6 +4663,7 @@ "azure/eu/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2.2e-05, "cache_read_input_token_cost": 2.75e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.00011, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", @@ -4672,6 +4673,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.00022, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -6306,6 +6308,7 @@ "azure/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2e-05, "cache_read_input_token_cost": 2.5e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.0001, "input_cost_per_token": 5e-06, "litellm_provider": "azure", @@ -6315,6 +6318,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.0002, "output_cost_per_token": 2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -6350,7 +6354,7 @@ "supports_tool_choice": true }, "azure/gpt-4o-transcribe": { - "deprecation_date": "2026-12-31", + "deprecation_date": "2026-10-15", "input_cost_per_audio_token": 2.5e-06, "input_cost_per_token": 2.5e-06, "litellm_provider": "azure", @@ -6358,7 +6362,7 @@ "max_output_tokens": 2000, "mode": "audio_transcription", "output_cost_per_token": 1e-05, - "source": "https://management.azure.com/subscriptions/c873328e-b572-4770-8dff-aaeb6f1f0e79/providers/Microsoft.CognitiveServices/locations/eastus2/models?api-version=2024-10-01", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule", "supported_endpoints": [ "/v1/audio/transcriptions" ] @@ -11093,6 +11097,7 @@ "azure/us/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2.2e-05, "cache_read_input_token_cost": 2.75e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.00011, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", @@ -11102,6 +11107,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.00022, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -12556,6 +12562,7 @@ "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models" }, "azure_ai/jamba-instruct": { + "deprecation_date": "2025-03-01", "input_cost_per_token": 5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 70000, @@ -12563,6 +12570,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 7e-07, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true }, "azure_ai/kimi-k2.5": { diff --git a/ruff-strict.toml b/ruff-strict.toml index 899a8ff3af5..9e0b11a1af4 100644 --- a/ruff-strict.toml +++ b/ruff-strict.toml @@ -30,6 +30,10 @@ external = [ # grows over time; typing it concretely (`object`) broke that forwarding call outright โ€” # basedpyright turned every named param into a reportArgumentType error. Any is correct here. "litellm/proxy/guardrails/guardrail_hooks/alice/alice.py" = ["ANN401"] +# Same reason: `**kwargs` forwards verbatim to CustomGuardrail.__init__, and the lifecycle +# hook signatures inherit `Any` for `response` from CustomLogger, so narrowing them here +# would break the override rather than describe it. +"litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py" = ["ANN401"] [lint.mccabe] max-complexity = 15 diff --git a/scripts/generate_trace_types.py b/scripts/generate_trace_types.py index 39c93dabf49..52a61452711 100644 --- a/scripts/generate_trace_types.py +++ b/scripts/generate_trace_types.py @@ -34,7 +34,7 @@ class GeneratorConfig(BaseModel): options: tuple[str, ...] -def export(crate: str) -> Mapping[str, Mapping[str, JsonValue]]: +def export(crate: str, extra_args: tuple[str, ...] = ()) -> Mapping[str, Mapping[str, JsonValue]]: result: Final = subprocess.run( ( "cargo", @@ -48,6 +48,7 @@ def export(crate: str) -> Mapping[str, Mapping[str, JsonValue]]: f"export-{crate}-schema", "--features", "schema", + *(("--", *extra_args) if extra_args else ()), ), check=True, stdout=subprocess.PIPE, @@ -163,8 +164,9 @@ def main() -> int: sys.stderr.write(f"requires datamodel-code-generator=={config.version}\n") return 1 domain: Final = export("traces") + requests: Final = export("traces", ("--requests",)) clickhouse: Final = export("traces-clickhouse") - exported: Final = tuple(schema_files(domain, clickhouse)) + exported: Final = tuple(schema_files(domain, clickhouse, requests)) schema_results: Final = tuple(publish(path, content, args.check) for path, content in exported) schema_set_matches: Final = reconcile_schemas(frozenset(path for path, _ in exported), args.check) with TemporaryDirectory(prefix="trace-codegen-") as temporary: @@ -176,9 +178,11 @@ def main() -> int: directory, config, ) + request_models: Final = generate(requests, "requests", directory, config) python_results: Final = ( publish(GENERATED / "types.py", types.read_text(), args.check), publish(GENERATED / "models.py", models.read_text(), args.check), + publish(GENERATED / "requests.py", request_models.read_text(), args.check), ) return 0 if all((schema_set_matches, *schema_results, *python_results)) else 1 @@ -186,8 +190,9 @@ def main() -> int: def schema_files( domain: Mapping[str, Mapping[str, JsonValue]], clickhouse: Mapping[str, Mapping[str, JsonValue]], + requests: Mapping[str, Mapping[str, JsonValue]], ) -> Iterator[tuple[Path, str]]: - for crate, schemas in (("traces", domain), ("traces-clickhouse", clickhouse)): + for crate, schemas in (("traces", domain), ("traces-clickhouse", clickhouse), ("traces", requests)): for name, schema in schemas.items(): yield TOOLING / "schemas" / crate / f"{name}.json", json.dumps(schema, indent=2, sort_keys=True) + "\n" diff --git a/scripts/trace_codegen/README.md b/scripts/trace_codegen/README.md index f9071001c78..bb4f511a569 100644 --- a/scripts/trace_codegen/README.md +++ b/scripts/trace_codegen/README.md @@ -1,9 +1,13 @@ -Run `uv run scripts/generate_trace_types.py` from the repository root to export Rust schemas and regenerate the Python trace contracts. Run the same command with `--check` to compare fresh output with the committed schemas and Python files +Run `uv run scripts/generate_trace_types.py` from the repository root to export Rust schemas and regenerate the Python trace contracts. Run `uv run scripts/generate_trace_types.py --check` to compare fresh output with the committed schemas and Python files The script pins datamodel-code-generator in its inline dependency metadata. Rust uses the workspace's locked Schemars version through each owning crate's optional `schema` feature. Neither tool is a Python runtime dependency -Each crate exports its own roots using JSON Schema 2020-12. Request parameters use Schemars' deserialization contract. Trace views and query help use its serialization contract. Lens rows use their ClickHouse deserialization schemas, including quoted numbers and numeric boolean flags +The `litellm-traces` Rust request types own the generated request models in `litellm/rust_bridge/trace/generated/requests.py`. The GET routes bind their query parameters directly to the generated models. GET request types allow unknown fields because existing clients' unknown query parameters are ignored. The SQL body model forbids extra fields + +Each crate exports its own roots using JSON Schema 2020-12. Request parameters use Schemars' deserialization contract and carry only explicitly declared constraints. Their schemas skip the integer-bounds transform because it would add i64 bounds to `start_ms` and `end_ms`, narrowing what Python accepts, and replace `page_size`'s explicit 1..500 range with 0..65535. Trace views and query help use the serialization contract. Lens rows use their ClickHouse deserialization schemas, including quoted numbers and numeric boolean flags The templates preserve tuple conversion, immutable tuple defaults, and bounded `ReadOnly` TypedDict fields. Pydantic models use the generator's frozen-model option and each schema's extra-field policy. ClickHouse numeric schemas select bounded, normalized Python scalar types through schema metadata consumed by the model template +Changing the public request contract requires a separate behavior-change PR + Edit the owning Rust contract, schema annotation, or generation configuration, then regenerate. Never edit `litellm/rust_bridge/trace/generated/` manually. The SQL response envelope remains handwritten in `queries.py` diff --git a/scripts/trace_codegen/schemas/traces/TraceDetailRequest.json b/scripts/trace_codegen/schemas/traces/TraceDetailRequest.json new file mode 100644 index 00000000000..60cb76dfedc --- /dev/null +++ b/scripts/trace_codegen/schemas/traces/TraceDetailRequest.json @@ -0,0 +1,29 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "properties": { + "cursor": { + "default": null, + "maxLength": 512, + "type": [ + "string", + "null" + ] + }, + "page_size": { + "default": null, + "format": "uint16", + "maximum": 500, + "minimum": 1, + "type": [ + "integer", + "null" + ] + }, + "trace_ref": { + "default": "", + "type": "string" + } + }, + "title": "TraceDetailRequest", + "type": "object" +} diff --git a/scripts/trace_codegen/schemas/traces/TraceErrorPageRequest.json b/scripts/trace_codegen/schemas/traces/TraceErrorPageRequest.json new file mode 100644 index 00000000000..950dab8e83a --- /dev/null +++ b/scripts/trace_codegen/schemas/traces/TraceErrorPageRequest.json @@ -0,0 +1,19 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "properties": { + "cursor": { + "default": null, + "maxLength": 512, + "type": [ + "string", + "null" + ] + }, + "trace_ref": { + "default": "", + "type": "string" + } + }, + "title": "TraceErrorPageRequest", + "type": "object" +} diff --git a/scripts/trace_codegen/schemas/traces/TraceListRequest.json b/scripts/trace_codegen/schemas/traces/TraceListRequest.json new file mode 100644 index 00000000000..a19c3965c0b --- /dev/null +++ b/scripts/trace_codegen/schemas/traces/TraceListRequest.json @@ -0,0 +1,33 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "properties": { + "cursor": { + "default": null, + "maxLength": 512, + "type": [ + "string", + "null" + ] + }, + "end_ms": { + "default": null, + "description": "Window end, unix ms. Default: now", + "format": "int64", + "type": [ + "integer", + "null" + ] + }, + "start_ms": { + "default": null, + "description": "Window start, unix ms. Default: 24h ago", + "format": "int64", + "type": [ + "integer", + "null" + ] + } + }, + "title": "TraceListRequest", + "type": "object" +} diff --git a/scripts/trace_codegen/schemas/traces/TraceQueryRequest.json b/scripts/trace_codegen/schemas/traces/TraceQueryRequest.json new file mode 100644 index 00000000000..1361b89ad90 --- /dev/null +++ b/scripts/trace_codegen/schemas/traces/TraceQueryRequest.json @@ -0,0 +1,14 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "additionalProperties": false, + "properties": { + "sql": { + "type": "string" + } + }, + "required": [ + "sql" + ], + "title": "TraceQueryRequest", + "type": "object" +} diff --git a/scripts/trace_codegen/schemas/traces/TraceSpanRequest.json b/scripts/trace_codegen/schemas/traces/TraceSpanRequest.json new file mode 100644 index 00000000000..4f5bb8fff69 --- /dev/null +++ b/scripts/trace_codegen/schemas/traces/TraceSpanRequest.json @@ -0,0 +1,11 @@ +{ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "properties": { + "trace_ref": { + "default": "", + "type": "string" + } + }, + "title": "TraceSpanRequest", + "type": "object" +} diff --git a/tests/code_coverage_tests/test_e2e_changed_gate.py b/tests/code_coverage_tests/test_e2e_changed_gate.py index 758d579f67d..c4d12d1282c 100644 --- a/tests/code_coverage_tests/test_e2e_changed_gate.py +++ b/tests/code_coverage_tests/test_e2e_changed_gate.py @@ -8,10 +8,6 @@ from typing import Final import pytest GATE: Final = Path(__file__).resolve().parents[2] / ".github/e2e-stack/assert_tests_ran.py" -SECRETS_TO_ENV: Final = GATE.with_name("secrets_to_env.py") -SELECT_TESTS: Final = GATE.with_name("select_tests.py") -REDACT_OUTPUT: Final = GATE.with_name("redact_output.py") -CANARY: Final = ("tests/e2e/access_control/test_a.py", "tests/e2e/access_control/test_b.py") SELECTED: Final = ("tests/e2e/access_control/test_a.py", "tests/e2e/access_control/test_b.py") @@ -102,238 +98,6 @@ def test_one_passing_management_case_cannot_hide_a_missing_actor(tmp_path: Path, assert f"test_actor_subject_and_database_role[{omitted_role}]" in result.stdout -def test_short_values_are_written_without_masking_every_digit_in_the_log(tmp_path: Path) -> None: - env_path: Final = tmp_path / ".env" - - result: Final = subprocess.run( - [sys.executable, "-I", str(SECRETS_TO_ENV), str(env_path)], - input='{"FLAG": "1", "API_KEY": "sk-0123456789abcdef"}', - capture_output=True, - text=True, - env={**os.environ, "GITHUB_ACTIONS": "true"}, - ) - - assert result.returncode == 0, result.stderr - assert result.stdout == "::add-mask::sk-0123456789abcdef\n" - assert env_path.read_text() == "FLAG='1'\nAPI_KEY='sk-0123456789abcdef'\n" - - -def test_outside_actions_no_value_is_printed(tmp_path: Path) -> None: - env_path: Final = tmp_path / ".env" - local_env: Final = {key: value for key, value in os.environ.items() if key != "GITHUB_ACTIONS"} - - result: Final = subprocess.run( - [sys.executable, "-I", str(SECRETS_TO_ENV), str(env_path)], - input='{"FLAG": "1", "API_KEY": "sk-0123456789abcdef"}', - capture_output=True, - text=True, - env=local_env, - ) - - assert result.returncode == 0, result.stderr - assert result.stdout == "" - assert "sk-0123456789abcdef" not in result.stderr - assert env_path.read_text() == "FLAG='1'\nAPI_KEY='sk-0123456789abcdef'\n" - - -def redact_output(tmp_path: Path, values: tuple[str, ...], text: str) -> tuple[subprocess.CompletedProcess[str], Path]: - env_path: Final = tmp_path / ".env" - _ = env_path.write_text("".join(f"{name}='{value}'\n" for name, value in zip(("A", "B", "C"), values))) - stack_env: Final = tmp_path / "stack.env" - _ = stack_env.write_text("LITELLM_MASTER_KEY=sk-e2e-master0123\nREDIS_PORT=6379\n") - log: Final = tmp_path / "e2e-pass-1.log" - _ = log.write_text(text) - out_dir: Final = tmp_path / "redacted" - result: Final = subprocess.run( # test-quality-ok: standalone script that imports its sibling by script directory - [ - sys.executable, - str(REDACT_OUTPUT), - "--values", - str(env_path), - "--values", - str(stack_env), - "--out", - str(out_dir), - str(log), - ], - capture_output=True, - text=True, - ) - return result, out_dir / log.name - - -def test_redacted_output_hides_every_masked_value_and_keeps_the_rest(tmp_path: Path) -> None: - text: Final = ( - "FAILED key=sk-0123456789abcdef master=sk-e2e-master0123 flag=1 port=6379 message=Missing credentials\n" - ) - - result, redacted = redact_output(tmp_path, ("sk-0123456789abcdef", "1"), text) - - assert result.returncode == 0, result.stderr - assert redacted.read_text() == "FAILED key=*** master=*** flag=1 port=6379 message=Missing credentials\n" - assert (redacted.stat().st_mode & 0o777) == 0o600 - assert (tmp_path / "e2e-pass-1.log").read_text() == text - assert "sk-" not in result.stdout + result.stderr - - -def test_a_masked_value_that_prefixes_a_longer_one_leaves_no_tail(tmp_path: Path) -> None: - result, redacted = redact_output(tmp_path, ("sk-0123456789", "sk-0123456789abcdef"), "token sk-0123456789abcdef\n") - - assert result.returncode == 0, result.stderr - assert redacted.read_text() == "token ***\n" - - -def test_a_json_secret_is_hidden_field_by_field_however_it_is_escaped(tmp_path: Path) -> None: - credentials: Final = ( - '{"type": "service_account", "signing_key": "MIIEvAIBADANBgkqhkiG9w0BAQEFAASC\\n' - 'c2VjcmV0LWtleS1ib2R5LWxpbmUtdHdv\\n", "client_id": "104857600000000000001"}' - ) - text: Final = ( - "decoded MIIEvAIBADANBgkqhkiG9w0BAQEFAASC\n" - "c2VjcmV0LWtleS1ib2R5LWxpbmUtdHdv\n" - "escaped MIIEvAIBADANBgkqhkiG9w0BAQEFAASC\\nc2VjcmV0LWtleS1ib2R5LWxpbmUtdHdv\\n\n" - "twice MIIEvAIBADANBgkqhkiG9w0BAQEFAASC\\\\nc2VjcmV0LWtleS1ib2R5LWxpbmUtdHdv\n" - "client 104857600000000000001 status 403\n" - ) - - result, redacted = redact_output(tmp_path, (credentials,), text) - - assert result.returncode == 0, result.stderr - assert redacted.read_text() == "decoded ***\n***\nescaped ***\\n***\\n\ntwice ***\\\\n***\nclient *** status 403\n" - - -def test_a_secret_with_xml_special_characters_is_hidden_in_the_junit_file(tmp_path: Path) -> None: - text: Final = 'body p&ss<w"rd-1\n' - - result, redacted = redact_output(tmp_path, ('p&ssbody ***\n' - - -def select_tests(changed: tuple[str, ...]) -> tuple[str, ...]: - result: Final = subprocess.run( - [sys.executable, str(SELECT_TESTS), *CANARY], - input="".join(f"{path}\n" for path in changed), - capture_output=True, - text=True, - ) - assert result.returncode == 0, result.stderr - return tuple(result.stdout.split()) - - -@pytest.mark.parametrize( - ("changed", "expected"), - ( - (("tests/e2e/logging/test_datadog_e2e.py", "litellm/router.py"), ("tests/e2e/logging/test_datadog_e2e.py",)), - (("tests/e2e/ui/test_keys.py", "tests/e2e/claude_code/test_cli.py", "tests/e2e/load/test_burst.py"), ()), - (("tests/e2e/migrations/test_startup.py", "tests/e2e/migrations/test_recovery.py"), ()), - (("tests/e2e/batches/test_managed_files_enforcement_e2e.py",), ()), - (("tests/e2e/guardrails/test_presidio_masking_e2e.py",), ()), - (("tests/e2e/llm_translation/realtime/test_realtime_pipecat_audio_e2e.py",), ()), - (("tests/e2e/logging/test_otel_v2_langfuse_generation_output_e2e.py",), ()), - ( - ("tests/e2e/logging/test_team_langfuse_callback_e2e.py",), - ("tests/e2e/logging/test_team_langfuse_callback_e2e.py",), - ), - ( - ("tests/e2e/llm_translation/realtime/test_realtime_e2e.py",), - ("tests/e2e/llm_translation/realtime/test_realtime_e2e.py",), - ), - ( - ("tests/e2e/guardrails/test_bedrock_guardrail_e2e.py",), - ("tests/e2e/guardrails/test_bedrock_guardrail_e2e.py",), - ), - (("tests/e2e/logging/helpers.py", "docs/my-website/docs/index.md", "tests/e2e/AGENTS.md"), ()), - ( - ("tests/e2e/logging/test_datadog_e2e.py", "tests/e2e/logging/test_datadog_e2e.py"), - ("tests/e2e/logging/test_datadog_e2e.py",), - ), - ), -) -def test_changed_suite_files_are_selected_unless_the_stack_cannot_run_them( - changed: tuple[str, ...], expected: tuple[str, ...] -) -> None: - assert select_tests(changed) == expected - - -@pytest.mark.parametrize( - "harness_file", - ( - "tests/e2e/proxy_client.py", - "tests/e2e/conftest.py", - "tests/e2e/management/management_client.py", - "tests/e2e/management/jwt_actors.py", - "tests/e2e/management/conftest.py", - "tests/e2e/coverage_registry/management_cases.py", - "tests/e2e/pytest.ini", - "tests/e2e/gateway/stage_mirror_ci_config.yml", - ".github/e2e-stack/up.sh", - ".github/e2e-stack/start-idp.sh", - "tests/e2e/idp_realm.json", - ".github/workflows/test-e2e-changed.yml", - ), -) -def test_harness_changes_run_the_canary_suite(harness_file: str) -> None: - assert select_tests((harness_file, "litellm/router.py")) == CANARY - - -def test_a_changed_canary_file_is_selected_once_alongside_a_harness_change() -> None: - assert select_tests((CANARY[1], "tests/e2e/proxy_client.py")) == CANARY - - -def test_dedicated_migration_tests_do_not_suppress_shared_harness_canaries() -> None: - assert select_tests(("tests/e2e/migrations/test_startup.py", "tests/e2e/conftest.py")) == CANARY - - -def test_the_canary_joins_directly_selected_files_in_sorted_order() -> None: - assert select_tests(("tests/e2e/logging/test_datadog_e2e.py", ".github/e2e-stack/up.sh")) == ( - *CANARY, - "tests/e2e/logging/test_datadog_e2e.py", - ) - - -def test_a_harness_unit_test_change_runs_itself_and_the_canary() -> None: - assert select_tests(("tests/e2e/test_proxy_client.py",)) == (*CANARY, "tests/e2e/test_proxy_client.py") - - -def test_a_canary_argument_the_shell_never_expanded_fails_the_selector() -> None: - result: Final = subprocess.run( - [sys.executable, str(SELECT_TESTS), "tests/e2e/access_control/test_*.py"], - input="tests/e2e/proxy_client.py\n", - capture_output=True, - text=True, - ) - - assert result.returncode == 1 - assert "tests/e2e/access_control/test_*.py" in result.stderr - assert result.stdout == "" - - -@pytest.mark.parametrize( - ("secrets", "offender", "unprintable"), - ( - ('{"AWS_ACCESS_KEY_ID": "AKIAEXAMPLE", "BAD-NAME": "shibboleth"}', "BAD-NAME", "shibboleth"), - ("""{"AWS_SECRET_ACCESS_KEY": "quote'shibboleth"}""", "AWS_SECRET_ACCESS_KEY", "shibboleth"), - ('{"DD_API_KEY": "line\\nshibboleth"}', "DD_API_KEY", "shibboleth"), - ), -) -def test_an_unusable_secret_is_named_without_printing_its_value( - tmp_path: Path, secrets: str, offender: str, unprintable: str -) -> None: - env_path: Final = tmp_path / ".env" - - result: Final = subprocess.run( - [sys.executable, str(SECRETS_TO_ENV), str(env_path)], input=secrets, capture_output=True, text=True - ) - - assert result.returncode == 1 - assert offender in result.stderr - assert unprintable not in result.stderr - assert result.stdout == "" - assert not env_path.exists() - - @pytest.mark.parametrize("phase", ("setup", "call", "teardown")) @pytest.mark.parametrize("required_count", ("1", "4")) def test_oauth_failure_diagnostics_do_not_publish_private_payloads( diff --git a/tests/e2e/AGENTS.md b/tests/e2e/AGENTS.md index 920ca8b02a3..3fd13439fd8 100644 --- a/tests/e2e/AGENTS.md +++ b/tests/e2e/AGENTS.md @@ -125,7 +125,7 @@ E2E_FIXTURE_MODE=replay E2E_FIXTURE_DIR=/tmp/e2e-fixtures E2E_RESET_SPEND_LOGS=1 Point the proxy at bogus provider credentials for the replay run and it still has to pass: that is the whole proof that nothing left the process. Bundles are never committed. `tests/e2e/.fixtures` is gitignored because a bundle holds verbatim provider response bodies and hard-fails after seven days. CI records and replays this lane on a schedule in `.github/workflows/e2e_record_replay.yml`, publishing the bundle as a private `e2e-fixtures-bundle` artifact instead of committing it, selecting the tests with the `@pytest.mark.replayable` marker, and proving the bogus-credentials replay hermetic by counting provider egress with `.github/scripts/e2e_egress_sentinel.py` -Current limits: Bedrock cannot be mounted in record or replay (SigV4 signs the Host header, so a rewritten api_base fails signature verification); a test that needs to observe the Converse body registers its own `LiveEdge` with `provider_edge_bedrock.bedrock_signer` re-signing the forwarded request, and carries the `provider_edge_host` opt-in marker because the gateway must reach the pytest host, which the Buildkite ephemeral stack cannot (the GitHub changed-e2e lane, whose gateways run on the runner, sets `E2E_PROVIDER_EDGE_HOST_REACHABLE`). Deployments baked into the proxy's config file cannot be edge-wired (only `/model/new` registrations can carry the edge api_base), and a file upload routed by `custom_llm_provider` through the proxy's `files_settings` block never passes a deployment at all, so the batches `model_param` and `provider_fallback` scenarios keep uploading live in every mode +Current limits: Bedrock cannot be mounted in record or replay (SigV4 signs the Host header, so a rewritten api_base fails signature verification); a test that needs to observe the Converse body registers its own `LiveEdge` with `provider_edge_bedrock.bedrock_signer` re-signing the forwarded request, and carries the `provider_edge_host` opt-in marker because the gateway must reach the pytest host, which the Buildkite ephemeral stack cannot, so those tests are deselected unless `E2E_PROVIDER_EDGE_HOST_REACHABLE` is set, as on a local run whose gateway can reach the pytest host. Deployments baked into the proxy's config file cannot be edge-wired (only `/model/new` registrations can carry the edge api_base), and a file upload routed by `custom_llm_provider` through the proxy's `files_settings` block never passes a deployment at all, so the batches `model_param` and `provider_fallback` scenarios keep uploading live in every mode ## Typing diff --git a/tests/e2e/CONTRIBUTING.md b/tests/e2e/CONTRIBUTING.md index fb2cf2dfa24..682bf41d7d9 100644 --- a/tests/e2e/CONTRIBUTING.md +++ b/tests/e2e/CONTRIBUTING.md @@ -73,7 +73,7 @@ The suites run against a live proxy, so bring one up first by running the litell uv run pytest tests/e2e/other/test_jwt_auth_e2e.py tests/e2e/management/test_jwt_management_e2e.py --reruns 0 -v ``` - Buildkite runs this suite against a Keycloak deployed beside the ephemeral stack by project-releaser. It fetches the realm from the test-runner revision even when it reuses a gateway image from another commit. The GitHub Actions changed-test stack starts the same digest-pinned Keycloak through `.github/e2e-stack/start-idp.sh`, imports the checked-out realm, and exports the IdP URL and credentials in `stack.env`. Both runners configure issuer/audience validation and store the realm, keys and users in a separate schema in the stack's PostgreSQL, so replacing Keycloak preserves token validity. Both wait for realm discovery before running tests. Losing the whole ephemeral database invalidates the stack. Keycloak skips imports into an existing realm, so changes to the realm export require a fresh stack (or deliberately replacing the local data volume). A stack without it fails the JWT tests rather than skipping them + Buildkite runs this suite against a Keycloak deployed beside the ephemeral stack by project-releaser. It fetches the realm from the test-runner revision even when it reuses a gateway image from another commit. The runner configures issuer/audience validation and stores the realm, keys and users in a separate schema in the stack's PostgreSQL, so replacing Keycloak preserves token validity. It waits for realm discovery before running tests. Losing the whole ephemeral database invalidates the stack. Keycloak skips imports into an existing realm, so changes to the realm export require a fresh stack (or deliberately replacing the local data volume). A stack without it fails the JWT tests rather than skipping them 4. Run a suite against it; the harness reads `LITELLM_PROXY_URL` (default `http://localhost:4000`). The suites' client dependencies (the provider SDKs, websockets) live in the `e2e-dev` dependency group; `make bootstrap` installs it, and naming the group on the run keeps the command working from any environment state: @@ -103,20 +103,6 @@ Alternatively, start the proxy with `STORE_PROMPTS_IN_SPEND_LOGS=true litellm -- A couple of logging destinations are configured on the proxy rather than by the test. The Weave tests scope their callback to the key they create, but litellm builds the `weave_otel` logger from `WANDB_API_KEY` and `WANDB_PROJECT_ID` before it applies the per-key vars, so the proxy needs both in its own environment or the key-scoped callback never initializes and nothing ships -### The pull request check - -Every same-repository PR that adds, modifies, or renames a `tests/e2e/**/test_*.py` file runs those changed files three times. A change to the harness itself, meaning a root-level `tests/e2e/*.py` file or `pytest.ini`, `tests/e2e/gateway/`, `.github/e2e-stack/`, or the workflow, also runs the `access_control` suite and both JWT suites as canaries, because those files have no test of their own that exercises the stack. `.github/e2e-stack/select_tests.py` applies both rules. The stack config at `tests/e2e/gateway/stage_mirror_ci_config.yml` must declare every model the selected suites use; a missing one shows up as a failed test id in the public log. The suite's own single rerun for network errors and 5xx responses (see `pytest.ini`) applies on every pass, so a transport blip does not fail the check while a race inside a test still does. The stage-mirror stack has a control-plane backend, two gateways behind nginx, Postgres, Keycloak, Jaeger, and TLS cluster-mode Valkey. Realm-only edits also trigger these canaries. The stack exports every gateway address in `LITELLM_PROXY_REPLICA_URLS`, so model registration waits until each gateway lists the new model rather than whichever one the load balancer answered from. The Buildkite PR stack exports its two gateway pods the same way and, because those pods sit behind one router base that also fronts the backend, names that base in `LITELLM_CONTROL_PLANE_REPLICA_URLS` so management read-backs poll the plane that serves them instead of the gateway pods, which trim management routes at startup and answer them 404. Documentation, deleted-file, and application-only changes do not start the stack or request environment approval. The `ui/`, `claude_code/`, `load/`, and `secret_manager/` directories, `batches/test_managed_files_enforcement_e2e.py`, `llm_translation/realtime/test_realtime_pipecat_audio_e2e.py`, and `guardrails/test_presidio_masking_e2e.py` remain outside this check because they use separate tooling or need a differently configured stack: the pipecat audio suite skips itself at import time unless the NLTK `punkt_tab` data is installed, and the presidio suite fails without the analyzer and anonymizer services this stack does not start. `logging/test_otel_v2_langfuse_generation_output_e2e.py` is marked `otel_v2` and deselects itself unless `E2E_OTEL_V2` is set, because it needs a gateway booted with `LITELLM_OTEL_V2=true` and Langfuse credentials, neither of which this stack provides, so run it with `E2E_OTEL_V2=1` against a local OTel v2 proxy. The Redis chaos test under `load/` needs a proxy it can pause the Redis of on the same host (`gateway/redis_chaos_ci_config.yml`), which `.github/workflows/test-e2e-redis-chaos.yml` boots, and which the Buildkite `e2e-redis-chaos` step in project-releaser runs co-located with Postgres and Valkey in one pod; it is deselected unless `E2E_REDIS_CHAOS` is set. The `secret_manager/` lanes each need a proxy configured against their own secret manager (see Secret manager lanes below) - -Every selected file must execute at least one passing test in each pass, and any test failure, collection error, or entirely skipped or deselected file fails the check. A file whose tests are all marked skip therefore cannot pass this check, so unskip at least one of them, or add the file to `UNSUPPORTED` in `select_tests.py` with the reason, before changing one. A failed pass stops the run. The public log prints pytest's one-line summary for each pass, including the rerun count, and names each failed or errored test as `classname::name`, so a retried network error or a failing test is visible without the raw output. The final `e2e-changed-tests` job succeeds only when no supported test files changed or the approved run completed all three passes. Fork PRs with selected tests fail this gate until a maintainer brings the reviewed change onto a same-repository branch - -Repository admins must require the `e2e-changed-tests` status check for merging and configure the `e2e-changed` environment with required reviewers, self-review disabled, and admin bypass disabled. Each push cancels the previous run; a new run that selects tests needs a fresh approval. Reviewers must inspect the entire executable PR diff, including application code, dependencies, tests, and workflow helpers, before approving the exact revision. Approved code executes with provider credentials, so environment approval is a trust decision about that code - -Credentials come from the existing AWS Secrets Manager secrets in us-east-1, `litellm-e2e-changed-provider-keys` and `litellm-e2e-changed-license`. The OIDC role must trust only `repo:BerriAI/litellm:environment:e2e-changed` with audience `sts.amazonaws.com` and have read access only to these secrets. The short-lived reader credentials are scoped to the fetch step. Provider credentials must cover the selected suites, including Datadog credentials when logging or MCP tests need them; missing credentials fail the run. `up.sh` refuses to start without `DD_API_KEY`, because the stack's gateway config enables the Datadog callback for every run and a gateway booted without the key fails readiness. Keep provider credentials dedicated to this lane with only the permissions those tests need - -Fetched values of eight characters or more are masked before use, while shorter values such as flags stay unmasked because masking a one-character value would blank every matching digit in the log, and credential files and raw output are private to the runner. Public logs contain selected file names, counts, pytest's summary line, failed test ids, and pass status; raw pytest output, reports, and stack logs are not uploaded or printed. The workflow removes them and the credential files during cleanup. To diagnose a failed pass, reproduce the selected files locally with the appropriate credentials and inspect the local logs - -To reproduce the CI topology on a dedicated machine, `bash .github/e2e-stack/up.sh` reads `tests/e2e/.env`, writes `stack.env` under `${E2E_STACK_DIR:-/tmp/litellm-e2e-stack}`, and `bash .github/e2e-stack/down.sh` stops it. Keep this directory private and remove its credential files and logs after use - ### Secret manager lanes `key_management_system` is global to the proxy, so the `secret_manager/` tests run once per backend, each against its own proxy. The backends are `hashicorp_vault` and `cyberark` (CyberArk Conjur). `E2E_SECRET_MANAGER` opts in and names the backend (a key of `secret_backends.BACKENDS`). The proxy boots from `gateway/secret_manager__ci_config.yml`, and the tests reach the same manager through that backend's `SecretStore`. The managers are enterprise features, so the proxy needs a license. `secret_manager/backend.sh` runs any backend in Docker and writes its env, so every lane runs the same way locally: @@ -319,8 +305,7 @@ same-repository branch for their verification. The workflow's path-filtered check is not configured here as a globally required branch-protection check. Provision `E2E_LINEAR_STORAGE_STATE_B64` as a secret there and retain the existing E2E license/AWS role configuration. A missing or expired session fails the job; collection, deselection and skips are not passes. -The generic changed-test job excludes this file because it requires an owned -proxy and consent UI. No LLM call is needed +No LLM call is needed Coverage remains limited to authorization-code OAuth over HTTP. M2M, OBO, PKCE passthrough, static/BYOK, ID-JAG, forwarding, SigV4 and stdio are outside this diff --git a/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py b/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py index c20b73c0d63..81470c21d51 100644 --- a/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py +++ b/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py @@ -38,6 +38,7 @@ pytestmark = [pytest.mark.e2e, pytest.mark.mcp_oauth_live, pytest.mark.provider_ class OAuthMetadata(BaseModel): + issuer: str authorization_endpoint: str token_endpoint: str registration_endpoint: str @@ -135,6 +136,7 @@ class TestMcpOauthHappyPath: auth_type="oauth2", oauth2_flow="authorization_code", per_server_oauth_discovery=route == "explicit_header_jwt", + issuer=metadata.issuer if metadata else None, authorization_url=metadata.authorization_endpoint if metadata else None, token_url=metadata.token_endpoint if metadata else None, registration_url=metadata.registration_endpoint if metadata else None, diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 90fe04ae896..ccb5cd35280 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -658,6 +658,7 @@ class McpServerCreateBody(BaseModel): auth_type: str | None = None oauth2_flow: Literal["client_credentials", "authorization_code"] | None = None per_server_oauth_discovery: bool | None = None + issuer: str | None = None authorization_url: str | None = None token_url: str | None = None registration_url: str | None = None diff --git a/tests/integration/README.md b/tests/integration/README.md index fe31e215f8a..c53220cc895 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -2,7 +2,7 @@ These tests exercise a running gateway, PostgreSQL and Redis with an owned local upstream. CircleCI owns this suite. Tests are grouped by behavior, with no automatic test retries or fallback to paid provider calls -The ROI database contracts in `database/test_roi_observed.py` run in the GitHub Actions `roi-database` Postgres shard and upload coverage on each PR. They own temporary databases and script only the external provider transport. `GITHUB_FILES` in `run.py` assigns these files to GitHub Actions and excludes them from the CircleCI selection +The ROI database contracts in `database/test_roi_observed.py` run as plain pytest outside `run_integration.sh` in CircleCI's `roi-database` Postgres job. The job uploads coverage with the `roi-postgres` flag. These contracts own temporary databases and script only the external provider transport The `cost` group is driven by `cost_tracking_cases.json`, which contains the cost map, literal requests, literal provider responses and expected accounting values. Each case has a name, contract ID, cost-map model, optional deployment overrides, request body, tagged response and exact or recount expectations. Request bodies use `$MODEL` for the registered proxy model, while responses use `$REQUEST_ID` for the per-run scenario ID. To add a case, add a cost-map entry when the model is new, add the request body and exact provider response data, and add hand-computed expected values. The upstream serves each stored response for any path under `/`, while the test-owned cost map is served over loopback through `LITELLM_MODEL_COST_MAP_URL` @@ -14,7 +14,7 @@ The generated lifecycle models use 20 examples, eight steps, generation and shri Reuse the existing canned provider handlers through `_support/upstream.py`. It rejects internal request fields and exposes actual received requests for independent assertions. Register every created resource for cleanup immediately, keep expected values independent of production calculations, and assert readback plus the runtime effect of a change -The CircleCI workflow starts its own database and Redis, restricts test-phase egress to its owned services and writes JUnit plus an executed-node manifest. Missing setup, failed cleanup or a selected test with neither a passed call nor a skip fail qualification. Skipped nodes are listed under `skipped` in `execution.json`, so the skip reasons double as the open bug list. GitHub Actions runs only the explicit `GITHUB_FILES` set in `run.py` +The CircleCI workflow starts its own database and Redis, restricts test-phase egress to its owned services and writes JUnit plus an executed-node manifest. Missing setup, failed cleanup or a selected test with neither a passed call nor a skip fail qualification. Skipped nodes are listed under `skipped` in `execution.json`, so the skip reasons double as the open bug list. `GITHUB_FILES` in `run.py` lists the files excluded from the `run_integration.sh` selection There is no per-node manifest. A positional argument is a file of the group or a pytest node id inside one (`path::test[param]`), so one cell of a parametrized file can run alone. The runner fails only when pytest fails, when collection errors, or when a selected file collects zero tests. Older tests still carry `@pytest.mark.covers(...)` decorators; the marker stays registered so they collect, but the IDs are not checked against anything and new tests should not use it. The GitHub Actions coverage census reads the `GROUPS` literal in `run.py` and treats every `tests/integration//test_*.py` file in a scheduled group as owned by CircleCI diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 3d76c7881e1..3fa7f4b0333 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -230,6 +230,60 @@ def owned_proxy_process( _stop(process) +def _is_ready(client: httpx.Client) -> bool: + try: + return client.get("/health/readiness", timeout=2).status_code == 200 + except httpx.TransportError: + return False + + +def refused_boot_log( + gateway: Gateway, + directory: Path, + overrides: Mapping[str, str], + *, + config: Path | None = None, +) -> str: + """Start the proxy and return its log once it exits non-zero instead of becoming ready.""" + root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3]) + environment: Final = { + **os.environ, + **proxy_database_environment(), + "LITELLM_MASTER_KEY": gateway.key, + "LITELLM_SALT_KEY": os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"), + "STORE_MODEL_IN_DB": "True", + **overrides, + } + output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(directory))) + output.mkdir(parents=True, exist_ok=True) + command: Final = ( + sys.executable, + "-m", + "integration._support.proxy", + "--config", + str(config or "tests/integration/proxy_config.yaml"), + "--host", + "127.0.0.1", + "--num_workers", + "1", + *DB_PUSH, + ) + launch: Final = _launch(command, root, environment, output) + try: + with httpx.Client(base_url=f"http://127.0.0.1:{launch.port}", timeout=15, trust_env=False) as client: + deadline: Final = time.monotonic() + 70 + while launch.process.poll() is None: + assert not _is_ready(client), ( + f"Proxy became ready instead of refusing to boot:\n{launch.log.read_text()}" + ) + assert time.monotonic() < deadline, "Proxy neither exited nor became ready within the deadline" + time.sleep(0.1) + assert launch.process.returncode != 0, f"Proxy exited 0 instead of refusing to boot:\n{launch.log.read_text()}" + return launch.log.read_text() + finally: + _stop(launch.process) + + _UPSTREAM_READY_SECONDS: Final = 60 diff --git a/tests/integration/authorization/test_team_scoped_models.py b/tests/integration/authorization/test_team_scoped_models.py index 2b347a6f4f8..aea6b9e4ba4 100644 --- a/tests/integration/authorization/test_team_scoped_models.py +++ b/tests/integration/authorization/test_team_scoped_models.py @@ -66,7 +66,7 @@ def test_a_team_model_is_listed_and_served_only_for_keys_of_its_team(gateway: Ga def _v2_team_public_names(gateway: Gateway, key: str, model: str) -> list[JsonValue]: - response: Final = gateway.request("GET", "/v2/model/info", key=key, params={"model_name": model}) + response: Final = gateway.request("GET", "/v2/model/info", key=key, params={"model": model}) assert response.status_code == 200, response.text return [entry["model_info"].get("team_public_model_name") for entry in response.json()["data"]] diff --git a/tests/integration/configuration/test_fips_mode_boot.py b/tests/integration/configuration/test_fips_mode_boot.py new file mode 100644 index 00000000000..80904c600e1 --- /dev/null +++ b/tests/integration/configuration/test_fips_mode_boot.py @@ -0,0 +1,62 @@ +"""LITELLM_FIPS_MODE is a boot gate: the proxy refuses to serve unless the process really enforces FIPS. + +Every leg launches the real proxy binary against the suite's Postgres and asserts on what an operator sees: +exit status and the refusal text in the log. Nothing is patched inside the proxy. +""" + +import hashlib +from pathlib import Path +from typing import Final + +import pytest +import yaml + +from tests.integration._support.client import Gateway +from tests.integration._support.process import owned_proxy, refused_boot_log + +REFUSAL: Final = "LiteLLM proxy refused to start" + + +def _this_python_enforces_fips() -> bool: + try: + hashlib.md5(b"probe", usedforsecurity=True) + except ValueError: + return True + return False + + +def test_fips_mode_refuses_to_serve_when_this_python_does_not_enforce_fips(gateway: Gateway, tmp_path: Path) -> None: + if _this_python_enforces_fips(): + pytest.skip("Runner OpenSSL enforces FIPS, so this leg cannot observe the non-enforcing refusal") + log: Final = refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true"}) + assert REFUSAL in log, log + assert "LITELLM_FIPS_MODE" in log and "does not enforce FIPS" in log, log + + +@pytest.mark.parametrize("source", ("environment", "config")) +def test_fips_mode_refuses_to_serve_with_tls_verification_disabled( + gateway: Gateway, tmp_path: Path, source: str +) -> None: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + path: Final = tmp_path / "ssl_verify_off.yaml" + settings: Final = {**config.get("litellm_settings", {}), "ssl_verify": False} + path.write_text(yaml.safe_dump({**config, "litellm_settings": settings})) + log: Final = ( + refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true", "SSL_VERIFY": "false"}) + if source == "environment" + else refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true"}, config=path) + ) + assert REFUSAL in log, log + assert "TLS certificate verification is disabled" in log, log + assert ("SSL_VERIFY" if source == "environment" else "litellm_settings.ssl_verify") in log, log + + +def test_fips_mode_refuses_to_serve_on_a_value_that_is_not_a_boolean(gateway: Gateway, tmp_path: Path) -> None: + log: Final = refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "enforced"}) + assert REFUSAL in log, log + assert "LITELLM_FIPS_MODE=enforced" in log and "true or false" in log, log + + +def test_fips_mode_off_serves_even_with_tls_verification_disabled(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy(gateway, tmp_path, {"LITELLM_FIPS_MODE": "false", "SSL_VERIFY": "false"}) as candidate: + assert candidate.client.get("/health/readiness").status_code == 200 diff --git a/tests/integration/cost_calculation/cost_tracking_cases.json b/tests/integration/cost_calculation/cost_tracking_cases.json index 8242c8683b5..f03893c4244 100644 --- a/tests/integration/cost_calculation/cost_tracking_cases.json +++ b/tests/integration/cost_calculation/cost_tracking_cases.json @@ -30790,7 +30790,8 @@ "role": "user", "content": "proxy behaviour probe" } - ] + ], + "max_tokens": 412 }, "response": { "content_type": "application/json", diff --git a/tests/integration/cost_calculation/test_cost_tracking.py b/tests/integration/cost_calculation/test_cost_tracking.py index cfe0a99ef4b..379b2c13f5b 100644 --- a/tests/integration/cost_calculation/test_cost_tracking.py +++ b/tests/integration/cost_calculation/test_cost_tracking.py @@ -281,6 +281,8 @@ def test_case_bills_expected_cost(gateway: Gateway, case: CostTrackingTestCase) { "path": f"/{scenario_id}/v1/decisions", "authorization": "Bearer sk-scripted-provider", + "method": "POST", + "api_key": "", "body": { "model": "pplx-decider-v1-27b", "state": {"source": "cost-tracking"}, diff --git a/tests/integration/management/test_vector_store_file_list_managed_ids.py b/tests/integration/management/test_vector_store_file_list_managed_ids.py index 7ad68e884e6..c6124519c9e 100644 --- a/tests/integration/management/test_vector_store_file_list_managed_ids.py +++ b/tests/integration/management/test_vector_store_file_list_managed_ids.py @@ -602,14 +602,23 @@ def test_two_different_after_values_forward_the_last_one(gateway: Gateway) -> No assert _ids(page) == (managed_b,), page -@pytest.mark.parametrize("status", (401, 404, 500)) -def test_provider_errors_reach_the_caller_and_other_models_keep_mapping(gateway: Gateway, status: int) -> None: +@pytest.mark.parametrize( + ("status", "expected_error_type"), + ((401, "authentication_error"), (404, "invalid_request_error"), (500, "internal_server_error")), +) +def test_provider_errors_reach_the_caller_and_other_models_keep_mapping( + gateway: Gateway, status: int, expected_error_type: str +) -> None: message: Final = f"provider refused listing {uuid.uuid4().hex[:8]}" with _rig(gateway, listing=_error_listing(status, message)) as failing, _rig(gateway, "a.txt") as healthy: member: Final = _member(failing.scenario, failing.model, healthy.model) managed_a: Final = healthy.upload(member.key, "a.txt") failed: Final = failing.list(member.key, {"model": failing.model}) - assert _json(failed) == _provider_error(status, message), failed.text + assert failed.status_code == status, failed.text + error: Final = object_value(_json(failed)["error"]) + assert message in string_value(error["message"]), failed.text + assert error["code"] == str(status), failed.text + assert error["type"] == expected_error_type, failed.text assert len(failing.list_requests()) == 1 assert _ids(healthy.listed(member.key)) == (managed_a,) liveliness: Final = gateway.request("GET", "/health/liveliness") diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index d377afb206c..1df4427bb4e 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -1,12 +1,20 @@ +import asyncio import json import os +import re import signal import socket +import threading import uuid +from collections.abc import Callable, Iterator, Mapping from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from dataclasses import dataclass from pathlib import Path +from types import MappingProxyType from typing import Final +import anthropic import httpx import psutil import pytest @@ -14,9 +22,10 @@ import yaml from integration._support.client import Gateway, eventually, object_value from integration._support.database import read_rows from integration._support.mcp import mcp_peer, register_mcp, tool_names -from integration._support.process import group_members, owned_proxy, owned_proxy_process -from integration._support.wire import Reply, Request, wire_server +from integration._support.process import OwnedProxy, group_members, owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server from openai import AsyncOpenAI, OpenAI +from pydantic import JsonValue @pytest.mark.covers("other.observability.guardrails.rewrite_reaches_correct_anthropic_positions") @@ -1796,3 +1805,739 @@ def test_responses_pre_call_denial_stream_survives_worker_kill(gateway: Gateway, for response in responses: assert response.status_code == 200, response.text assert response.headers["content-type"].startswith("text/event-stream"), response.text + + +_TOKEN: Final = re.compile(rb"token-[0-9a-f]{32}-\d+") + + +def _secret_for(request: Request) -> str: + token: Final = _TOKEN.search(request.body) + assert token is not None, request.body + return "synthetic-leaked-secret-" + token.group().decode() + + +def _chat_frame(identity: str, choices: tuple[dict[str, JsonValue], ...]) -> bytes: + payload: Final = { + "id": identity, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": list(choices), + } + return b"data: " + json.dumps(payload).encode() + b"\n\n" + + +def _chat_choice(index: int, delta: dict[str, JsonValue], finish: str | None = None) -> dict[str, JsonValue]: + return {"index": index, "delta": delta, "finish_reason": finish} + + +def _chat_stream_frames(secret: str, shape: str) -> tuple[bytes, ...]: + identity: Final = "chatcmpl-" + uuid.uuid4().hex + tool_call: Final = { + "index": 0, + "id": "call_" + identity, + "type": "function", + "function": {"name": "lookup", "arguments": json.dumps({"query": secret})}, + } + released: Final = { + "text": (_chat_choice(0, {"role": "assistant", "content": secret}),), + "empty": (_chat_choice(0, {"role": "assistant", "content": ""}),), + "tool_call": (_chat_choice(0, {"role": "assistant", "tool_calls": [tool_call]}),), + "two_choices": ( + _chat_choice(0, {"role": "assistant", "content": secret + "-first"}), + _chat_choice(1, {"role": "assistant", "content": secret + "-second"}), + ), + }[shape] + finish: Final = "tool_calls" if shape == "tool_call" else "stop" + tail: Final = tuple(_chat_choice(int(str(choice["index"])), {"content": " tail"}, finish) for choice in released) + return (_chat_frame(identity, released), _chat_frame(identity, tail), b"data: [DONE]\n\n") + + +def _chat_completion_body(secret: str) -> bytes: + return json.dumps( + { + "id": "chatcmpl-" + uuid.uuid4().hex, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": secret}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 9, "completion_tokens": 4, "total_tokens": 13}, + } + ).encode() + + +def _responses_tool_call_frames(secret: str, *, with_text: bool) -> tuple[bytes, ...]: + identity: Final = "resp_" + uuid.uuid4().hex + arguments: Final = json.dumps({"query": secret}) + pending: Final = {"type": "function_call", "id": "fc_" + identity, "call_id": "call_" + identity, "name": "lookup"} + finished: Final = {**pending, "arguments": arguments, "status": "completed"} + envelope: Final = {"id": identity, "object": "response", "created_at": 1, "model": "gpt-4o-mini", "output": []} + message: Final = { + "type": "message", + "id": "msg_" + identity, + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": secret, "annotations": []}], + } + text_delta: Final = { + "type": "response.output_text.delta", + "item_id": "msg_" + identity, + "output_index": 0, + "content_index": 0, + "delta": secret, + } + text_events: Final = (text_delta,) if with_text else () + tool_index: Final = len(text_events) + output: Final = [*((message,) if with_text else ()), finished] + events: Final = ( + {"type": "response.created", "response": {**envelope, "status": "in_progress"}}, + *text_events, + { + "type": "response.output_item.added", + "output_index": tool_index, + "item": {**pending, "arguments": "", "status": "in_progress"}, + }, + { + "type": "response.function_call_arguments.delta", + "item_id": "fc_" + identity, + "output_index": tool_index, + "delta": arguments, + }, + {"type": "response.output_item.done", "output_index": tool_index, "item": finished}, + {"type": "response.completed", "response": {**envelope, "status": "completed", "output": output}}, + ) + encoded: Final = tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events) + released_through: Final = len(events) - 1 if with_text else 3 + return (b"".join(encoded[:released_through]), b"".join(encoded[released_through:])) + + +def _responses_stream_frames(secret: str) -> tuple[bytes, ...]: + identity: Final = "resp_" + uuid.uuid4().hex + completed: Final = { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": "msg_" + identity, + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": secret, "annotations": []}], + } + ], + "usage": { + "input_tokens": 11, + "output_tokens": 4, + "total_tokens": 15, + "input_tokens_details": {"cached_tokens": 0}, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + } + events: Final = ( + {"type": "response.created", "response": {**completed, "status": "in_progress", "output": [], "usage": None}}, + { + "type": "response.output_text.delta", + "item_id": "msg_" + identity, + "output_index": 0, + "content_index": 0, + "delta": secret, + }, + {"type": "response.completed", "response": completed}, + ) + encoded: Final = tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events) + return (encoded[0] + encoded[1], encoded[2]) + + +def _gemini_stream_frames(secret: str) -> tuple[bytes, ...]: + def frame(text: str, finish: str | None) -> bytes: + candidate: Final = { + "content": {"parts": [{"text": text}], "role": "model"}, + "index": 0, + **({"finishReason": finish} if finish else {}), + } + payload: Final = { + "candidates": [candidate], + "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 5, "totalTokenCount": 15}, + "modelVersion": "gemini-2.5-flash", + } + return b"data: " + json.dumps(payload).encode() + b"\r\n\r\n" + + return (frame(secret, None), frame(" tail", "STOP")) + + +def _scripted_provider(gate: threading.Event | None, pause: float, shape: str) -> Callable[[Request], Reply]: + def provider(request: Request) -> Reply: + secret: Final = _secret_for(request) + path: Final = request.target.split("?")[0] + if path.endswith("/chat/completions") and not json.loads(request.body).get("stream"): + return Reply(body=_chat_completion_body(secret)) + frames: Final = ( + _gemini_stream_frames(secret) + if "streamGenerateContent" in path + else ( + _responses_tool_call_frames(secret, with_text=shape == "text_then_tool_call") + if shape in ("tool_call", "text_then_tool_call") + else _responses_stream_frames(secret) + ) + if path.endswith("/responses") + else _chat_stream_frames(secret, shape) + ) + return Reply(content_type="text/event-stream", chunks=frames, gate_after_first=gate, pause_between_chunks=pause) + + return provider + + +def _allowing_guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + +def _failing_response_scans(reply: Reply) -> Callable[[Request], Reply]: + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + if json.loads(request.body)["input_type"] == "response": + return reply + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + return guardrail + + +def _post_call_config( + tmp_path: Path, identity: str, policy_url: str, params: Mapping[str, JsonValue], default_on: bool +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "post_call", + "default_on": default_on, + "api_base": policy_url, + "api_key": "synthetic-guardrail-key", + **params, + }, + } + ] + path: Final = tmp_path / f"{identity}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +class _ScanLog: + def __init__(self, policy: Wire) -> None: + self.policy: Final = policy + self.seen: tuple[dict[str, JsonValue], ...] = () + + def response_scans(self, secret: str) -> tuple[dict[str, JsonValue], ...]: + self.seen = (*self.seen, *(object_value(json.loads(request.body)) for request in self.policy.drain())) + return tuple(body for body in self.seen if body["input_type"] == "response" and secret in json.dumps(body)) + + +@dataclass(frozen=True, slots=True) +class _DisconnectRig: + owned: OwnedProxy + model: str + gemini: str + scans: _ScanLog + identity: str + gate: threading.Event + upstream: Wire + + @property + def candidate(self) -> Gateway: + return self.owned.gateway + + def token(self, index: int = 0) -> str: + return f"token-{self.identity.removeprefix('guardrail')}-{index}" + + def secret(self, index: int = 0) -> str: + return "synthetic-leaked-secret-" + self.token(index) + + +_END_OF_STREAM_ONLY: Final = MappingProxyType({"streaming_end_of_stream_only": True}) + + +@contextmanager +def _disconnect_rig( + gateway: Gateway, + tmp_path: Path, + *, + params: Mapping[str, JsonValue] = _END_OF_STREAM_ONLY, + guardrail: Callable[[Request], Reply] = _allowing_guardrail, + gated: bool = True, + pause: float = 0, + shape: str = "text", + default_on: bool = True, + workers: int = 1, +) -> Iterator[_DisconnectRig]: + identity: Final = "guardrail" + uuid.uuid4().hex + gate: Final = threading.Event() + with ( + wire_server(guardrail) as policy, + wire_server(_scripted_provider(gate if gated else None, pause, shape)) as upstream, + ): + config: Final = _post_call_config(tmp_path, identity, policy.url, params, default_on) + try: + with ( + owned_proxy_process(gateway, tmp_path, {}, config=config, workers=workers) as owned, + owned.gateway.scenario() as scenario, + ): + yield _DisconnectRig( + owned, + scenario.model(api_base=upstream.url + "/v1", api_key="synthetic-openai-key"), + scenario.model( + model="gemini/gemini-2.5-flash", api_base=upstream.url, api_key="synthetic-gemini-key" + ), + _ScanLog(policy), + identity, + gate, + upstream, + ) + finally: + gate.set() + + +def _close_on(response: httpx.Response, marker: str) -> str: + assert response.status_code == 200, response.read() + for line in response.iter_lines(): + if marker in line: + return line + raise AssertionError(f"The stream ended before the client received {marker}") + + +def _stream_and_close(rig: _DisconnectRig, path: str, body: Mapping[str, JsonValue], marker: str) -> str: + with rig.candidate.client.stream( + "POST", path, json=dict(body), headers={"Authorization": f"Bearer {rig.candidate.key}"} + ) as response: + return _close_on(response, marker) + + +def _chat_body(rig: _DisconnectRig, index: int, stream: bool = True) -> dict[str, JsonValue]: + return { + "model": rig.model, + "messages": [{"role": "user", "content": "synthetic prompt " + rig.token(index)}], + "stream": stream, + } + + +def _chat_httpx(rig: _DisconnectRig, index: int = 0) -> str: + return _stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, index), rig.secret(index)) + + +def _responses_httpx(rig: _DisconnectRig, index: int = 0) -> str: + body: Final = {"model": rig.model, "input": "synthetic prompt " + rig.token(index), "stream": True} + return _stream_and_close(rig, "/v1/responses", body, rig.secret(index)) + + +def _messages_httpx(rig: _DisconnectRig, index: int = 0) -> str: + body: Final = {**_chat_body(rig, index), "max_tokens": 64} + return _stream_and_close(rig, "/v1/messages", body, rig.secret(index)) + + +def _gemini_httpx(rig: _DisconnectRig, index: int = 0) -> str: + body: Final = {"contents": [{"role": "user", "parts": [{"text": "synthetic prompt " + rig.token(index)}]}]} + path: Final = f"/v1beta/models/{rig.gemini}:streamGenerateContent?alt=sse" + return _stream_and_close(rig, path, body, rig.secret(index)) + + +def _chat_async_openai_sdk(rig: _DisconnectRig, index: int = 0) -> str: + async def read() -> str: + client: Final = AsyncOpenAI( + base_url=str(rig.candidate.client.base_url) + "/v1", + api_key=rig.candidate.key, + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=30), + ) + async with client: + stream: Final = await client.chat.completions.create( + model=rig.model, + messages=[{"role": "user", "content": "synthetic prompt " + rig.token(index)}], + stream=True, + ) + async for chunk in stream: + if chunk.choices and rig.secret(index) in (chunk.choices[0].delta.content or ""): + await stream.close() + return chunk.choices[0].delta.content or "" + raise AssertionError("The stream ended before the client received the streamed content") + + return asyncio.run(read()) + + +def _responses_openai_sdk(rig: _DisconnectRig, index: int = 0) -> str: + client: Final = OpenAI( + base_url=str(rig.candidate.client.base_url) + "/v1", + api_key=rig.candidate.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) + with client: + stream: Final = client.responses.create( + model=rig.model, input="synthetic prompt " + rig.token(index), stream=True + ) + for event in stream: + if event.type == "response.output_text.delta" and rig.secret(index) in event.delta: + stream.close() + return event.delta + raise AssertionError("The stream ended before the client received the streamed content") + + +def _messages_anthropic_sdk(rig: _DisconnectRig, index: int = 0) -> str: + client: Final = anthropic.Anthropic( + base_url=str(rig.candidate.client.base_url), + api_key=rig.candidate.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) + with client: + stream: Final = client.messages.create( + model=rig.model, + max_tokens=64, + messages=[{"role": "user", "content": "synthetic prompt " + rig.token(index)}], + stream=True, + ) + for event in stream: + text: Final = ( + event.delta.text if event.type == "content_block_delta" and event.delta.type == "text_delta" else "" + ) + if rig.secret(index) in text: + stream.close() + return text + raise AssertionError("The stream ended before the client received the streamed content") + + +def _scanned_while_upstream_is_held( + rig: _DisconnectRig, disconnect: Callable[[_DisconnectRig, int], str], index: int = 0 +) -> tuple[dict[str, JsonValue], ...]: + try: + received: Final = disconnect(rig, index) + assert rig.secret(index) in received, received + return eventually( + lambda: rig.scans.response_scans(rig.secret(index)), lambda values: len(values) >= 1, seconds=4 + ) + finally: + rig.gate.set() + + +def _post_call_statuses(rig: _DisconnectRig, model: str, rows: int = 1) -> tuple[tuple[str, ...], ...]: + found: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda values: len(values) == rows, + seconds=70, + ) + + def statuses(metadata: JsonValue) -> tuple[str, ...]: + entries: Final = object_value(metadata).get("guardrail_information") or [] + assert isinstance(entries, list), metadata + post_call: Final = tuple( + entry + for entry in (object_value(value) for value in entries) + if entry.get("guardrail_name") == rig.identity and entry.get("guardrail_mode") == "post_call" + ) + return tuple(str(entry["guardrail_status"]) for entry in post_call) + + return tuple(statuses(row["metadata"]) for row in found) + + +_ENDPOINT_CLIENTS: Final = ( + pytest.param(_chat_httpx, id="chat-httpx"), + pytest.param(_chat_async_openai_sdk, id="chat-async-openai-sdk"), + pytest.param(_responses_httpx, id="responses-httpx"), + pytest.param(_responses_openai_sdk, id="responses-openai-sdk"), + pytest.param(_messages_httpx, id="messages-httpx"), + pytest.param(_messages_anthropic_sdk, id="messages-anthropic-sdk"), + pytest.param(_gemini_httpx, id="native-gemini-stream-generate-content"), +) + + +_NO_DISCONNECT_ROW_LIT_8603: Final = pytest.mark.skip( + reason="BUG: LIT-8603 a mid-stream disconnect writes no spend row" +) + + +@pytest.mark.parametrize("disconnect", _ENDPOINT_CLIENTS) +def test_client_disconnect_mid_stream_still_scans_the_content_it_already_received( + gateway: Gateway, tmp_path: Path, disconnect: Callable[[_DisconnectRig, int], str] +) -> None: + with _disconnect_rig(gateway, tmp_path) as rig: + scans: Final = _scanned_while_upstream_is_held(rig, disconnect) + assert len(scans) == 1, scans + + +@pytest.mark.parametrize( + "disconnect", + ( + pytest.param(_chat_httpx, id="chat-httpx"), + pytest.param(_chat_async_openai_sdk, id="chat-async-openai-sdk"), + pytest.param(_responses_httpx, id="responses-httpx", marks=_NO_DISCONNECT_ROW_LIT_8603), + pytest.param(_messages_anthropic_sdk, id="messages-anthropic-sdk", marks=_NO_DISCONNECT_ROW_LIT_8603), + pytest.param( + _gemini_httpx, + id="native-gemini-stream-generate-content", + marks=pytest.mark.skip(reason="BUG: LIT-9087 a mid-stream disconnect writes no spend row"), + ), + ), +) +def test_client_disconnect_mid_stream_records_the_post_call_verdict_on_the_spend_row( + gateway: Gateway, tmp_path: Path, disconnect: Callable[[_DisconnectRig, int], str] +) -> None: + with _disconnect_rig(gateway, tmp_path) as rig: + _scanned_while_upstream_is_held(rig, disconnect) + model: Final = rig.gemini if disconnect is _gemini_httpx else rig.model + assert _post_call_statuses(rig, model) == (("success",),) + + +@pytest.mark.parametrize( + "reply", + ( + pytest.param(Reply(status=500, body=b'{"error": "synthetic guardrail outage"}'), id="guardrail-500"), + pytest.param(Reply(body=b"synthetic non-json guardrail body"), id="guardrail-malformed-200"), + ), +) +def test_client_disconnect_mid_stream_records_a_failed_scan_when_the_guardrail_errors( + gateway: Gateway, tmp_path: Path, reply: Reply +) -> None: + with _disconnect_rig(gateway, tmp_path, guardrail=_failing_response_scans(reply)) as rig: + _scanned_while_upstream_is_held(rig, _chat_httpx) + assert _post_call_statuses(rig, rig.model) == (("guardrail_failed_to_respond",),) + + +def test_client_disconnect_mid_stream_records_a_blocking_verdict_and_keeps_serving( + gateway: Gateway, tmp_path: Path +) -> None: + blocked: Final = Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic leak"}).encode()) + with _disconnect_rig(gateway, tmp_path, guardrail=_failing_response_scans(blocked)) as rig: + _scanned_while_upstream_is_held(rig, _chat_httpx) + statuses: Final = _post_call_statuses(rig, rig.model) + assert len(statuses) == 1 and len(statuses[0]) == 1 and statuses[0][0] != "success", statuses + health: Final = rig.candidate.request("GET", "/health/liveliness") + assert health.status_code == 200, health.text + + +def test_client_disconnect_while_end_of_stream_scan_is_in_flight_still_records_the_verdict( + gateway: Gateway, tmp_path: Path +) -> None: + scan_started: Final = threading.Event() + scan_released: Final = threading.Event() + + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + if json.loads(request.body)["input_type"] == "response": + scan_started.set() + assert scan_released.wait(timeout=30), "The in-flight scan was never released" + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + with _disconnect_rig(gateway, tmp_path, guardrail=guardrail, gated=False) as rig: + with rig.candidate.client.stream( + "POST", + "/v1/chat/completions", + json=_chat_body(rig, 0), + headers={"Authorization": f"Bearer {rig.candidate.key}"}, + ) as response: + try: + assert rig.secret() in _close_on(response, rig.secret()) + assert scan_started.wait(timeout=10), "The end-of-stream scan never started" + finally: + pass + scan_released.set() + assert _post_call_statuses(rig, rig.model) == (("success",),) + assert len(rig.scans.response_scans(rig.secret())) == 1 + + +@pytest.mark.parametrize( + ("params", "shape", "expected"), + ( + pytest.param( + {"streaming_buffer_until_moderated": False}, + "text", + ("synthetic-leaked-secret-",), + id="sampled-before-the-sampling-threshold", + ), + pytest.param( + {"streaming_buffer_until_moderated": False, "streaming_transform_mode": "incremental_diff"}, + "tool_call", + ('\\"query\\": \\"synthetic-leaked-secret-',), + id="incremental-diff-tool-call-in-flight", + ), + pytest.param(dict(_END_OF_STREAM_ONLY), "two_choices", ("-first", "-second"), id="two-choices"), + ), +) +def test_client_disconnect_mid_stream_scans_what_each_streaming_mode_released( + gateway: Gateway, tmp_path: Path, params: dict[str, JsonValue], shape: str, expected: tuple[str, ...] +) -> None: + with _disconnect_rig(gateway, tmp_path, params=params, shape=shape) as rig: + marker: Final = rig.secret() + ("-first" if shape == "two_choices" else "") + try: + received: Final = _stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, 0), rig.token()) + assert rig.token() in received, received + scans: Final = eventually( + lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) >= 1, seconds=4 + ) + finally: + rig.gate.set() + payload: Final = json.dumps(scans[-1]) + assert all(fragment in payload for fragment in expected), (marker, scans) + assert _post_call_statuses(rig, rig.model)[0][-1:] == ("success",) + + +def _tool_call_request(rig: _DisconnectRig, path: str) -> dict[str, JsonValue]: + if path == "/v1/responses": + return {"model": rig.model, "input": "synthetic prompt " + rig.token(), "stream": True} + if path == "/v1/messages": + return {**_chat_body(rig, 0), "max_tokens": 64} + return _chat_body(rig, 0) + + +@pytest.mark.parametrize( + "path", + ( + pytest.param("/v1/chat/completions", id="chat"), + pytest.param("/v1/responses", id="responses"), + pytest.param("/v1/messages", id="messages"), + ), +) +def test_client_disconnect_mid_tool_call_scans_the_tool_call_it_already_received( + gateway: Gateway, tmp_path: Path, path: str +) -> None: + with _disconnect_rig(gateway, tmp_path, shape="tool_call") as rig: + try: + received: Final = _stream_and_close(rig, path, _tool_call_request(rig, path), rig.token()) + assert rig.token() in received, received + scans: Final = eventually( + lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) >= 1, seconds=4 + ) + finally: + rig.gate.set() + assert rig.secret() in json.dumps(scans[-1].get("tool_calls")), scans + + +def test_client_disconnect_after_a_finished_responses_tool_call_scans_the_text_and_tool_call_it_received( + gateway: Gateway, tmp_path: Path +) -> None: + with _disconnect_rig(gateway, tmp_path, shape="text_then_tool_call") as rig: + try: + received: Final = _stream_and_close( + rig, "/v1/responses", _tool_call_request(rig, "/v1/responses"), "response.output_item.done" + ) + assert rig.token() in received, received + scans: Final = eventually( + lambda: tuple(scan for scan in rig.scans.response_scans(rig.secret()) if scan.get("texts")), + lambda values: len(values) >= 1, + seconds=4, + ) + finally: + rig.gate.set() + assert rig.secret() in json.dumps(scans[-1].get("texts")), scans + assert rig.secret() in json.dumps(scans[-1].get("tool_calls")), scans + + +def test_client_disconnect_mid_stream_scans_for_a_guardrail_the_request_opted_into( + gateway: Gateway, tmp_path: Path +) -> None: + with _disconnect_rig(gateway, tmp_path, default_on=False) as rig: + body: Final = {**_chat_body(rig, 0), "guardrails": [rig.identity]} + try: + received: Final = _stream_and_close(rig, "/v1/chat/completions", body, rig.secret()) + assert rig.secret() in received, received + eventually(lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) == 1, seconds=4) + finally: + rig.gate.set() + assert _post_call_statuses(rig, rig.model) == (("success",),) + + +def test_client_disconnect_before_any_content_sends_no_response_scan(gateway: Gateway, tmp_path: Path) -> None: + with _disconnect_rig(gateway, tmp_path, shape="empty") as rig: + try: + _stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, 0), "data: ") + finally: + rig.gate.set() + rows: Final = _post_call_statuses(rig, rig.model) + assert rig.scans.response_scans(rig.secret()) == (), rig.scans.seen + assert len(rows) == 1 and "success" not in rows[0], rows + + +@pytest.mark.parametrize( + "params", + ( + pytest.param(dict(_END_OF_STREAM_ONLY), id="end-of-stream-only"), + pytest.param({"streaming_buffer_until_moderated": True}, id="buffered"), + ), +) +def test_a_fully_read_stream_is_scanned_exactly_once( + gateway: Gateway, tmp_path: Path, params: dict[str, JsonValue] +) -> None: + with _disconnect_rig(gateway, tmp_path, params=params, gated=False) as rig: + response: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0)) + assert response.status_code == 200, response.text + assert rig.secret() in response.text and "[DONE]" in response.text, response.text + assert _post_call_statuses(rig, rig.model) == (("success",),) + assert len(rig.scans.response_scans(rig.secret())) == 1, rig.scans.seen + + +def _cached_twin_rows(rig: _DisconnectRig) -> tuple[tuple[str, ...], ...]: + first: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0, stream=False)) + second: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0, stream=False)) + assert first.status_code == second.status_code == 200, (first.text, second.text) + assert first.json()["choices"] == second.json()["choices"], (first.text, second.text) + assert len(rig.upstream.drain()) == 1, "the second request must be served from the cache" + return _post_call_statuses(rig, rig.model, rows=2) + + +def test_a_non_streaming_response_and_its_cache_hit_are_each_scanned_once(gateway: Gateway, tmp_path: Path) -> None: + with _disconnect_rig(gateway, tmp_path, gated=False) as rig: + rows: Final = _cached_twin_rows(rig) + assert rows[0] == ("success",), rows + assert len(rig.scans.response_scans(rig.secret())) == 2, (rows, rig.scans.seen) + + +def test_a_cache_hit_row_records_the_post_call_verdict_of_its_scan(gateway: Gateway, tmp_path: Path) -> None: + pytest.skip("BUG: LIT-9088 the cache-hit spend row drops the post_call verdict of the scan that ran on it") + with _disconnect_rig(gateway, tmp_path, gated=False) as rig: + assert _cached_twin_rows(rig) == (("success",), ("success",)) + + +def test_concurrent_disconnects_during_a_guardrail_outage_each_record_exactly_one_verdict( + gateway: Gateway, tmp_path: Path +) -> None: + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + body: Final = json.loads(request.body) + index: Final = int(_secret_for(request).rsplit("-", 1)[1]) + if body["input_type"] == "response" and index % 3 == 0: + return Reply(status=503, body=b'{"error": "synthetic guardrail outage"}') + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + clients: Final = (_chat_httpx, _responses_httpx, _messages_httpx) + with _disconnect_rig(gateway, tmp_path, guardrail=guardrail, gated=False, pause=3, workers=2) as rig: + with ThreadPoolExecutor(max_workers=30) as pool: + received: Final = tuple(pool.map(lambda index: clients[index % 3](rig, index), range(30))) + assert all(rig.secret(index) in line for index, line in enumerate(received)), received + scanned: Final = eventually( + lambda: tuple(len(rig.scans.response_scans(rig.secret(index) + '"')) for index in range(30)), + lambda counts: all(count >= 1 for count in counts), + seconds=20, + ) + assert scanned == (1,) * 30, scanned + chat_rows: Final = _post_call_statuses(rig, rig.model, rows=10) + assert sorted(chat_rows) == sorted( + ("guardrail_failed_to_respond",) if index % 3 == 0 else ("success",) for index in range(0, 30, 3) + ), chat_rows + + +def test_disconnect_scans_keep_recording_after_a_worker_is_killed(gateway: Gateway, tmp_path: Path) -> None: + with _disconnect_rig(gateway, tmp_path, gated=False, pause=3, workers=2) as rig: + members: Final = tuple( + member for member in group_members(rig.owned.process.pid) if member.pid != rig.owned.process.pid + ) + workers: Final = tuple(member for member in members if any("spawn_main" in part for part in member.cmdline())) + assert len(workers) >= 2, members + workers[0].send_signal(signal.SIGKILL) + psutil.wait_procs((workers[0],), timeout=10) + with ThreadPoolExecutor(max_workers=8) as pool: + received: Final = tuple(pool.map(lambda index: _chat_httpx(rig, index), range(8))) + assert all(rig.secret(index) in line for index, line in enumerate(received)), received + assert _post_call_statuses(rig, rig.model, rows=8) == (("success",),) * 8 + assert rig.owned.process.poll() is None diff --git a/tests/integration/observability/test_otel_excluded_services.py b/tests/integration/observability/test_otel_excluded_services.py index 8768052b488..09f40630391 100644 --- a/tests/integration/observability/test_otel_excluded_services.py +++ b/tests/integration/observability/test_otel_excluded_services.py @@ -170,7 +170,9 @@ def _assert_tenant_keeps_redis_without_postgres( tenant_start, _ = recorded_spans(audit_sinks.tenant) operator_start, _ = recorded_spans(audit_sinks.operator) traffic: Final = _drive(candidate, langfuse_vars) - _await_db_span(audit_sinks.operator, None, "batch_write_to_db", seconds=60, since=operator_start) + _await_db_span( + audit_sinks.operator, None, "postgres.update LiteLLM_VerificationToken", seconds=60, since=operator_start + ) tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic) _await_db_span(audit_sinks.tenant, tenant_trace, "redis", seconds=60) systems: Final = _db_systems(_trace_spans(audit_sinks.tenant, tenant_trace, seconds=15)) @@ -232,7 +234,9 @@ def test_excluded_services_drops_db_spans_at_tenant_only( _, all_tenant = recorded_spans(audit_sinks.tenant, ten_start) names: Final = sorted(str(span["name"]) for span in all_tenant) assert _db_systems(all_tenant) == set(), f"aux db spans reached tenant: {names}" - assert not any("batch_write_to_db" in name for name in names), f"spend writer reached tenant: {names}" + assert not any("postgres.update LiteLLM_VerificationToken" in name for name in names), ( + f"spend writer reached tenant: {names}" + ) @pytest.mark.timeout(180) @@ -247,8 +251,8 @@ def test_without_excluded_services_the_tenant_still_gets_redis_and_postgres_span with owned_proxy(gateway, tmp_path, {"LITELLM_OTEL_V2": "1"}, config=config, workers=2) as candidate: tenant_start, _ = recorded_spans(audit_sinks.tenant) traffic: Final = _drive(candidate, langfuse_vars) - _await_db_span(audit_sinks.tenant, None, "batch_write_to_db", seconds=60, since=tenant_start) tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic) + _await_db_span(audit_sinks.tenant, tenant_trace, "postgresql", seconds=60, since=tenant_start) _await_db_span(audit_sinks.tenant, tenant_trace, "redis", seconds=60) _assert_core_spans_present(_trace_spans(audit_sinks.tenant, tenant_trace, seconds=15)) _, all_tenant = recorded_spans(audit_sinks.tenant, tenant_start) @@ -574,7 +578,7 @@ def test_bogus_excluded_services_env_logs_and_drops_without_otel_callback( _assert_tenant_keeps_redis_without_postgres(owned.gateway, audit_sinks, langfuse_vars) -def test_postgres_exclusion_covers_batch_write_to_db( +def test_postgres_exclusion_covers_spend_flush( gateway: Gateway, audit_sinks: SpanSinks, otel_audit_config: AuditConfigWriter, @@ -586,11 +590,19 @@ def test_postgres_exclusion_covers_batch_write_to_db( op_start, _ = recorded_spans(audit_sinks.operator) ten_start, _ = recorded_spans(audit_sinks.tenant) traffic: Final = _drive(candidate, langfuse_vars) - _await_db_span(audit_sinks.operator, None, "batch_write_to_db", seconds=60, since=op_start) + _await_db_span(audit_sinks.operator, None, "postgres.update LiteLLM_VerificationToken", seconds=60, since=op_start) + operator_trace: Final = _trace_id(audit_sinks.operator, traffic) + _, operator_spans = recorded_spans(audit_sinks.operator, op_start) + assert any( + span["name"] == "postgres.update LiteLLM_VerificationToken" and span["trace_id"] != operator_trace + for span in operator_spans + ), tuple((span["name"], span["trace_id"]) for span in operator_spans) tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic) _await_db_span(audit_sinks.tenant, tenant_trace, "redis", seconds=60) tenant_spans: Final = _trace_spans(audit_sinks.tenant, tenant_trace, seconds=15) _, all_tenant = recorded_spans(audit_sinks.tenant, ten_start) names: Final = sorted(str(span["name"]) for span in all_tenant) assert "redis" in _db_systems(tenant_spans), f"redis spans missing at tenant: {names}" - assert not any("batch_write_to_db" in name for name in names), f"spend writer reached tenant: {names}" + assert not any("postgres.update LiteLLM_VerificationToken" in name for name in names), ( + f"spend writer reached tenant: {names}" + ) diff --git a/tests/integration/observability/test_otel_v1_request_trace.py b/tests/integration/observability/test_otel_v1_request_trace.py index c99ddc11ced..468aceb4860 100644 --- a/tests/integration/observability/test_otel_v1_request_trace.py +++ b/tests/integration/observability/test_otel_v1_request_trace.py @@ -4,7 +4,7 @@ from pathlib import Path from typing import Final import pytest -from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.client import Gateway, eventually, gateway_from_environment, object_value from integration._support.otlp_sink import Span, SpanSinks, recorded_spans from integration._support.process import owned_proxy from pydantic import JsonValue @@ -25,7 +25,7 @@ def _traces(spans: tuple[Span, ...]) -> dict[str, frozenset[str]]: return {trace: frozenset(span["name"] for span in spans if span["trace_id"] == trace) for trace in trace_ids} -def test_default_otel_logger_puts_datastore_model_and_spend_writer_spans_in_the_request_trace( +def test_default_otel_logger_keeps_spend_flush_outside_the_request_trace( gateway: Gateway, audit_sinks: SpanSinks, otel_audit_config: AuditConfigWriter, tmp_path: Path ) -> None: config: Final = otel_audit_config(tmp_path, {}) @@ -41,7 +41,7 @@ def test_default_otel_logger_puts_datastore_model_and_spend_writer_spans_in_the_ key=key, ) assert response.status_code == 200, response.text - expected: Final = frozenset({"postgres", "redis", "raw_gen_ai_request", "batch_write_to_db"}) + expected: Final = frozenset({"postgres", "redis", "raw_gen_ai_request"}) traces: Final = eventually( lambda: _traces(recorded_spans(audit_sinks.operator, start)[1]), lambda grouped: any(expected <= names for names in grouped.values()), @@ -51,3 +51,22 @@ def test_default_otel_logger_puts_datastore_model_and_spend_writer_spans_in_the_ assert any(expected <= names for names in traces.values()), { trace: sorted(names) for trace, names in traces.items() } + request_trace: Final = next(trace for trace, names in traces.items() if expected <= names) + key_info: Final = eventually( + lambda: candidate.request("GET", "/key/info", key=key, params={"key": key}), + lambda response: response.status_code == 200 + and float(str(object_value(response.json()["info"])["spend"])) > 0, + seconds=60, + ) + assert key_info.status_code == 200, key_info.text + assert float(str(object_value(key_info.json()["info"])["spend"])) > 0 + request_spans: Final = tuple( + span for span in recorded_spans(audit_sinks.operator, start)[1] if span["trace_id"] == request_trace + ) + assert not any(span["name"] == "batch_write_to_db" for span in request_spans), request_spans + assert not any( + span["name"] == "postgres" + and span["attributes"].get("call_type") == "commit_spend_updates" + and span["attributes"].get("table_name") == "LiteLLM_VerificationToken" + for span in request_spans + ), request_spans diff --git a/tests/integration/pricing/test_per_second_pricing.py b/tests/integration/pricing/test_per_second_pricing.py index ad44a631054..a4958bd18c5 100644 --- a/tests/integration/pricing/test_per_second_pricing.py +++ b/tests/integration/pricing/test_per_second_pricing.py @@ -70,8 +70,11 @@ def _clear_observations(upstream: httpx.Client) -> None: def _observed_request_body(upstream: httpx.Client) -> dict[str, JsonValue]: observations: Final = JSON_OBJECT.validate_json(upstream.get("/__observations").content)["requests"] assert isinstance(observations, list) - assert len(observations) == 1 - return object_value(object_value(observations[0])["body"]) + post_observations: Final = tuple( + observation for observation in observations if isinstance(observation, dict) and observation.get("method") == "POST" + ) + assert len(post_observations) == 1 + return object_value(object_value(post_observations[0])["body"]) @pytest.mark.parametrize( diff --git a/tests/integration/routing/test_stale_cost_map_boot.py b/tests/integration/routing/test_stale_cost_map_boot.py index d2eeb2bdc6c..13b3d777587 100644 --- a/tests/integration/routing/test_stale_cost_map_boot.py +++ b/tests/integration/routing/test_stale_cost_map_boot.py @@ -81,6 +81,8 @@ def test_config_deployment_dropped_by_stale_boot_cost_map_is_restored_after_relo { "path": "/v1/chat/completions", "authorization": "Bearer sk-upstream", + "method": "POST", + "api_key": "", "body": {"model": model, "messages": [{"role": "user", "content": "stale cost map control"}]}, } ] diff --git a/tests/integration/spend/test_roi_branch_spend.py b/tests/integration/spend/test_roi_branch_spend.py index 9efb22c43be..cc9cfad8cd2 100644 --- a/tests/integration/spend/test_roi_branch_spend.py +++ b/tests/integration/spend/test_roi_branch_spend.py @@ -123,18 +123,19 @@ def test_documented_header_and_body_tags_reach_recorded_branch_and_pr_cost(gatew ({"tags": tags}, {}), ({}, {"x-litellm-tags": ", ".join(tags + tags)}), ) - for payload, headers in examples: + for index, (payload, headers) in enumerate(examples): response: Final = gateway.request( "POST", "/v1/chat/completions", { "model": model, - "messages": [{"role": "user", "content": "tag attribution"}], + "messages": [{"role": "user", "content": f"tag attribution {index}"}], **payload, }, headers=headers, ) assert response.status_code == 200, response.text + assert len(upstream.drain()) == 3 rows: Final = eventually( lambda: read_rows( 'SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_tags @> %s::jsonb', (json.dumps(tags),) diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index b8af0f0c374..43950a2f011 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -167,13 +167,18 @@ async def test_from_env_reads_with_clickhouse_url( @pytest.mark.asyncio async def test_schema_setup_uses_configured_retention(recording_server: RecordingServer) -> None: recording_server.expected_requests = None + recording_server.default_response = ResponseSpec(body=b"") storage: Final = _native_storage("trace_test", recording_server.base_url, 7) await storage.ensure_schema() ttl_statements: Final = tuple( request.raw_body for request in recording_server.requests if b"MODIFY TTL" in request.raw_body ) - assert len(ttl_statements) == 3 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:]) == ( + 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", + ) @pytest.mark.asyncio @@ -181,7 +186,7 @@ async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement recording_server: RecordingServer, ) -> None: recording_server.expected_requests = 2 - recording_server.enqueue(ResponseSpec(body="")) + recording_server.enqueue(ResponseSpec(body=b"")) recording_server.enqueue(ResponseSpec(status=403, body="denied")) writer_url: Final = recording_server.base_url.replace("http://", "http://writer:p%40ss%2Fword%25@") storage: Final = _native_storage("trace_test", writer_url + "?database=wrong&readonly=1", 7) diff --git a/tests/unit/a2a_protocol/providers/pydantic_ai_agents/test_pydantic_ai_agent_headers.py b/tests/unit/a2a_protocol/providers/pydantic_ai_agents/test_pydantic_ai_agent_headers.py index db561fd1dc2..8d5e15b4b71 100644 --- a/tests/unit/a2a_protocol/providers/pydantic_ai_agents/test_pydantic_ai_agent_headers.py +++ b/tests/unit/a2a_protocol/providers/pydantic_ai_agents/test_pydantic_ai_agent_headers.py @@ -2,10 +2,15 @@ Tests for Pydantic AI agents header forwarding via agent_extra_headers. """ -from unittest.mock import AsyncMock, MagicMock, patch +from typing import Final +from unittest.mock import ANY, AsyncMock, MagicMock, patch +import httpx import pytest +import respx +import litellm +from litellm.a2a_protocol.providers.pydantic_ai_agents.config import PydanticAIProviderConfig from litellm.a2a_protocol.providers.pydantic_ai_agents.transformation import ( PydanticAITransformation, ) @@ -198,3 +203,67 @@ async def test_provider_config_threads_agent_extra_headers(): sent_headers = mock_client.post.await_args.kwargs["headers"] assert sent_headers["x-trace-id"] == "abc-123" assert sent_headers["Content-Type"] == "application/json" + + +COMPLETED_TASK: Final = { + "jsonrpc": "2.0", + "id": "req-5", + "result": { + "id": "task-5", + "kind": "task", + "status": {"state": "completed"}, + "history": [], + "artifacts": [{"artifactId": "a-5", "parts": [{"kind": "text", "text": "ok"}]}], + }, +} + + +def _user_params() -> dict[str, object]: + return {"message": {"role": "user", "parts": [{"kind": "text", "text": "hello"}], "messageId": "msg-user-5"}} + + +@pytest.fixture +def agent_route(monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter) -> respx.Route: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + return respx_mock.post("http://agent.test/").mock(return_value=httpx.Response(200, json=COMPLETED_TASK)) + + +@pytest.mark.asyncio +async def test_provider_config_non_streaming_defaults_to_sixty_second_timeout(agent_route: respx.Route) -> None: + response: Final = await PydanticAIProviderConfig().handle_non_streaming( + "req-5", _user_params(), "http://agent.test" + ) + + request: Final = agent_route.calls.last.request + assert response == { + "jsonrpc": "2.0", + "id": "req-5", + "result": {"kind": "message", "role": "agent", "parts": [{"kind": "text", "text": "ok"}], "messageId": ANY}, + } + assert request.extensions["timeout"] == {"connect": 60.0, "read": 60.0, "write": 60.0, "pool": 60.0} + assert request.headers["content-type"] == "application/json" + + +@pytest.mark.asyncio +async def test_provider_config_non_streaming_forwards_timeout_and_ignores_unrelated_keywords( + agent_route: respx.Route, +) -> None: + response: Final = await PydanticAIProviderConfig().handle_non_streaming( + request_id="req-5", + params=_user_params(), + api_base="http://agent.test", + timeout=5.5, + agent_extra_headers={"x-trace-id": "abc-123"}, + litellm_params={"custom_llm_provider": "pydantic_ai_agents"}, + ) + + request: Final = agent_route.calls.last.request + assert response["id"] == "req-5" + assert request.extensions["timeout"] == {"connect": 5.5, "read": 5.5, "write": 5.5, "pool": 5.5} + assert request.headers["x-trace-id"] == "abc-123" + + +@pytest.mark.asyncio +async def test_provider_config_non_streaming_requires_api_base() -> None: + with pytest.raises(ValueError, match="api_base is required for PydanticAIProviderConfig"): + await PydanticAIProviderConfig().handle_non_streaming("req-5", _user_params()) diff --git a/tests/unit/a2a_protocol/providers/watsonx_orchestrate/test_config.py b/tests/unit/a2a_protocol/providers/watsonx_orchestrate/test_config.py new file mode 100644 index 00000000000..38b896723f7 --- /dev/null +++ b/tests/unit/a2a_protocol/providers/watsonx_orchestrate/test_config.py @@ -0,0 +1,161 @@ +import json +import re +from typing import Final +from unittest.mock import ANY + +import httpx +import pytest +import respx + +import litellm +from litellm.a2a_protocol.providers.watsonx_orchestrate.config import WatsonxOrchestrateA2AConfig +from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import WXOLitellmParams + +MISSING_LITELLM_PARAMS: Final = re.escape( + "litellm_params is required for WatsonxOrchestrateA2AConfig " + "(must contain cp4d_host, instance_id, wxo_agent_id, api_key)" +) +MISSING_HOST: Final = re.escape("'cp4d_host' is required in litellm_params for WXO agents") +RUNS_URL: Final = "https://wxo-config.test/orchestrate/cpd/instances/inst-1/v1/orchestrate/runs" +LITELLM_PARAMS: Final[WXOLitellmParams] = { + "cp4d_host": "https://wxo-config.test", + "instance_id": "inst-1", + "wxo_agent_id": "agent-1", + "api_key": "config-test-key", + "auth_mode": "ibm_cloud", +} +EXPECTED_RUN_BODY: Final = { + "agent_id": "agent-1", + "message": {"role": "user", "content": [{"response_type": "text", "text": "hello"}]}, +} + + +def _a2a_params() -> dict[str, object]: + return {"message": {"role": "user", "parts": [{"kind": "text", "text": "hello"}], "messageId": "m-1"}} + + +def _mock_wxo_route(respx_mock: respx.MockRouter, url: str) -> respx.Route: + respx_mock.post("https://iam.cloud.ibm.com/identity/token").mock( + return_value=httpx.Response(200, json={"access_token": "tok", "expires_in": 0}) + ) + return respx_mock.post(url).mock( + return_value=httpx.Response(200, json={"status": "completed", "results": "wxo says hi"}) + ) + + +@pytest.fixture +def httpx_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + +@pytest.mark.parametrize("litellm_params", [None, {}]) +async def test_handle_non_streaming_rejects_empty_litellm_params(litellm_params: WXOLitellmParams | None) -> None: + with pytest.raises(ValueError, match=MISSING_LITELLM_PARAMS): + await WatsonxOrchestrateA2AConfig().handle_non_streaming("req-1", _a2a_params(), litellm_params=litellm_params) + + +async def test_handle_non_streaming_requires_litellm_params_when_only_other_keywords_are_given() -> None: + with pytest.raises(ValueError, match=MISSING_LITELLM_PARAMS): + await WatsonxOrchestrateA2AConfig().handle_non_streaming( + "req-1", _a2a_params(), "https://ignored.test", agent_extra_headers={"x-tenant-id": "acme"} + ) + + +async def test_handle_non_streaming_hands_litellm_params_to_the_wxo_handler() -> None: + with pytest.raises(ValueError, match=MISSING_HOST): + await WatsonxOrchestrateA2AConfig().handle_non_streaming( + "req-1", _a2a_params(), litellm_params={"instance_id": "inst-1"} + ) + + +@pytest.mark.usefixtures("httpx_transport") +async def test_handle_non_streaming_runs_the_agent_and_ignores_unrelated_keywords( + respx_mock: respx.MockRouter, +) -> None: + runs_route: Final = _mock_wxo_route(respx_mock, RUNS_URL) + + response: Final = await WatsonxOrchestrateA2AConfig().handle_non_streaming( + request_id="req-1", + params=_a2a_params(), + api_base="https://ignored.test", + litellm_params=LITELLM_PARAMS, + agent_extra_headers={"x-tenant-id": "acme"}, + ) + + run_request: Final = runs_route.calls.last.request + assert response == { + "jsonrpc": "2.0", + "id": "req-1", + "result": { + "kind": "message", + "role": "agent", + "parts": [{"kind": "text", "text": "wxo says hi"}], + "messageId": ANY, + }, + } + assert json.loads(run_request.content) == EXPECTED_RUN_BODY + assert run_request.headers["authorization"] == "Bearer tok" + assert "x-tenant-id" not in run_request.headers + + +@pytest.mark.parametrize("litellm_params", [None, {}]) +async def test_handle_streaming_rejects_empty_litellm_params(litellm_params: WXOLitellmParams | None) -> None: + stream: Final = WatsonxOrchestrateA2AConfig().handle_streaming( + "req-1", _a2a_params(), litellm_params=litellm_params + ) + + with pytest.raises(ValueError, match=MISSING_LITELLM_PARAMS): + await anext(stream) + + +async def test_handle_streaming_requires_litellm_params_when_only_other_keywords_are_given() -> None: + stream: Final = WatsonxOrchestrateA2AConfig().handle_streaming( + "req-1", _a2a_params(), "https://ignored.test", agent_extra_headers={"x-tenant-id": "acme"} + ) + + with pytest.raises(ValueError, match=MISSING_LITELLM_PARAMS): + await anext(stream) + + +async def test_handle_streaming_hands_litellm_params_to_the_wxo_handler() -> None: + stream: Final = WatsonxOrchestrateA2AConfig().handle_streaming( + "req-1", _a2a_params(), litellm_params={"instance_id": "inst-1"} + ) + + with pytest.raises(ValueError, match=MISSING_HOST): + await anext(stream) + + +@pytest.mark.usefixtures("httpx_transport") +async def test_handle_streaming_runs_the_agent_and_ignores_unrelated_keywords(respx_mock: respx.MockRouter) -> None: + stream_route: Final = _mock_wxo_route(respx_mock, f"{RUNS_URL}/stream") + + chunks: Final = [ + chunk + async for chunk in WatsonxOrchestrateA2AConfig().handle_streaming( + request_id="req-1", + params=_a2a_params(), + api_base="https://ignored.test", + litellm_params=LITELLM_PARAMS, + agent_extra_headers={"x-tenant-id": "acme"}, + ) + ] + + stream_request: Final = stream_route.calls.last.request + assert [chunk["id"] for chunk in chunks] == ["req-1", "req-1", "req-1", "req-1"] + assert chunks[2]["result"] == { + "contextId": ANY, + "kind": "artifact-update", + "taskId": ANY, + "artifact": {"artifactId": ANY, "parts": [{"kind": "text", "text": "wxo says hi"}]}, + } + assert chunks[3]["result"] == { + "contextId": ANY, + "final": True, + "kind": "status-update", + "status": {"state": "completed"}, + "taskId": ANY, + } + assert json.loads(stream_request.content) == EXPECTED_RUN_BODY + assert stream_request.headers["authorization"] == "Bearer tok" + assert "x-tenant-id" not in stream_request.headers diff --git a/tests/unit/caching/test_unit_test_caching.py b/tests/unit/caching/test_unit_test_caching.py index fd9f4bb9e89..d720838c507 100644 --- a/tests/unit/caching/test_unit_test_caching.py +++ b/tests/unit/caching/test_unit_test_caching.py @@ -251,6 +251,33 @@ def test_get_model_param_value(): assert cache._get_model_param_value(kwargs) == "not-in-caching-group-gpt-3.5-turbo" +def test_get_model_param_value_reads_model_group_from_litellm_metadata(): + cache = Cache() + request = { + "model": "openai/gpt-5.6", + "input": "search this text", + "tools": [{"type": "web_search_preview", "search_context_size": "medium"}], + } + + assert cache._get_model_param_value({**request, "litellm_metadata": {"model_group": "group-a"}}) == "group-a" + assert ( + cache._get_model_param_value({**request, "litellm_params": {"litellm_metadata": {"model_group": "group-a"}}}) + == "group-a" + ) + assert cache._get_model_param_value( + { + **request, + "litellm_metadata": { + "model_group": "group-a", + "caching_groups": [("group-a", "group-b")], + }, + } + ) == "('group-a', 'group-b')" + assert cache.get_cache_key(**request, litellm_metadata={"model_group": "group-a"}) != cache.get_cache_key( + **request, litellm_metadata={"model_group": "group-b"} + ) + + def test_preset_cache_key(): """ Test that the preset cache key is used if it is set in kwargs["litellm_params"] diff --git a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py index 52e44ca5448..6cfdb75fd01 100644 --- a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py +++ b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py @@ -2,6 +2,8 @@ import asyncio import json import os import unittest.mock as mock +from types import SimpleNamespace +from typing import Final from unittest.mock import patch import pytest @@ -19,7 +21,12 @@ from litellm_enterprise.types.enterprise_callbacks.send_emails import ( ) from litellm.integrations.email_templates.email_footer import EMAIL_FOOTER -from litellm.proxy._types import CallInfo, Litellm_EntityType, WebhookEvent +from litellm.proxy._types import ( + CallInfo, + LiteLLM_UserTable, + Litellm_EntityType, + WebhookEvent, +) from litellm.constants import EMAIL_BUDGET_ALERT_TTL @@ -1417,3 +1424,69 @@ async def test_budget_alert_release_failure_does_not_propagate(base_email_logger ) mock_cache.async_delete_cache.assert_awaited_once() + + +class _RecordingEmailLogger(BaseEmailLogger): + def __init__(self): + super().__init__() + self.recipients = [] + + async def send_email(self, from_email, to_email, subject, html_body): + self.recipients.append(to_email) + + +class _UserTable: + def __init__(self, rows): + self._rows = rows + + async def find_unique(self, where): + return self._rows.get(where["user_id"]) + + +def _prisma_client_with_users(*rows): + table: Final = _UserTable({row.user_id: row for row in rows}) + return SimpleNamespace(db=SimpleNamespace(litellm_usertable=table)) + + +def _key_created_event(user_id): + return SendKeyCreatedEmailEvent( + user_id=user_id, + user_email=None, + virtual_key="sk-test", + max_budget=None, + spend=0.0, + event_group=Litellm_EntityType.USER, + event="key_created", + event_message="Key Created", + ) + + +@pytest.mark.asyncio +async def test_key_created_email_goes_to_the_address_stored_for_the_user(monkeypatch): + stored_user: Final = LiteLLM_UserTable( + user_id="user-1", user_email="stored@example.com" + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.prisma_client", + _prisma_client_with_users(stored_user), + ) + email_logger: Final = _RecordingEmailLogger() + + await email_logger.send_key_created_email(_key_created_event("user-1")) + + assert email_logger.recipients == [["stored@example.com"]] + + +@pytest.mark.asyncio +async def test_key_created_email_is_refused_for_a_user_the_database_does_not_know( + monkeypatch, +): + monkeypatch.setattr( + "litellm.proxy.proxy_server.prisma_client", _prisma_client_with_users() + ) + email_logger: Final = _RecordingEmailLogger() + + with pytest.raises(ValueError, match="User email not found for user_id: user-2"): + await email_logger.send_key_created_email(_key_created_event("user-2")) + + assert email_logger.recipients == [] diff --git a/tests/unit/files/test_main.py b/tests/unit/files/test_main.py index 2704bdfb6ff..cb70b39f4d5 100644 --- a/tests/unit/files/test_main.py +++ b/tests/unit/files/test_main.py @@ -3,9 +3,12 @@ from urllib.parse import parse_qs, urlparse import httpx import pytest +import respx +from pydantic import ValidationError import litellm from litellm.llms.custom_httpx.http_handler import HTTPHandler +from litellm.types.llms.openai import OpenAIFileObject NATIVE_VERTEX_ROWS: Final = ( b'{"request": {"contents": [{"role": "user", "parts": [{"text": "Who won the 2024 Tour de France?"}]}],' @@ -69,3 +72,52 @@ def test_create_file_passthrough_kwarg_ships_native_rows_byte_for_byte_under_the assert upload.read() == NATIVE_VERTEX_ROWS assert object_name.startswith("litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/") assert file_object.id == f"gs://my-bucket/{object_name}" + + +OPENAI_FILES_API_BASE: Final = "https://files.test/v1" +PROVIDER_FILE: Final = { + "id": "file-abc123", + "bytes": 120, + "created_at": 1700000000, + "filename": "batch.jsonl", + "object": "file", + "purpose": "batch", +} +UNSET_OPTIONAL_FIELDS: Final = {"status": None, "expires_at": None, "status_details": None} + + +@pytest.mark.parametrize( + "provider_extras, expected", + [ + ({}, {**PROVIDER_FILE, **UNSET_OPTIONAL_FIELDS}), + ( + {"status": "processed", "expires_at": 1800000000, "status_details": "ok", "provider_only": {"a": [1]}}, + {**PROVIDER_FILE, "status": "processed", "expires_at": 1800000000, "status_details": "ok"}, + ), + ], + ids=["required-fields-only", "optional-and-unknown-fields"], +) +@respx.mock +async def test_afile_retrieve_returns_the_provider_file_as_an_openai_file_object(provider_extras, expected): + respx.get(f"{OPENAI_FILES_API_BASE}/files/file-abc123").respond(200, json={**PROVIDER_FILE, **provider_extras}) + + file_object: Final = await litellm.afile_retrieve( + file_id="file-abc123", custom_llm_provider="openai", api_key="sk-test", api_base=OPENAI_FILES_API_BASE + ) + + assert type(file_object) is OpenAIFileObject + assert file_object.model_dump() == expected + + +@respx.mock +async def test_afile_retrieve_rejects_a_provider_file_without_its_size(): + provider_file: Final = {key: value for key, value in PROVIDER_FILE.items() if key != "bytes"} + respx.get(f"{OPENAI_FILES_API_BASE}/files/file-abc123").respond(200, json=provider_file) + + with pytest.raises(ValidationError) as exc_info: + await litellm.afile_retrieve( + file_id="file-abc123", custom_llm_provider="openai", api_key="sk-test", api_base=OPENAI_FILES_API_BASE + ) + + assert exc_info.value.title == "OpenAIFileObject" + assert [error["loc"] for error in exc_info.value.errors()] == [("bytes",)] diff --git a/tests/unit/integrations/arize/test_arize_utils.py b/tests/unit/integrations/arize/test_arize_utils.py index 167b083e147..0d833b90148 100644 --- a/tests/unit/integrations/arize/test_arize_utils.py +++ b/tests/unit/integrations/arize/test_arize_utils.py @@ -5,6 +5,7 @@ from typing import Optional import asyncio +import httpx import pytest import litellm @@ -13,6 +14,7 @@ from litellm.integrations._types.open_inference import ( SpanAttributes, ToolCallAttributes, ) +from litellm.integrations.arize._utils import _coerce_response_obj_for_attrs, _parse_passthrough_response from litellm.integrations.arize.arize import ArizeLogger from litellm.integrations.custom_logger import CustomLogger from litellm.types.utils import Choices, StandardCallbackDynamicParams @@ -1513,3 +1515,42 @@ def test_arize_mcp_emitter_is_inert_without_a_standard_logging_object(): written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} assert SpanAttributes.TOOL_NAME not in written + + +@pytest.mark.parametrize( + ("text", "expected"), + [ + ('{"id": "msg_1", "usage": {"input_tokens": 3}}', {"id": "msg_1", "usage": {"input_tokens": 3}}), + ("{}", {}), + ], +) +def test_coerce_response_obj_decodes_a_json_object_body(text: str, expected: dict[str, object]): + assert _coerce_response_obj_for_attrs(httpx.Response(200, text=text)) == expected + + +@pytest.mark.parametrize("text", ["[1, 2]", '"text"', "7", "null", "true", "not json"]) +def test_coerce_response_obj_keeps_the_response_when_the_body_is_not_a_json_object(text: str): + response = httpx.Response(200, text=text) + + assert _coerce_response_obj_for_attrs(response) is response + + +@pytest.mark.parametrize( + ("raw", "coerced", "kwargs", "expected"), + [ + (None, {"response": '{"id": "wrapped"}'}, {}, {"id": "wrapped"}), + (None, {"response": "[1, 2]"}, {}, None), + ({"id": "raw"}, {"response": "[1, 2]"}, {}, {"id": "raw"}), + ({"id": "raw"}, {"response": "not json"}, {}, {"id": "raw"}), + (None, {"response": "7"}, {"original_response": '{"id": "original"}'}, {"id": "original"}), + (None, None, {"original_response": '{"id": "original"}'}, {"id": "original"}), + (None, None, {"original_response": "[1, 2]"}, None), + (None, None, {"original_response": '"text"'}, None), + (None, None, {"original_response": "null"}, None), + (None, None, {"original_response": "not json"}, None), + ], +) +def test_parse_passthrough_response_reads_only_json_objects_from_text( + raw: object, coerced: object, kwargs: dict[str, object], expected: dict[str, object] | None +): + assert _parse_passthrough_response(raw, coerced, kwargs) == expected diff --git a/tests/unit/integrations/datadog/test_datadog_llm_obs.py b/tests/unit/integrations/datadog/test_datadog_llm_obs.py index 77d2518696c..a7826e11483 100644 --- a/tests/unit/integrations/datadog/test_datadog_llm_obs.py +++ b/tests/unit/integrations/datadog/test_datadog_llm_obs.py @@ -1139,3 +1139,31 @@ def test_reasoning_content_survives_the_mapping(logger: DataDogLLMObsLogger) -> ) assert payload["meta"]["output"]["messages"][0]["reasoning_content"] == "thinking" + + +@pytest.mark.parametrize( + ("raw_arguments", "shipped"), + [ + ('{"city": "Paris", "days": [1, 2]}', {"city": "Paris", "days": [1, 2]}), + ("{}", {}), + ("[1, 2]", "[1, 2]"), + ("null", "null"), + ("true", "true"), + ('"text"', '"text"'), + ("1.5", "1.5"), + ("", ""), + ], +) +def test_tool_arguments_ship_as_an_object_only_when_they_decode_to_one( + logger: DataDogLLMObsLogger, raw_arguments: str, shipped: object +) -> None: + payload = build( + logger, + response_message={ + "role": "assistant", + "content": None, + "tool_calls": [{"id": "c1", "type": "function", "function": {"name": "f", "arguments": raw_arguments}}], + }, + ) + + assert payload["meta"]["output"]["messages"][0]["tool_calls"][0]["arguments"] == shipped diff --git a/tests/unit/integrations/opik/opik_payload_builder/__init__.py b/tests/unit/integrations/opik/opik_payload_builder/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/integrations/opik/opik_payload_builder/test_api.py b/tests/unit/integrations/opik/opik_payload_builder/test_api.py new file mode 100644 index 00000000000..d2e20909d3c --- /dev/null +++ b/tests/unit/integrations/opik/opik_payload_builder/test_api.py @@ -0,0 +1,230 @@ +from collections import OrderedDict +from collections.abc import Mapping +from datetime import datetime, timezone +from types import MappingProxyType +from typing import Final +from uuid import UUID + +import pytest +from pydantic import ValidationError + +from litellm.integrations.opik.opik_payload_builder import build_opik_payload +from litellm.integrations.opik.opik_payload_builder.types import SpanPayload, TracePayload +from litellm.types.utils import ModelResponse, Usage + +_START: Final = datetime(2026, 1, 2, 3, 4, 5, tzinfo=timezone.utc) +_END: Final = datetime(2026, 1, 2, 3, 4, 6, tzinfo=timezone.utc) +_MESSAGES: Final = [{"role": "user", "content": "hi"}] +_RESPONSE: Final = {"id": "chatcmpl-1", "choices": []} +_HIDDEN_PARAMS: Final = {"model_id": "deployment-1"} +_LOGGING_METADATA: Final = { + "user_api_key_alias": "team-key", + "requester_metadata": {"opik": {"thread_id": "thread-1"}}, +} +_STANDARD_LOGGING_OBJECT: Final = { + "call_type": "acompletion", + "status": "success", + "model": "gpt-4o", + "metadata": _LOGGING_METADATA, + "messages": _MESSAGES, + "response": _RESPONSE, + "hidden_params": _HIDDEN_PARAMS, + "trace_id": "not-forwarded-to-opik", +} +_RESPONSE_OBJ: Final = ModelResponse( + model="gpt-4o", + created=1767323045, + usage=Usage(prompt_tokens=1, completion_tokens=2, total_tokens=3), +) +_OPIK_FIELDS_OF_THE_LOGGING_OBJECT: Final = { + "type": "acompletion", + "status": "success", + "model": "gpt-4o", + "hidden_params": _HIDDEN_PARAMS, +} +_EXISTING_TRACE: Final = {"metadata": {"opik": {"current_span_data": {"trace_id": "trace-1", "id": "span-0"}}}} + + +def _payloads_attached_to_trace_1(standard_logging_object: object) -> tuple[TracePayload | None, SpanPayload]: + return build_opik_payload( + kwargs={"standard_logging_object": standard_logging_object, "litellm_params": _EXISTING_TRACE}, + response_obj=_RESPONSE_OBJ, + start_time=_START, + end_time=_END, + project_name="default-project", + ) + + +def _span_attached_to_trace_1(span_id: str, metadata: Mapping[str, object]) -> SpanPayload: + return SpanPayload( + id=span_id, + project_name="default-project", + trace_id="trace-1", + parent_span_id="span-0", + name="gpt-4o_chat.completion_1767323045", + type="llm", + model="gpt-4o", + start_time="2026-01-02T03:04:05Z", + end_time="2026-01-02T03:04:06Z", + input=_MESSAGES, + output=_RESPONSE, + metadata=metadata, + tags=[], + usage={"completion_tokens": 2, "prompt_tokens": 1, "total_tokens": 3}, + ) + + +def test_build_opik_payload_creates_a_trace_and_its_span_from_the_standard_logging_object() -> None: + trace, span = build_opik_payload( + kwargs={ + "standard_logging_object": _STANDARD_LOGGING_OBJECT, + "custom_llm_provider": "openai", + "response_cost": 0.25, + }, + response_obj=_RESPONSE_OBJ, + start_time=_START, + end_time=_END, + project_name="default-project", + ) + metadata: Final = { + "thread_id": "thread-1", + "created_from": "litellm", + **_LOGGING_METADATA, + **_OPIK_FIELDS_OF_THE_LOGGING_OBJECT, + "cost": {"total_tokens": 0.25, "currency": "USD"}, + } + + assert trace is not None + assert (trace, span) == ( + TracePayload( + project_name="default-project", + id=trace.id, + name="chat.completion", + start_time="2026-01-02T03:04:05Z", + end_time="2026-01-02T03:04:06Z", + input=_MESSAGES, + output=_RESPONSE, + metadata=metadata, + tags=["openai"], + thread_id="thread-1", + ), + SpanPayload( + id=span.id, + project_name="default-project", + trace_id=trace.id, + name="gpt-4o_chat.completion_1767323045", + type="llm", + model="gpt-4o", + start_time="2026-01-02T03:04:05Z", + end_time="2026-01-02T03:04:06Z", + input=_MESSAGES, + output=_RESPONSE, + metadata=metadata, + tags=["openai"], + usage={"completion_tokens": 2, "prompt_tokens": 1, "total_tokens": 3}, + provider="openai", + total_cost=0.25, + ), + ) + assert [UUID(trace.id).version, UUID(span.id).version, trace.id != span.id] == [7, 7, True] + assert [span.input is _MESSAGES, span.output is _RESPONSE, span.metadata["hidden_params"] is _HIDDEN_PARAMS] == [ + True, + True, + True, + ] + assert list(span.metadata) == [ + "thread_id", + "created_from", + "user_api_key_alias", + "requester_metadata", + "type", + "status", + "model", + "hidden_params", + "cost", + ] + + +@pytest.mark.parametrize( + "standard_logging_object", + [ + _STANDARD_LOGGING_OBJECT, + OrderedDict(_STANDARD_LOGGING_OBJECT), + MappingProxyType(_STANDARD_LOGGING_OBJECT), + {**_STANDARD_LOGGING_OBJECT, "metadata": MappingProxyType(_LOGGING_METADATA)}, + ], +) +def test_build_opik_payload_reads_any_string_keyed_mapping_as_the_standard_logging_object( + standard_logging_object: object, +) -> None: + trace, span = _payloads_attached_to_trace_1(standard_logging_object) + + assert (trace, span) == ( + None, + _span_attached_to_trace_1( + span.id, + { + "thread_id": "thread-1", + "created_from": "litellm", + **_LOGGING_METADATA, + **_OPIK_FIELDS_OF_THE_LOGGING_OBJECT, + }, + ), + ) + + +@pytest.mark.parametrize( + "standard_logging_object", + [ + {key: value for key, value in _STANDARD_LOGGING_OBJECT.items() if key != "metadata"}, + {**_STANDARD_LOGGING_OBJECT, "metadata": None}, + {**_STANDARD_LOGGING_OBJECT, "metadata": {}}, + {**_STANDARD_LOGGING_OBJECT, "metadata": ""}, + ], +) +def test_build_opik_payload_without_standard_logging_metadata_keeps_only_the_logging_object_fields( + standard_logging_object: object, +) -> None: + trace, span = _payloads_attached_to_trace_1(standard_logging_object) + + assert (trace, span) == ( + None, + _span_attached_to_trace_1(span.id, {"created_from": "litellm", **_OPIK_FIELDS_OF_THE_LOGGING_OBJECT}), + ) + + +def test_build_opik_payload_of_an_empty_standard_logging_object_has_empty_input_and_output() -> None: + _, span = _payloads_attached_to_trace_1({}) + + assert (span.input, span.output, span.metadata) == ({}, {}, {"created_from": "litellm"}) + + +@pytest.mark.parametrize( + "standard_logging_object", + [ + None, + "standard_logging_object", + ["messages", "response"], + list(_STANDARD_LOGGING_OBJECT.items()), + 7, + {**_STANDARD_LOGGING_OBJECT, 7: "keys must be strings"}, + {**_STANDARD_LOGGING_OBJECT, "metadata": "metadata"}, + {**_STANDARD_LOGGING_OBJECT, "metadata": ["user_api_key_alias"]}, + {**_STANDARD_LOGGING_OBJECT, "metadata": 7}, + {**_STANDARD_LOGGING_OBJECT, "metadata": {7: "keys must be strings"}}, + ], +) +def test_build_opik_payload_rejects_a_standard_logging_object_that_is_not_a_string_keyed_mapping( + standard_logging_object: object, +) -> None: + with pytest.raises(ValidationError) as raised: + _payloads_attached_to_trace_1(standard_logging_object) + + assert "input_value" not in str(raised.value) + + +def test_build_opik_payload_without_a_standard_logging_object_raises_a_key_error() -> None: + with pytest.raises(KeyError, match="standard_logging_object"): + build_opik_payload( + kwargs={}, response_obj=_RESPONSE_OBJ, start_time=_START, end_time=_END, project_name="default-project" + ) diff --git a/tests/unit/integrations/opik/test_opik_extractors.py b/tests/unit/integrations/opik/test_opik_extractors.py index 6f85a1c6090..4ef9f500b54 100644 --- a/tests/unit/integrations/opik/test_opik_extractors.py +++ b/tests/unit/integrations/opik/test_opik_extractors.py @@ -1,4 +1,7 @@ +import pytest + from litellm.integrations.opik.opik_payload_builder.extractors import ( + apply_proxy_header_overrides, extract_opik_metadata, ) @@ -82,3 +85,32 @@ def test_extract_opik_metadata_requester_metadata_overrides_all_other_sources(): "workspace": "requester-workspace", "thread_id": "requester-thread", } + + +@pytest.mark.parametrize( + ("opik_tags_header", "expected_tags"), + [ + ('["from-header", "second"]', ["from-request", "from-header", "second"]), + ('["text", 7, null, {"nested": [1]}]', ["from-request", "text", 7, None, {"nested": [1]}]), + ("[]", ["from-request"]), + ('{"not": "a list"}', ["from-request"]), + ('"not-a-list"', ["from-request"]), + ("null", ["from-request"]), + ("not json", ["from-request"]), + ], +) +def test_opik_tags_header_adds_tags_only_when_it_is_a_json_list(opik_tags_header: str, expected_tags: list[object]): + overrides = apply_proxy_header_overrides("project", ["from-request"], None, {"opik_tags": opik_tags_header}) + + assert overrides == ("project", expected_tags, None) + + +def test_opik_headers_override_the_project_name_and_thread_id(): + overrides = apply_proxy_header_overrides( + "project", + ["from-request"], + "thread-from-request", + {"opik_project_name": "header-project", "opik_thread_id": "header-thread", "opik_tags": "", "x-other": "1"}, + ) + + assert overrides == ("header-project", ["from-request"], "header-thread") diff --git a/tests/unit/integrations/test_galileo.py b/tests/unit/integrations/test_galileo.py index d0709b966d4..d067b31a651 100644 --- a/tests/unit/integrations/test_galileo.py +++ b/tests/unit/integrations/test_galileo.py @@ -1,7 +1,11 @@ +from dataclasses import dataclass from datetime import datetime, timezone +from decimal import Decimal +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest +from pydantic import BaseModel, ValidationError from litellm.integrations.galileo import GalileoObserve @@ -905,3 +909,78 @@ async def test_galileo_async_log_success_appends_and_flushes(galileo_v2_env): assert "/ingest/traces/" in flushed_url["url"] assert logger.in_memory_records == [] + + +class RealtimeEvent(BaseModel): + type: str + + +@dataclass(frozen=True) +class ThirdPartyEvent: + model_dump: object + + +@pytest.mark.parametrize( + ("event", "expected_output"), + [ + (RealtimeEvent(type="response.done"), '[{"type": "response.done"}]'), + (ThirdPartyEvent(model_dump=lambda: {"type": "custom"}), '[{"type": "custom"}]'), + (ThirdPartyEvent(model_dump=lambda: [Decimal("2.5")]), '[["2.5"]]'), + (Decimal("1.5"), '["1.5"]'), + ], +) +def test_galileo_realtime_output_serializes_events_that_json_cannot_encode( + galileo_v2_env: None, event: object, expected_output: str +) -> None: + assert GalileoObserve().get_output_str_from_response([event], {"call_type": "_arealtime"}) == expected_output + + +@pytest.mark.parametrize("model_dump", [None, "not callable"]) +def test_galileo_realtime_output_with_a_model_dump_that_cannot_be_called_raises_a_validation_error( + galileo_v2_env: None, model_dump: object +) -> None: + with pytest.raises(ValidationError) as raised: + GalileoObserve().get_output_str_from_response( + [ThirdPartyEvent(model_dump=model_dump)], {"call_type": "_arealtime"} + ) + + assert "input_value" not in str(raised.value) + + +@dataclass(frozen=True) +class ThirdPartyMessage: + json: object + + +def _chat_response_carrying(message: object) -> ModelResponse: + response: Final = ModelResponse(choices=[Choices(message=Message(content="replaced", role="assistant"))]) + response.choices[0].message = message + return response + + +@pytest.mark.parametrize( + ("message", "expected_output"), + [ + ( + ThirdPartyMessage(json=lambda: '{"role":"assistant","content":"hi"}'), + '{"role": "assistant", "content": "hi"}', + ), + ( + ThirdPartyMessage(json=lambda: {"role": "assistant", "content": "hi"}), + '{"role": "assistant", "content": "hi"}', + ), + (ThirdPartyMessage(json=lambda: '"just text"'), "just text"), + (ThirdPartyMessage(json=lambda: None), ""), + ({"role": "assistant", "content": "hi"}, '{"role": "assistant", "content": "hi"}'), + ("plain reply", "plain reply"), + (None, ""), + ], +) +def test_galileo_output_str_of_a_chat_response_whose_message_is_not_a_litellm_message( + galileo_v2_env: None, message: object, expected_output: str +) -> None: + output: Final = GalileoObserve().get_output_str_from_response( + _chat_response_carrying(message), {"call_type": "acompletion"} + ) + + assert output == expected_output diff --git a/tests/unit/integrations/test_opentelemetry.py b/tests/unit/integrations/test_opentelemetry.py index 175bd95c263..2ce3561d441 100644 --- a/tests/unit/integrations/test_opentelemetry.py +++ b/tests/unit/integrations/test_opentelemetry.py @@ -6762,3 +6762,37 @@ class TestOpenTelemetryNonInferenceUsage(unittest.TestCase): self.assertEqual( self._time_per_output_token_calls("aget_responses", response_obj=self.BACKGROUND_RESPONSE_OBJ), 1 ) + + +def _raw_response_span_attributes(original_response: str) -> dict[str, object]: + span_exporter = InMemorySpanExporter() + tracer_provider = TracerProvider() + tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter)) + span = tracer_provider.get_tracer(__name__).start_span("raw_gen_ai_request") + + OpenTelemetry(tracer_provider=tracer_provider).set_raw_request_attributes( + span, + {"litellm_params": {"custom_llm_provider": "vertex_ai"}, "original_response": original_response}, + None, + ) + span.end() + + return dict(span_exporter.get_finished_spans()[0].attributes or {}) + + +@pytest.mark.parametrize( + ("original_response", "expected"), + [ + ('{"id": "r1", "model": "m"}', {"llm.vertex_ai.id": "r1", "llm.vertex_ai.model": "m"}), + ("{}", {}), + ("not json", {"llm.vertex_ai.stringified_raw_response": "not json"}), + ("[1, 2]", {}), + ('"text"', {}), + ("7", {}), + ("null", {}), + ], +) +def test_set_raw_request_attributes_stamps_only_json_object_responses( + original_response: str, expected: dict[str, object] +): + assert _raw_response_span_attributes(original_response) == expected diff --git a/tests/unit/integrations/test_opik_utils.py b/tests/unit/integrations/test_opik_utils.py index a4250acf1dc..376c70a4ed9 100644 --- a/tests/unit/integrations/test_opik_utils.py +++ b/tests/unit/integrations/test_opik_utils.py @@ -4,7 +4,7 @@ import uuid from datetime import datetime, timezone from unittest.mock import patch -from litellm.integrations.opik.utils import create_uuid7 +from litellm.integrations.opik.utils import create_uuid7, get_traces_and_spans_from_payload def _timestamp_ms(uuid_str: str) -> int: @@ -27,3 +27,22 @@ def test_create_uuid7_encodes_timestamp_in_milliseconds(): value = create_uuid7() assert _timestamp_ms(value) == int(fixed.timestamp() * 1000) + + +def test_queued_opik_payloads_are_split_into_traces_and_spans_without_their_null_fields(): + traces, spans = get_traces_and_spans_from_payload( + [ + {"id": "trace-1", "name": "chat.completion", "thread_id": None, "input": {"kept": None}}, + {"id": "span-1", "type": "llm", "parent_span_id": None, "tags": [], "total_cost": 0}, + {"id": "span-2", "type": None}, + ] + ) + + assert (traces, spans) == ( + [{"id": "trace-1", "name": "chat.completion", "input": {"kept": None}}], + [{"id": "span-1", "type": "llm", "tags": [], "total_cost": 0}, {"id": "span-2"}], + ) + + +def test_an_empty_opik_queue_has_no_traces_and_no_spans(): + assert get_traces_and_spans_from_payload([]) == ([], []) diff --git a/tests/unit/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py b/tests/unit/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py index a0766ac3d58..726ce2f7f1d 100644 --- a/tests/unit/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py +++ b/tests/unit/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py @@ -1,5 +1,5 @@ import logging -from collections.abc import Iterator, Mapping +from collections.abc import Callable, Iterable, Iterator, Mapping from dataclasses import dataclass, field from types import MappingProxyType from typing import Literal, Protocol @@ -33,6 +33,7 @@ from litellm.types.utils import ( ) from litellm.types.vector_stores import ( VectorStoreResultContent, + VectorStoreSearchFailure, VectorStoreSearchResponse, VectorStoreSearchResult, ) @@ -457,6 +458,166 @@ async def test_a_failing_vector_store_is_reported_on_the_streaming_chunk(registr ) +@dataclass +class ThirdPartyMessage: + provider_specific_fields: dict[str, object] | None = None + + +@dataclass +class ThirdPartyChoice: + message: ThirdPartyMessage | None = None + delta: ThirdPartyMessage | None = None + + +@dataclass +class ThirdPartyResponse: + choices: object + + +_SEARCH_FAILURES = ( + VectorStoreSearchFailure(vector_store_id="vs-broken", custom_llm_provider="bedrock", error="search timed out"), +) + + +def _logging_obj_with_search_failures() -> FakeLoggingObj: + logging_obj = FakeLoggingObj({}) + logging_obj.model_call_details["vector_store_search_failures"] = _SEARCH_FAILURES + return logging_obj + + +@pytest.mark.asyncio +async def test_search_failures_join_the_provider_fields_the_message_already_carries() -> None: + existing_fields: dict[str, object] = {"citations": ["doc-1"]} + response = ModelResponse( + choices=[Choices(message=Message(content="an answer", provider_specific_fields=existing_fields))] + ) + + returned = await VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=None) + ).async_post_call_success_deployment_hook( + request_data={"litellm_logging_obj": _logging_obj_with_search_failures()}, + response=response, + call_type=CallTypes.acompletion, + ) + + assert returned is response + assert _first_message(response).provider_specific_fields is existing_fields + assert existing_fields == {"citations": ["doc-1"], "vector_store_search_failures": _SEARCH_FAILURES} + + +@pytest.mark.asyncio +async def test_search_failures_join_the_provider_fields_the_streaming_delta_already_carries() -> None: + existing_fields: dict[str, object] = {"citations": ["doc-1"]} + chunk = ModelResponseStream( + choices=[StreamingChoices(delta=Delta(content="an answer", provider_specific_fields=existing_fields))] + ) + + returned = await VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=None) + ).async_post_call_streaming_deployment_hook( + request_data=_logging_obj_with_search_failures().model_call_details, + response_chunk=chunk, + call_type=CallTypes.acompletion, + ) + + assert returned is chunk + assert chunk.choices[0].delta.provider_specific_fields is existing_fields + assert existing_fields == {"citations": ["doc-1"], "vector_store_search_failures": _SEARCH_FAILURES} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("as_choices", [list, tuple, iter]) +async def test_a_chunk_that_only_looks_like_a_chat_completion_chunk_is_annotated_too( + as_choices: Callable[[list[ThirdPartyChoice]], Iterable[ThirdPartyChoice]], +) -> None: + delta = ThirdPartyMessage() + chunk = ThirdPartyResponse(choices=as_choices([ThirdPartyChoice(delta=None), ThirdPartyChoice(delta=delta)])) + + returned = await VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=None) + ).async_post_call_streaming_deployment_hook( + request_data=_logging_obj_with_search_failures().model_call_details, + response_chunk=chunk, + call_type=CallTypes.acompletion, + ) + + assert returned is chunk + assert delta.provider_specific_fields == {"vector_store_search_failures": _SEARCH_FAILURES} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("choices", [None, [], ()]) +async def test_a_response_without_choices_is_returned_untouched( + choices: object, warnings: list[logging.LogRecord] +) -> None: + response = ThirdPartyResponse(choices=choices) + + returned = await VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=None) + ).async_post_call_success_deployment_hook( + request_data={"litellm_logging_obj": _logging_obj_with_search_failures()}, + response=response, + call_type=CallTypes.acompletion, + ) + + assert returned is response + assert warnings == [] + + +@pytest.mark.asyncio +async def test_a_chunk_whose_choices_cannot_be_iterated_is_logged_and_passed_through( + warnings: list[logging.LogRecord], +) -> None: + chunk = ThirdPartyResponse(choices=7) + + returned = await VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=None) + ).async_post_call_streaming_deployment_hook( + request_data=_logging_obj_with_search_failures().model_call_details, + response_chunk=chunk, + call_type=CallTypes.acompletion, + ) + + assert returned is chunk + assert [record.levelname for record in warnings] == ["ERROR"] + assert warnings[0].getMessage().startswith("Error adding search results to streaming chunk: ") + assert "input_value" not in warnings[0].getMessage() + + +@pytest.mark.asyncio +async def test_the_search_receives_the_requests_own_metadata_object(registry_with: RegisterStores) -> None: + registry_with("vs-router") + router = RecordingRouter() + metadata = {"user_api_key_team_id": "team-a"} + + await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)), + ["vs-router"], + FakeLoggingObj(metadata), + ) + + assert [call["metadata"] is metadata for call in router.calls] == [True] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("litellm_params", [None, "metadata", ["metadata"], {"other": "value"}]) +async def test_a_request_without_metadata_in_its_litellm_params_searches_with_empty_metadata( + registry_with: RegisterStores, litellm_params: object +) -> None: + registry_with("vs-router") + router = RecordingRouter() + logging_obj = FakeLoggingObj({}) + logging_obj.model_call_details["litellm_params"] = litellm_params + + await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)), + ["vs-router"], + logging_obj, + ) + + assert [call["metadata"] for call in router.calls] == [{}] + + @pytest.mark.asyncio async def test_error_mode_fails_the_request_instead_of_answering_without_the_knowledge_base( registry_with: RegisterStores, diff --git a/tests/unit/llms/a2a/chat/test_a2a_chat_transformation.py b/tests/unit/llms/a2a/chat/test_a2a_chat_transformation.py index 6440825e135..a419e73a2e1 100644 --- a/tests/unit/llms/a2a/chat/test_a2a_chat_transformation.py +++ b/tests/unit/llms/a2a/chat/test_a2a_chat_transformation.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock import pytest +from litellm.llms.a2a.chat.streaming_iterator import A2AModelResponseIterator from litellm.llms.a2a.chat.transformation import A2AConfig from litellm.types.utils import ModelResponse @@ -85,3 +86,38 @@ def test_transform_request_tags_the_message_with_its_kind(optional_params: dict) ) assert request["params"]["message"]["kind"] == "message" + + +def test_get_model_response_iterator_parses_the_agent_stream(): + iterator = A2AConfig().get_model_response_iterator( + streaming_response=iter( + [ + '{"jsonrpc":"2.0","id":"1","result":{"kind":"task","status":{"state":"completed"},' + '"artifacts":[{"parts":[{"kind":"text","text":"7"}]}]}}' + ] + ), + sync_stream=True, + ) + + chunk = next(iterator) + + assert isinstance(iterator, A2AModelResponseIterator) + assert chunk["text"] == "7" + assert chunk["finish_reason"] == "stop" + + +def test_resolve_agent_config_from_registry_returns_the_explicit_headers_object_untouched(): + headers: dict[str, object] = {"X-Test": "value"} + optional_params: dict[str, object] = {"stream": True} + + resolved = A2AConfig.resolve_agent_config_from_registry( + agent_name="test-agent", + api_base="http://explicit.example", + api_key="explicit-key", + headers=headers, + optional_params=optional_params, + ) + + assert resolved[2] is headers + assert resolved[:2] == ("http://explicit.example", "explicit-key") + assert optional_params == {"stream": True} diff --git a/tests/unit/llms/aiohttp_openai/__init__.py b/tests/unit/llms/aiohttp_openai/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/aiohttp_openai/chat/__init__.py b/tests/unit/llms/aiohttp_openai/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/aiohttp_openai/chat/test_transformation.py b/tests/unit/llms/aiohttp_openai/chat/test_transformation.py new file mode 100644 index 00000000000..253f9dbe746 --- /dev/null +++ b/tests/unit/llms/aiohttp_openai/chat/test_transformation.py @@ -0,0 +1,101 @@ +from unittest.mock import AsyncMock, Mock + +import pytest +from aiohttp import ClientResponse +from pydantic import ValidationError + +from litellm.llms.aiohttp_openai.chat.transformation import AiohttpOpenAIChatConfig +from litellm.types.utils import ModelResponse + + +async def _transform(body: object) -> ModelResponse: + raw_response = Mock(spec=ClientResponse) + raw_response.json = AsyncMock(return_value=body) + return await AiohttpOpenAIChatConfig().transform_response( + model="gpt-4o", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=Mock(), + request_data={}, + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +async def test_transform_response_copies_the_openai_body_onto_the_model_response(): + response = await _transform( + { + "id": "chatcmpl-1", + "created": 1700000000, + "model": "gpt-4o-2024", + "object": "chat.completion", + "system_fingerprint": "fp_1", + "choices": [ + {"index": 0, "finish_reason": "length", "message": {"role": "assistant", "content": "Hi"}}, + { + "index": 1, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + ], + }, + }, + ], + } + ) + + assert response.id == "chatcmpl-1" + assert response.created == 1700000000 + assert response.model == "gpt-4o-2024" + assert response.object == "chat.completion" + assert response.system_fingerprint == "fp_1" + assert [(choice.index, choice.finish_reason, choice.message.content) for choice in response.choices] == [ + (0, "length", "Hi"), + (1, "tool_calls", None), + ] + assert response.choices[1].message.tool_calls[0].function.name == "lookup" + + +async def test_transform_response_fills_choice_defaults_for_an_empty_choice(): + response = await _transform({"choices": [{}]}) + + assert [(choice.index, choice.finish_reason, choice.message.role) for choice in response.choices] == [ + (0, "stop", "assistant") + ] + assert response.id is None + + +async def test_transform_response_returns_no_choices_for_an_empty_choices_list(): + response = await _transform({"id": "chatcmpl-1", "choices": []}) + + assert response.choices == [] + assert response.id == "chatcmpl-1" + + +@pytest.mark.parametrize( + "body", + [ + {"id": "chatcmpl-1"}, + {"choices": None}, + {"choices": 7}, + {"choices": "secret-completion"}, + {"choices": {"message": "secret-completion"}}, + {"choices": ["secret-completion"]}, + {"choices": [{"index": 0}, ["secret-completion"]]}, + ], +) +async def test_transform_response_rejects_choices_that_are_not_a_list_of_objects_without_echoing_them(body: object): + with pytest.raises(ValidationError) as exc_info: + await _transform(body) + + assert "secret-completion" not in str(exc_info.value) + + +async def test_transform_response_rejects_a_choice_whose_message_is_not_an_object(): + with pytest.raises(ValidationError, match="validation error for Choices"): + await _transform({"choices": [{"message": "text"}]}) diff --git a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 09921a9a85c..85e4c055f07 100644 --- a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -2489,6 +2489,28 @@ class TestAnthropicMessagesHandlerStreamingScanKey: assert ended_key.tool_calls_in_flight is False assert ended_key != open_key + def test_released_stream_as_ended_keys_the_tool_use_the_client_already_received(self): + handler = AnthropicMessagesHandler() + tool_use = self._sse( + "content_block_start", + { + "type": "content_block_start", + "index": 1, + "content_block": {"type": "tool_use", "id": "toolu_1", "name": "get_weather", "input": {}}, + }, + ) + stopped_key = handler.get_streaming_scan_key([self._text_delta("hi"), tool_use, self._stop("tool_use")]) + released_key = handler.get_streaming_scan_key( + handler.released_stream_as_ended([self._text_delta("hi"), tool_use]) + ) + assert released_key.stream_ended is True + assert released_key == stopped_key + + def test_released_stream_as_ended_leaves_a_text_only_stream_as_released(self): + released = (self._text_delta("hi"), self._text_delta(" there")) + ended = AnthropicMessagesHandler().released_stream_as_ended(released) + assert ended == released and all(a is b for a, b in zip(ended, released, strict=True)) + class PerRowTextGuardrail(CustomGuardrail): """Answers one redacted text per chat row it was shown, the way a guardrail diff --git a/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py b/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py index f385c2f2211..e9736d47537 100644 --- a/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py +++ b/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py @@ -12,6 +12,8 @@ import time import httpx import pytest +from openai.types.file_deleted import FileDeleted +from pydantic import ValidationError from unittest.mock import Mock, patch from litellm.llms.anthropic.files.transformation import ( @@ -196,6 +198,45 @@ class TestAnthropicFilesConfig: assert result.purpose == "messages" assert result.status == "uploaded" + def test_create_file_response_maps_the_anthropic_file_onto_an_openai_file(self) -> None: + result = self.config.transform_create_file_response( + model=None, + raw_response=httpx.Response( + 200, + json={ + "id": "file-abc123", + "type": "file", + "filename": "document.pdf", + "mime_type": "application/pdf", + "size_bytes": 12345, + "created_at": "2025-01-15T10:30:00Z", + }, + ), + logging_obj=Mock(), + litellm_params={}, + ) + + assert result == OpenAIFileObject( + id="file-abc123", + bytes=12345, + created_at=1736937000, + filename="document.pdf", + object="file", + purpose="messages", + status="uploaded", + status_details=None, + ) + + @pytest.mark.parametrize("body", [b'["file-abc123"]', b'"file-abc123"', b"null", b"7"]) + def test_create_file_response_rejects_a_body_that_is_not_a_json_object(self, body: bytes) -> None: + with pytest.raises(ValidationError): + self.config.transform_create_file_response( + model=None, + raw_response=httpx.Response(200, content=body), + logging_obj=Mock(), + litellm_params={}, + ) + def test_transform_retrieve_file_request(self): url, params = self.config.transform_retrieve_file_request( file_id="file-abc123", @@ -245,6 +286,43 @@ class TestAnthropicFilesConfig: assert result.id == "file-abc123" assert result.bytes == 5000 + def test_retrieve_file_response_maps_the_anthropic_file_onto_an_openai_file(self) -> None: + result = self.config.transform_retrieve_file_response( + raw_response=httpx.Response( + 200, + json={ + "id": "file-abc123", + "type": "file", + "filename": "document.pdf", + "mime_type": "application/pdf", + "size_bytes": 5000, + "created_at": "2025-06-01T12:00:00Z", + }, + ), + logging_obj=Mock(), + litellm_params={}, + ) + + assert result == OpenAIFileObject( + id="file-abc123", + bytes=5000, + created_at=1748779200, + filename="document.pdf", + object="file", + purpose="messages", + status="uploaded", + status_details=None, + ) + + @pytest.mark.parametrize("body", [b'["file-abc123"]', b'"file-abc123"', b"null", b"7"]) + def test_retrieve_file_response_rejects_a_body_that_is_not_a_json_object(self, body: bytes) -> None: + with pytest.raises(ValidationError): + self.config.transform_retrieve_file_response( + raw_response=httpx.Response(200, content=body), + logging_obj=Mock(), + litellm_params={}, + ) + def test_transform_delete_file_request(self): url, params = self.config.transform_delete_file_request( file_id="file-abc123", @@ -271,6 +349,35 @@ class TestAnthropicFilesConfig: assert result.deleted is True assert result.object == "file" + @pytest.mark.parametrize( + ("payload", "expected_id"), + [ + ({"id": "file-abc123", "type": "file_deleted"}, "file-abc123"), + ({"id": "file-abc123", "unknown": [1, {"nested": None}]}, "file-abc123"), + ({"type": "error", "error": {"type": "not_found_error", "message": "File not found"}}, ""), + ({}, ""), + ], + ) + def test_delete_file_response_reports_the_id_anthropic_returned(self, payload: object, expected_id: str) -> None: + result = self.config.transform_delete_file_response( + raw_response=httpx.Response(200, json=payload), + logging_obj=Mock(), + litellm_params={}, + ) + + assert result == FileDeleted(id=expected_id, deleted=True, object="file") + + @pytest.mark.parametrize( + "body", [b'["file-abc123"]', b'"file-abc123"', b"null", b"7", b'{"id": null}', b'{"id": 7}'] + ) + def test_delete_file_response_rejects_a_body_without_a_string_id(self, body: bytes) -> None: + with pytest.raises(ValidationError): + self.config.transform_delete_file_response( + raw_response=httpx.Response(200, content=body), + logging_obj=Mock(), + litellm_params={}, + ) + def test_transform_list_files_request(self): url, params = self.config.transform_list_files_request( purpose=None, diff --git a/tests/unit/llms/azure/test_azure_speech_audio_transcription.py b/tests/unit/llms/azure/test_azure_speech_audio_transcription.py index b447645bae8..4330f05e79c 100644 --- a/tests/unit/llms/azure/test_azure_speech_audio_transcription.py +++ b/tests/unit/llms/azure/test_azure_speech_audio_transcription.py @@ -3,6 +3,7 @@ from unittest.mock import MagicMock import httpx import pytest +from pydantic import ValidationError import litellm from litellm.llms.azure.audio_transcription.transformation import ( @@ -226,3 +227,22 @@ def test_azure_speech_transcription_routes_through_provider_config(monkeypatch): assert audio_handler.call_args.kwargs["custom_llm_provider"] == "azure" +def test_azure_speech_audio_transcription_response_keeps_the_raw_payload_as_hidden_params(): + payload = {"RecognitionStatus": "Success", "NBest": [{"Lexical": "hello world", "Confidence": 0.9}], "Offset": 3} + + response = AzureSpeechAudioTranscriptionConfig().transform_audio_transcription_response( + httpx.Response(200, json=payload) + ) + + assert response.text == "hello world" + assert response._hidden_params == payload + + +@pytest.mark.parametrize("payload", [7, "spoken secret", [{"DisplayText": "spoken secret"}]]) +def test_azure_speech_audio_transcription_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + AzureSpeechAudioTranscriptionConfig().transform_audio_transcription_response( + httpx.Response(200, json=payload) + ) + + assert "spoken secret" not in str(exc_info.value) diff --git a/tests/unit/llms/bedrock/image_edit/test_stability_transformation.py b/tests/unit/llms/bedrock/image_edit/test_stability_transformation.py new file mode 100644 index 00000000000..6ca70324162 --- /dev/null +++ b/tests/unit/llms/bedrock/image_edit/test_stability_transformation.py @@ -0,0 +1,14 @@ +from litellm.llms.bedrock.image_edit.stability_transformation import BedrockStabilityImageEditConfig + + +def test_transform_image_edit_request_returns_the_json_body_and_no_files(): + request = BedrockStabilityImageEditConfig().transform_image_edit_request( + model="stability.stable-image-inpaint-v1:0", + prompt="add a red hat", + image=b"\x89PNG-bytes", + image_edit_optional_request_params={}, + litellm_params={}, + headers={}, + ) + + assert request == ({"output_format": "png", "prompt": "add a red hat", "image": "iVBORy1ieXRlcw=="}, {}) diff --git a/tests/unit/llms/bedrock/image_generation/__init__.py b/tests/unit/llms/bedrock/image_generation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/bedrock/image_generation/test_amazon_titan_transformation.py b/tests/unit/llms/bedrock/image_generation/test_amazon_titan_transformation.py new file mode 100644 index 00000000000..eb8b9de6721 --- /dev/null +++ b/tests/unit/llms/bedrock/image_generation/test_amazon_titan_transformation.py @@ -0,0 +1,57 @@ +import pytest + +from litellm.llms.bedrock.image_generation.amazon_titan_transformation import AmazonTitanImageGenerationConfig + + +@pytest.mark.parametrize( + ("non_default_params", "expected"), + [ + ( + {"size": "1024x512", "n": 2, "quality": "hd"}, + {"imageGenerationConfig": {"width": 1024, "height": 512, "numberOfImages": 2, "quality": "premium"}}, + ), + ({"quality": "low"}, {"imageGenerationConfig": {"quality": "standard"}}), + ({"quality": "auto", "size": None, "n": None}, {}), + ], +) +def test_map_openai_params_builds_the_image_generation_config( + non_default_params: dict[str, object], expected: dict[str, object] +): + optional_params = AmazonTitanImageGenerationConfig.map_openai_params( + non_default_params=non_default_params, optional_params={} + ) + + assert optional_params == expected + + +@pytest.mark.parametrize( + ("optional_params", "expected"), + [ + ( + {"imageGenerationConfig": {"width": 512, "numberOfImages": 2}, "negativeText": "blurry"}, + { + "taskType": "TEXT_IMAGE", + "textToImageParams": {"text": "a cat", "negativeText": "blurry"}, + "imageGenerationConfig": {"width": 512, "numberOfImages": 2}, + }, + ), + ( + {"taskType": "COLOR_GUIDED_GENERATION", "negativeText": ""}, + { + "taskType": "COLOR_GUIDED_GENERATION", + "textToImageParams": {"text": "a cat"}, + "imageGenerationConfig": {}, + }, + ), + ], +) +def test_transform_request_body_builds_the_titan_request( + optional_params: dict[str, object], expected: dict[str, object] +): + request_body = AmazonTitanImageGenerationConfig.transform_request_body( + text="a cat", optional_params=optional_params + ) + + assert request_body == expected + assert list(request_body["textToImageParams"]) == list(expected["textToImageParams"]) + assert optional_params == {} diff --git a/tests/unit/llms/bedrock/rerank/test_bedrock_rerank_header_forwarding.py b/tests/unit/llms/bedrock/rerank/test_bedrock_rerank_header_forwarding.py index c40830b238f..495e0f56c84 100644 --- a/tests/unit/llms/bedrock/rerank/test_bedrock_rerank_header_forwarding.py +++ b/tests/unit/llms/bedrock/rerank/test_bedrock_rerank_header_forwarding.py @@ -8,7 +8,9 @@ forward_client_headers_to_llm_api were not being passed to Bedrock rerank provid import json from unittest.mock import AsyncMock, MagicMock, Mock, patch +import httpx import pytest +from pydantic import ValidationError import litellm from litellm.llms.bedrock.base_aws_llm import Boto3CredentialsInfo @@ -506,3 +508,60 @@ async def test_bedrock_rerank_records_llm_api_duration(): assert response._hidden_params["litellm_overhead_time_ms"] is not None assert response._hidden_params["_response_ms"] >= response._hidden_params["litellm_overhead_time_ms"] + + +RERANK_MODEL_ARN = "arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0" +NON_OBJECT_BODIES = [7, "sensitive-document", [{"index": 0, "relevanceScore": 0.9, "id": "sensitive-document"}]] + + +def _rerank_through_transport(payload: object, *, is_async: bool): + transport = httpx.MockTransport(lambda request: httpx.Response(200, json=payload)) + client = ( + AsyncHTTPHandler(transport=transport) if is_async else HTTPHandler(client=httpx.Client(transport=transport)) + ) + return BedrockRerankHandler().rerank( + model=RERANK_MODEL_ARN, + query=test_query, + documents=test_documents, + optional_params={ + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "example-secret", + "aws_region_name": "us-west-2", + }, + logging_obj=Mock(), + _is_async=is_async, + client=client, + ) + + +def _assert_is_the_bedrock_ranking(response: litellm.RerankResponse) -> None: + assert response.results == [ + {"index": 2, "relevance_score": 0.95}, + {"index": 0, "relevance_score": 0.1}, + {"index": 1, "relevance_score": 0.05}, + ] + assert response.meta == {"billed_units": {"search_units": 1}, "tokens": {}} + + +def test_bedrock_rerank_maps_the_upstream_ranking(): + _assert_is_the_bedrock_ranking(_rerank_through_transport(bedrock_rerank_response, is_async=False)) + + +async def test_bedrock_arerank_maps_the_upstream_ranking(): + _assert_is_the_bedrock_ranking(await _rerank_through_transport(bedrock_rerank_response, is_async=True)) + + +@pytest.mark.parametrize("payload", NON_OBJECT_BODIES) +def test_bedrock_rerank_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _rerank_through_transport(payload, is_async=False) + + assert "sensitive-document" not in str(exc_info.value) + + +@pytest.mark.parametrize("payload", NON_OBJECT_BODIES) +async def test_bedrock_arerank_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + await _rerank_through_transport(payload, is_async=True) + + assert "sensitive-document" not in str(exc_info.value) diff --git a/tests/unit/llms/bedrock/search/test_agentcore_search_transformation.py b/tests/unit/llms/bedrock/search/test_agentcore_search_transformation.py index 20bf65ee385..4526c428f1a 100644 --- a/tests/unit/llms/bedrock/search/test_agentcore_search_transformation.py +++ b/tests/unit/llms/bedrock/search/test_agentcore_search_transformation.py @@ -9,10 +9,12 @@ the sharded CI (coverage collection runs against this tree). import json import os +import httpx import pytest from unittest.mock import AsyncMock, patch, MagicMock import litellm +from litellm.llms.base_llm.search.transformation import SearchResponse from litellm.llms.bedrock.search.transformation import ( AGENTCORE_DEFAULT_MCP_PROTOCOL_VERSION, AgentCoreSearchConfig, @@ -635,3 +637,44 @@ class TestAgentCoreSearchEdgeCases: monkeypatch.setattr(litellm, "model_cost", GetModelCostMap.load_local_model_cost_map()) assert search_provider_cost_per_query(model="agentcore/search", custom_llm_provider="agentcore") == (0.0, 0.0) + + +def _transform_text(text: str) -> SearchResponse: + return AgentCoreSearchConfig().transform_search_response( + raw_response=httpx.Response(200, text=text), + logging_obj=MagicMock(), + ) + + +def _as_tuples(response: SearchResponse) -> list[tuple[str, str, str, str | None, str | None]]: + return [(r.title, r.url, r.snippet, r.date, r.last_updated) for r in response.results] + + +EXPECTED_MCP_RESULTS = [ + ("Test Result 1", "https://example.com/1", "Snippet for result 1", "2026-06-16", None), + ("Test Result 2", "https://example.com/2", "Snippet for result 2", None, None), +] + + +@pytest.mark.parametrize( + "text", + [ + json.dumps(_mcp_response_body()), + f"event: message\ndata: {json.dumps(_mcp_response_body())}\n\n", + f"data: 5\n\ndata: [1]\n\ndata: {json.dumps(_mcp_response_body())}\n\n", + json.dumps({"result": {"content": [{"type": "text", "text": json.dumps({"results": MCP_RESULTS})}]}}), + json.dumps({"result": {"content": [{"type": "text", "text": "prose"}], "structuredContent": MCP_RESULTS}}), + ], +) +def test_transform_search_response_reads_results_from_a_real_http_response(text: str): + assert _as_tuples(_transform_text(text)) == EXPECTED_MCP_RESULTS + + +@pytest.mark.parametrize( + "block_text", + ["null", "5", "true", '"text"', "{}", '{"results": null}', '{"results": "text"}', '["scalar", 5, null]'], +) +def test_transform_search_response_ignores_text_blocks_without_result_objects(block_text: str): + body = {"result": {"content": [{"type": "text", "text": block_text}]}} + + assert _transform_text(json.dumps(body)).results == [] diff --git a/tests/unit/llms/black_forest_labs/image_generation/test_bfl_image_generation_transformation.py b/tests/unit/llms/black_forest_labs/image_generation/test_bfl_image_generation_transformation.py index d6e2c4a3e06..c947967417d 100644 --- a/tests/unit/llms/black_forest_labs/image_generation/test_bfl_image_generation_transformation.py +++ b/tests/unit/llms/black_forest_labs/image_generation/test_bfl_image_generation_transformation.py @@ -10,6 +10,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest +from pydantic import ValidationError from litellm.llms.black_forest_labs.image_generation.transformation import ( @@ -355,3 +356,48 @@ class TestBlackForestLabsImageGenerationTransformation: config = get_black_forest_labs_image_generation_config("flux-pro-1.1") assert isinstance(config, BlackForestLabsImageGenerationConfig) + + +def _transform(payload: object) -> ImageResponse: + return BlackForestLabsImageGenerationConfig().transform_image_generation_response( + model="flux-pro-1.1", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize( + ("result", "expected_urls"), + [ + ({"sample": "https://bfl.example/a.png"}, ["https://bfl.example/a.png"]), + ( + ["https://bfl.example/a.png", {"url": "https://bfl.example/b.png"}, {"seed": 1}, 7], + ["https://bfl.example/a.png", "https://bfl.example/b.png"], + ), + ], +) +def test_transform_image_generation_response_reads_urls_from_the_result(result: object, expected_urls: list[str]): + response = _transform({"status": "Ready", "result": result}) + + assert [image.url for image in response.data] == expected_urls + + +@pytest.mark.parametrize("payload", [{}, {"result": None}, {"result": "https://bfl.example/a.png"}, {"result": []}]) +def test_transform_image_generation_response_without_a_url_is_a_provider_error(payload: dict[str, object]): + with pytest.raises(BlackForestLabsError) as exc_info: + _transform(payload) + + assert exc_info.value.status_code == 500 + + +@pytest.mark.parametrize("payload", [7, "https://bfl.example/a.png", [{"sample": "https://bfl.example/a.png"}]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "bfl.example" not in str(exc_info.value) diff --git a/tests/unit/llms/clarifai/__init__.py b/tests/unit/llms/clarifai/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/clarifai/chat/__init__.py b/tests/unit/llms/clarifai/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/clarifai/chat/test_transformation.py b/tests/unit/llms/clarifai/chat/test_transformation.py new file mode 100644 index 00000000000..e81b2a9558d --- /dev/null +++ b/tests/unit/llms/clarifai/chat/test_transformation.py @@ -0,0 +1,70 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.clarifai.chat.transformation import ClarifaiConfig +from litellm.llms.openai.common_utils import OpenAIError +from litellm.types.utils import ModelResponse + + +def _transform(raw_response: httpx.Response) -> ModelResponse: + return ClarifaiConfig().transform_response( + model="user.app.model", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=Mock(), + request_data={}, + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_response_builds_the_model_response_from_the_body(): + response = _transform( + httpx.Response( + 200, + json={ + "id": "chatcmpl-1", + "created": 1700000000, + "model": "upstream-model", + "system_fingerprint": "fp_1", + "choices": [{"index": 0, "finish_reason": "length", "message": {"role": "assistant", "content": "Hi"}}], + "usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, + "vendor_field": {"kept": True}, + }, + ) + ) + + assert response.id == "chatcmpl-1" + assert response.created == 1700000000 + assert response.model == "clarifai/user.app.model" + assert response.system_fingerprint == "fp_1" + assert [(choice.finish_reason, choice.message.content) for choice in response.choices] == [("length", "Hi")] + assert response.usage.model_dump()["total_tokens"] == 5 + assert response.vendor_field == {"kept": True} + + +def test_transform_response_keeps_a_missing_model_unset(): + response = _transform(httpx.Response(200, json={"choices": [{"message": {"content": "Hi"}}]})) + + assert response.model is None + assert response.choices[0].message.content == "Hi" + + +@pytest.mark.parametrize("body", [b'["prompt-text"]', b'"prompt-text"', b"7", b"null"]) +def test_transform_response_rejects_a_body_that_is_not_an_object_without_echoing_it(body: bytes): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, content=body)) + + assert "prompt-text" not in str(exc_info.value) + + +def test_transform_response_reports_an_unparseable_body_as_an_openai_error(): + with pytest.raises(OpenAIError, match="Failed to parse Clarifai response") as exc_info: + _transform(httpx.Response(502, content=b"bad gateway")) + + assert exc_info.value.status_code == 502 diff --git a/tests/unit/llms/claude_code/harness/test_transformation.py b/tests/unit/llms/claude_code/harness/test_transformation.py index 6348b4f6a3e..1ae4497ef18 100644 --- a/tests/unit/llms/claude_code/harness/test_transformation.py +++ b/tests/unit/llms/claude_code/harness/test_transformation.py @@ -318,6 +318,29 @@ def test_parse_skips_subagent_messages_and_maps_errors(): ] +def test_user_message_yields_only_its_tool_result_objects(): + line = { + "type": "user", + "message": { + "content": [ + "plain text", + None, + ["tool_result"], + {"type": "text", "text": "not a result"}, + {"type": "tool_result", "tool_use_id": "t1", "content": "ok"}, + {"type": "tool_result", "content": [{"type": "text", "text": "a"}, "b"], "is_error": 1}, + {"type": "tool_result", "tool_use_id": 7, "content": None, "extra": {"kept": "out"}}, + ] + }, + } + + assert ClaudeCodeHarnessConfig().transform_stream_line(line, ClaudeCodeStreamState()) == [ + ToolResult(id="t1", output="ok", is_error=False), + ToolResult(id="", output="a\nb", is_error=True), + ToolResult(id="7", output="", is_error=False), + ] + + def test_parse_thinking_and_mcp_tools(): state = ClaudeCodeStreamState() msg = { diff --git a/tests/unit/llms/codex/harness/test_transformation.py b/tests/unit/llms/codex/harness/test_transformation.py index 6371f74fba6..6df322db1ae 100644 --- a/tests/unit/llms/codex/harness/test_transformation.py +++ b/tests/unit/llms/codex/harness/test_transformation.py @@ -11,7 +11,7 @@ from pathlib import Path from typing import Optional import pytest -from pydantic import BaseModel +from pydantic import BaseModel, ValidationError from litellm.harness.context import SessionContext from litellm.harness.errors import HarnessError, HarnessInstallFailed, OptionsMismatch @@ -272,6 +272,36 @@ def test_parse_file_change_and_mcp_and_web_search(): assert events[0].name == "web_search" and events[0].input == {"query": "litellm"} +@pytest.mark.parametrize( + ("changes", "output"), + [ + ([{"path": "a.txt", "kind": "add"}, {"path": "b.txt", "kind": "update", "diff": "@@"}], "add a.txt\nupdate b.txt"), + ([{"path": "only-path.txt"}, {"kind": "delete"}, {}], "only-path.txt\ndelete\n"), + ([], ""), + (None, ""), + ], +) +def test_completed_file_change_lists_each_change(changes: object, output: str): + state = CodexStreamState(started={"i1"}) + item = {"id": "i1", "type": "file_change", "changes": changes, "status": "completed"} + + assert parse_event({"type": "item.completed", "item": item}, state) == [ + ToolResult(id="i1", output=output, is_error=False) + ] + + +@pytest.mark.parametrize( + "changes", + [["a.txt"], [{"path": "a.txt"}, None], "a.txt", {"a.txt": {"kind": "add"}}, 7], +) +def test_completed_file_change_rejects_changes_that_are_not_a_list_of_objects(changes: object): + state = CodexStreamState(started={"i1"}) + item = {"id": "i1", "type": "file_change", "changes": changes, "status": "completed"} + + with pytest.raises(ValidationError): + parse_event({"type": "item.completed", "item": item}, state) + + def test_parse_failed_command_is_error_and_unknown_events_ignored(): state = CodexStreamState() item = { diff --git a/tests/unit/llms/cohere/rerank/test_transformation.py b/tests/unit/llms/cohere/rerank/test_transformation.py new file mode 100644 index 00000000000..ac5acf08d0d --- /dev/null +++ b/tests/unit/llms/cohere/rerank/test_transformation.py @@ -0,0 +1,72 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.cohere.rerank.transformation import CohereRerankConfig +from litellm.types.rerank import RerankResponse + + +def _transform(payload: object) -> RerankResponse: + return CohereRerankConfig().transform_rerank_response( + model="rerank-english-v3.0", + raw_response=httpx.Response(200, json=payload), + model_response=RerankResponse(), + logging_obj=Mock(), + ) + + +def test_transform_rerank_response_keeps_the_cohere_payload(): + response = _transform( + { + "id": "rerank-1", + "results": [{"index": 1, "relevance_score": 0.9, "document": {"text": "sensitive-document"}}], + "meta": {"billed_units": {"search_units": 1}, "tokens": {"input_tokens": 4}}, + "warnings": ["ignored"], + } + ) + + assert response.model_dump() == { + "id": "rerank-1", + "results": [{"index": 1, "relevance_score": 0.9, "document": {"text": "sensitive-document"}}], + "meta": {"billed_units": {"search_units": 1}, "tokens": {"input_tokens": 4}}, + } + + +def test_transform_rerank_response_of_an_empty_object_has_no_results(): + assert _transform({}).model_dump() == {"id": None, "results": None, "meta": None} + + +@pytest.mark.parametrize("payload", [7, "sensitive-document", [{"id": "sensitive-document"}]]) +def test_transform_rerank_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "sensitive-document" not in str(exc_info.value) + + +@pytest.mark.parametrize("payload", [{"id": 7}, {"results": "not-a-list"}, {"results": [{"index": 0}]}, {"meta": []}]) +def test_transform_rerank_response_rejects_malformed_fields(payload: dict[str, object]): + with pytest.raises(ValidationError): + _transform(payload) + + +def test_map_cohere_rerank_params_returns_every_cohere_param(): + params = CohereRerankConfig().map_cohere_rerank_params( + non_default_params=None, + model="rerank-english-v3.0", + drop_params=False, + query="capital of france", + documents=["Paris", {"text": "Berlin"}], + top_n=1, + ) + + assert params == { + "query": "capital of france", + "documents": ["Paris", {"text": "Berlin"}], + "top_n": 1, + "rank_fields": None, + "return_documents": True, + "max_chunks_per_doc": None, + } diff --git a/tests/unit/llms/cometapi/image_generation/__init__.py b/tests/unit/llms/cometapi/image_generation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/cometapi/image_generation/test_transformation.py b/tests/unit/llms/cometapi/image_generation/test_transformation.py new file mode 100644 index 00000000000..0a14b955339 --- /dev/null +++ b/tests/unit/llms/cometapi/image_generation/test_transformation.py @@ -0,0 +1,54 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.cometapi.image_generation.transformation import CometAPIImageGenerationConfig +from litellm.types.utils import ImageResponse + + +def _transform(payload: object) -> ImageResponse: + return CometAPIImageGenerationConfig().transform_image_generation_response( + model="dall-e-3", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_each_image_object(): + response = _transform({"created": 1, "data": [{"url": "https://img.cometapi.com/a.png"}, {"b64_json": "QUJD"}, {}]}) + + assert [(image.url, image.b64_json) for image in response.data] == [ + ("https://img.cometapi.com/a.png", None), + (None, "QUJD"), + (None, None), + ] + + +@pytest.mark.parametrize("payload", [{}, {"created": 1}, {"data": []}, {"data": ""}, {"data": {}}, []]) +def test_transform_image_generation_response_without_images_has_no_data(payload: object): + assert _transform(payload).data == [] + + +@pytest.mark.parametrize( + "payload", + [ + ["data", "https://img.cometapi.com/a.png"], + {"data": None}, + {"data": "https://img.cometapi.com/a.png"}, + {"data": {"url": "https://img.cometapi.com/a.png"}}, + {"data": ["https://img.cometapi.com/a.png"]}, + {"data": [{"url": "https://img.cometapi.com/a.png"}, 7]}, + ], +) +def test_transform_image_generation_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "img.cometapi.com" not in str(exc_info.value) diff --git a/tests/unit/llms/dashscope/test_dashscope_rerank_transformation.py b/tests/unit/llms/dashscope/test_dashscope_rerank_transformation.py index 4466e5b8767..a913c79af31 100644 --- a/tests/unit/llms/dashscope/test_dashscope_rerank_transformation.py +++ b/tests/unit/llms/dashscope/test_dashscope_rerank_transformation.py @@ -7,6 +7,7 @@ from unittest.mock import MagicMock import httpx import pytest +from pydantic import ValidationError from litellm.llms.dashscope.common_utils import DashScopeError @@ -347,3 +348,89 @@ class TestProviderConfigManagerDispatch: present_version_params=[], ) assert isinstance(cfg, DashScopeRerankConfig) + + +def _transform(payload: object, status_code: int = 200) -> RerankResponse: + return DashScopeRerankConfig().transform_rerank_response( + model="qwen3-rerank", + raw_response=httpx.Response(status_code, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + ) + + +@pytest.mark.parametrize( + ("usage", "expected_total_tokens"), + [ + ({"total_tokens": 79}, 79), + ({}, None), + (None, None), + ], +) +def test_transform_rerank_response_reads_total_tokens_from_usage(usage: object, expected_total_tokens: int | None): + response = _transform({"id": "rerank-1", "results": [{"index": 0, "relevance_score": 0.5}], "usage": usage}) + + assert response.meta == { + "billed_units": {"total_tokens": expected_total_tokens}, + "tokens": {"input_tokens": expected_total_tokens}, + } + + +def test_transform_rerank_response_keeps_provider_id_and_drops_unknown_result_fields(): + response = _transform( + { + "id": "rerank-1", + "results": [{"index": 2, "relevance_score": 0.25, "document": {"text": "doc", "extra": 1}, "extra": 2}], + } + ) + + assert response.id == "rerank-1" + assert response.results == [{"index": 2, "relevance_score": 0.25, "document": {"text": "doc"}}] + + +def test_transform_rerank_response_empty_results_list_yields_no_results(): + assert _transform({"id": "rerank-1", "results": []}).results == [] + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"results": 7}, + {"results": ["not an object"]}, + {"results": [{"index": 0, "relevance_score": 0.5}], "usage": "seventy nine"}, + {"results": [{"index": 0, "relevance_score": 0.5}], "usage": {"total_tokens": 1.5}}, + {"results": [{"index": 0, "relevance_score": 0.5}], "id": 7}, + ], +) +def test_transform_rerank_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError): + _transform(payload) + + +def test_transform_rerank_response_result_without_index_raises_key_error(): + with pytest.raises(KeyError, match="index"): + _transform({"results": [{"relevance_score": 0.5}]}) + + +def test_transform_rerank_response_error_envelope_raises_with_the_provider_status_and_message(): + with pytest.raises(DashScopeError) as exc_info: + _transform({"code": "Throttling", "message": "slow down"}, status_code=429) + + assert exc_info.value.status_code == 429 + assert exc_info.value.message == "slow down" + + +def test_transform_rerank_response_without_results_raises_dashscope_error_naming_the_body(): + with pytest.raises(DashScopeError) as exc_info: + _transform({"id": "rerank-1"}, status_code=502) + + assert exc_info.value.status_code == 502 + assert exc_info.value.message == "No results in DashScope rerank response: {'id': 'rerank-1'}" + + +def test_transform_rerank_response_shape_errors_do_not_echo_the_payload(): + with pytest.raises(ValidationError) as exc_info: + _transform({"results": [["leaked document text"]]}) + + assert "leaked document text" not in str(exc_info.value) diff --git a/tests/unit/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py b/tests/unit/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py index 7c1f5256deb..366413df1c8 100644 --- a/tests/unit/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py +++ b/tests/unit/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py @@ -3,6 +3,7 @@ import os import pathlib from unittest.mock import MagicMock +import httpx import pytest @@ -525,3 +526,82 @@ def test_reconstruct_diarized_transcript_multiple_speaker_changes(): assert "Hello" in result assert "back" in result assert "Thanks" in result + + +def _deepgram_payload(alternative: dict[str, object], channel_fields: dict[str, object]) -> dict[str, object]: + return { + "metadata": {"duration": 2.5}, + "results": {"channels": [{"alternatives": [alternative], **channel_fields}]}, + } + + +def _transform_deepgram_response(payload: object) -> TranscriptionResponse: + return DeepgramAudioTranscriptionConfig().transform_audio_transcription_response( + httpx.Response(200, json=payload) + ) + + +@pytest.mark.parametrize( + ("channel_fields", "expected_language"), + [ + ({}, "en"), + ({"detected_language": None}, "en"), + ({"detected_language": ""}, "en"), + ({"detected_language": "fr"}, "fr"), + ({"detected_language": "de", "language_confidence": 0.98}, "de"), + ({"detected_language": ["es", "en"]}, ["es", "en"]), + ], +) +def test_transform_response_reports_the_detected_language_or_english( + channel_fields: dict[str, object], expected_language: object +): + response = _transform_deepgram_response(_deepgram_payload({"transcript": "bonjour"}, channel_fields)) + + assert response.text == "bonjour" + assert response["language"] == expected_language + assert response["duration"] == 2.5 + assert "words" not in response + + +@pytest.mark.parametrize( + ("words", "expected"), + [ + ([], []), + ("", []), + ({}, []), + ( + [{"word": "hello", "start": 0.0, "end": 0.5, "confidence": 0.9}], + [{"word": "hello", "start": 0.0, "end": 0.5}], + ), + ( + [{"word": None, "start": "0.1", "end": [2]}, {"word": "b", "start": 1, "end": 2}], + [{"word": None, "start": "0.1", "end": [2]}, {"word": "b", "start": 1, "end": 2}], + ), + ], +) +def test_transform_response_maps_words_to_openai_word_timestamps(words: object, expected: list[dict[str, object]]): + payload = _deepgram_payload({"transcript": "hello", "words": words}, {}) + + response = _transform_deepgram_response(payload) + + assert response["words"] == expected + assert response._hidden_params == payload + + +@pytest.mark.parametrize( + "words", + [ + "not a list", + ["not an object"], + [{"word": "hello", "start": 0.0, "end": 0.5}, None], + [{"word": "hello", "start": 0.0, "end": 0.5}, ["nested"]], + [{"word": "hello", "start": 0.0}], + ], +) +def test_transform_response_wraps_malformed_words_with_the_raw_body(words: object): + raw_response = httpx.Response(200, json=_deepgram_payload({"transcript": "hello", "words": words}, {})) + + with pytest.raises(ValueError, match="Error transforming Deepgram response: ") as exc_info: + DeepgramAudioTranscriptionConfig().transform_audio_transcription_response(raw_response) + + assert str(exc_info.value).endswith(f"\nResponse: {raw_response.text}") diff --git a/tests/unit/llms/e2b/__init__.py b/tests/unit/llms/e2b/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/e2b/sandbox/__init__.py b/tests/unit/llms/e2b/sandbox/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/e2b/sandbox/test_transformation.py b/tests/unit/llms/e2b/sandbox/test_transformation.py new file mode 100644 index 00000000000..44c0626bc27 --- /dev/null +++ b/tests/unit/llms/e2b/sandbox/test_transformation.py @@ -0,0 +1,121 @@ +import json + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.e2b.sandbox.transformation import E2BSandboxConfig + + +def _client_answering(body: bytes) -> AsyncHTTPHandler: + return AsyncHTTPHandler(transport=httpx.MockTransport(lambda request: httpx.Response(200, content=body))) + + +def _lines(*messages: object) -> list[str]: + return [json.dumps(message) for message in messages] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("created", "expected_domain", "expected_tokens"), + [ + ( + {"sandboxID": "sbx_1", "domain": "eu.e2b.app", "envdAccessToken": "envd", "trafficAccessToken": "traffic"}, + "eu.e2b.app", + ("envd", "traffic"), + ), + ({"sandboxID": "sbx_1"}, "e2b.app", (None, None)), + ({"sandboxID": "sbx_1", "domain": None, "envdAccessToken": "envd"}, "e2b.app", ("envd", None)), + ({"sandboxID": "sbx_1", "domain": ""}, "e2b.app", (None, None)), + ], +) +async def test_acreate_sandbox_builds_the_handle_from_the_create_response( + created: dict[str, object], expected_domain: str, expected_tokens: tuple[str | None, str | None] +): + handle = await E2BSandboxConfig().acreate_sandbox( + api_key="e2b_key", client=_client_answering(json.dumps(created).encode()) + ) + + assert (handle.id, handle.provider, handle.domain) == ("sbx_1", "e2b", expected_domain) + assert handle._hidden_params == { + "envd_access_token": expected_tokens[0], + "traffic_access_token": expected_tokens[1], + "api_key": "e2b_key", + "api_base": "https://api.e2b.app", + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("body", [b'["secret-token"]', b'"secret-token"', b"7", b"null"]) +async def test_acreate_sandbox_rejects_a_create_response_that_is_not_an_object_without_echoing_it(body: bytes): + with pytest.raises(ValidationError) as exc_info: + await E2BSandboxConfig().acreate_sandbox(api_key="e2b_key", client=_client_answering(body)) + + assert "secret-token" not in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_acreate_sandbox_requires_a_sandbox_id(): + with pytest.raises(KeyError, match="sandboxID"): + await E2BSandboxConfig().acreate_sandbox(api_key="e2b_key", client=_client_answering(b'{"domain": "e2b.app"}')) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("created", [{"sandboxID": 5}, {"sandboxID": "sbx_1", "domain": 5}]) +async def test_acreate_sandbox_rejects_non_string_handle_fields(created: dict[str, object]): + with pytest.raises(ValidationError, match="ContainerHandle"): + await E2BSandboxConfig().acreate_sandbox( + api_key="e2b_key", client=_client_answering(json.dumps(created).encode()) + ) + + +def test_parse_lines_maps_every_message_type_onto_the_result(): + result = E2BSandboxConfig._parse_lines( + [ + *_lines( + {"type": "stdout", "text": "1\n", "timestamp": 1}, + {"type": "stderr", "text": "warn\n"}, + {"type": "stdout"}, + {"type": "result", "png": "BASE64", "is_main_result": True}, + {"type": "error", "name": "ValueError", "value": "bad", "traceback": "tb", "ignored": 1}, + {"type": "number_of_executions", "execution_count": 3}, + {"type": "stdout", "text": "2\n"}, + None, + ), + "", + "not json", + ] + ) + + assert result.model_dump() == { + "stdout": "1\n2\n", + "stderr": "warn\n", + "results": [{"png": "BASE64", "is_main_result": True}], + "error": {"name": "ValueError", "value": "bad", "traceback": "tb"}, + "execution_count": 3, + "object": "code_execution", + } + + +def test_parse_lines_rejects_an_execution_count_that_is_not_a_number(): + with pytest.raises(ValidationError, match="execution_count"): + E2BSandboxConfig._parse_lines(_lines({"type": "number_of_executions", "execution_count": "many"})) + + +@pytest.mark.parametrize( + "message", + [ + ["secret-output"], + "secret-output", + 7, + False, + {"type": "stdout", "text": ["secret-output"]}, + {"type": "stderr", "text": {"secret-output": 1}}, + ], +) +def test_parse_lines_rejects_malformed_messages_without_echoing_them(message: object): + with pytest.raises(ValidationError) as exc_info: + E2BSandboxConfig._parse_lines(_lines({"type": "stdout", "text": "ok"}, message)) + + assert "secret-output" not in str(exc_info.value) diff --git a/tests/unit/llms/elevenlabs/audio_transcription/__init__.py b/tests/unit/llms/elevenlabs/audio_transcription/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/elevenlabs/audio_transcription/test_transformation.py b/tests/unit/llms/elevenlabs/audio_transcription/test_transformation.py new file mode 100644 index 00000000000..52a10f0a327 --- /dev/null +++ b/tests/unit/llms/elevenlabs/audio_transcription/test_transformation.py @@ -0,0 +1,100 @@ +import httpx +import pytest + +from litellm.llms.elevenlabs.audio_transcription.transformation import ElevenLabsAudioTranscriptionConfig +from litellm.types.utils import TranscriptionResponse + + +def _transform(payload: object) -> TranscriptionResponse: + return ElevenLabsAudioTranscriptionConfig().transform_audio_transcription_response( + raw_response=httpx.Response(200, json=payload) + ) + + +def test_transform_audio_transcription_response_keeps_only_spoken_words(): + payload = { + "language_code": "en", + "text": "Hello world", + "words": [ + {"type": "word", "text": "Hello", "start": 0.0, "end": 0.4, "speaker_id": "speaker_0"}, + {"type": "spacing", "text": " ", "start": 0.4, "end": 0.5}, + {"type": "audio_event", "text": "(laughter)", "start": 0.5, "end": 0.9}, + {"type": "word", "text": "world", "start": 0.9, "end": 1.3}, + ], + } + + response = _transform(payload) + + assert response.text == "Hello world" + assert response["task"] == "transcribe" + assert response["language"] == "en" + assert response["words"] == [ + {"word": "Hello", "start": 0.0, "end": 0.4}, + {"word": "world", "start": 0.9, "end": 1.3}, + ] + assert response._hidden_params == payload + + +@pytest.mark.parametrize( + ("word", "expected"), + [ + ({"type": "word"}, [{"word": "", "start": 0, "end": 0}]), + ({"type": "word", "text": None, "start": None, "end": None}, [{"word": None, "start": None, "end": None}]), + ({"type": "word", "text": 7, "start": "0.1", "end": [2]}, [{"word": 7, "start": "0.1", "end": [2]}]), + ({"text": "untyped"}, []), + ({"type": None, "text": "untyped"}, []), + ({}, []), + ], +) +def test_transform_audio_transcription_response_maps_one_word( + word: dict[str, object], expected: list[dict[str, object]] +): + assert _transform({"text": "t", "words": [word]})["words"] == expected + + +@pytest.mark.parametrize("words", [[], "", {}]) +def test_transform_audio_transcription_response_with_empty_words_has_empty_word_list(words: object): + assert _transform({"text": "t", "words": words})["words"] == [] + + +@pytest.mark.parametrize( + ("payload", "expected_text", "expected_language"), + [ + ({}, "", "unknown"), + ({"text": None, "language_code": None}, None, None), + ({"text": "bonjour", "language_code": "fr"}, "bonjour", "fr"), + ({"text": "hola", "language_code": ["es"]}, "hola", ["es"]), + ], +) +def test_transform_audio_transcription_response_without_words_key_has_no_word_list( + payload: dict[str, object], expected_text: str | None, expected_language: object +): + response = _transform(payload) + + assert response.text == expected_text + assert response["language"] == expected_language + assert "words" not in response + assert response._hidden_params == payload + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + "plain text", + 7, + {"text": ["not", "text"]}, + {"text": "t", "words": None}, + {"text": "t", "words": 7}, + {"text": "t", "words": "not a list"}, + {"text": "t", "words": ["not an object"]}, + {"text": "t", "words": [{"type": "word", "text": "ok"}, None]}, + ], +) +def test_transform_audio_transcription_response_wraps_malformed_payloads_with_the_raw_body(payload: object): + raw_response = httpx.Response(200, json=payload) + + with pytest.raises(ValueError, match="Error transforming ElevenLabs response: ") as exc_info: + ElevenLabsAudioTranscriptionConfig().transform_audio_transcription_response(raw_response=raw_response) + + assert str(exc_info.value).endswith(f"\nResponse: {raw_response.text}") diff --git a/tests/unit/llms/fal_ai/image_generation/test_flux_pro_v11_ultra_transformation.py b/tests/unit/llms/fal_ai/image_generation/test_flux_pro_v11_ultra_transformation.py new file mode 100644 index 00000000000..960d927851e --- /dev/null +++ b/tests/unit/llms/fal_ai/image_generation/test_flux_pro_v11_ultra_transformation.py @@ -0,0 +1,56 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.fal_ai.image_generation.flux_pro_v11_ultra_transformation import FalAIFluxProV11UltraConfig +from litellm.types.utils import ImageResponse + + +def _transform(payload: object) -> ImageResponse: + return FalAIFluxProV11UltraConfig().transform_image_generation_response( + model="fal-ai/flux-pro/v1.1-ultra", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_images_and_metadata(): + response = _transform( + { + "images": [{"url": "https://fal.media/a.png", "width": 2752, "height": 1536}, "https://fal.media/b.png"], + "seed": 42, + "timings": {"inference": 2.5}, + "has_nsfw_concepts": [False, False], + } + ) + + assert [image.url for image in response.data] == ["https://fal.media/a.png", "https://fal.media/b.png"] + assert response.data[0].provider_specific_fields == {"width": 2752, "height": 1536} + assert response._hidden_params["seed"] == 42 + assert response._hidden_params["timings"] == {"inference": 2.5} + assert response._hidden_params["has_nsfw_concepts"] == [False, False] + + +@pytest.mark.parametrize("payload", [{}, {"images": None}, {"images": []}]) +def test_transform_image_generation_response_without_images_is_empty(payload: dict[str, object]): + response = _transform(payload) + + assert response.data == [] + assert "seed" not in response._hidden_params + assert "timings" not in response._hidden_params + assert "has_nsfw_concepts" not in response._hidden_params + + +@pytest.mark.parametrize("payload", [7, "https://fal.media/a.png", [{"url": "https://fal.media/a.png"}]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "fal.media" not in str(exc_info.value) diff --git a/tests/unit/llms/fal_ai/image_generation/test_ideogram_v3_transformation.py b/tests/unit/llms/fal_ai/image_generation/test_ideogram_v3_transformation.py new file mode 100644 index 00000000000..716440004cd --- /dev/null +++ b/tests/unit/llms/fal_ai/image_generation/test_ideogram_v3_transformation.py @@ -0,0 +1,71 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.fal_ai.image_generation.ideogram_v3_transformation import FalAIIdeogramV3Config +from litellm.types.utils import ImageResponse + + +def _transform(payload: object) -> ImageResponse: + return FalAIIdeogramV3Config().transform_image_generation_response( + model="fal-ai/ideogram/v3", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_files_and_seed(): + response = _transform( + {"images": [{"url": "https://fal.media/a.png", "file_name": "a.png"}, "https://fal.media/b.png"], "seed": 42} + ) + + assert [image.url for image in response.data] == ["https://fal.media/a.png", "https://fal.media/b.png"] + assert [image.b64_json for image in response.data] == [None, None] + assert response._hidden_params["seed"] == 42 + + +@pytest.mark.parametrize("payload", [{}, {"images": None}, {"images": "https://fal.media/a.png"}]) +def test_transform_image_generation_response_without_an_image_list_is_empty(payload: dict[str, object]): + response = _transform(payload) + + assert response.data == [] + assert "seed" not in response._hidden_params + + +@pytest.mark.parametrize("payload", [7, "https://fal.media/a.png", [{"url": "https://fal.media/a.png"}]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "fal.media" not in str(exc_info.value) + + +@pytest.mark.parametrize( + ("size", "expected"), + [ + ("1024x1024", "square_hd"), + (" 1536x1024 ", "landscape_16_9"), + ("640x480", {"width": 640, "height": 480}), + ("wide", "square_hd"), + ("axb", "square_hd"), + ({"width": 640, "height": 480, "unit": "px"}, {"width": 640, "height": 480}), + ({"width": "640"}, {"width": "640"}), + (512, 512), + ], +) +def test_map_openai_params_translates_size_to_image_size(size: object, expected: object): + optional_params = FalAIIdeogramV3Config().map_openai_params( + non_default_params={"size": size, "n": 2, "response_format": "url"}, + optional_params={}, + model="fal-ai/ideogram/v3", + drop_params=False, + ) + + assert optional_params == {"image_size": expected, "num_images": 2} diff --git a/tests/unit/llms/fal_ai/image_generation/test_recraft_v3_transformation.py b/tests/unit/llms/fal_ai/image_generation/test_recraft_v3_transformation.py new file mode 100644 index 00000000000..09591fe24d5 --- /dev/null +++ b/tests/unit/llms/fal_ai/image_generation/test_recraft_v3_transformation.py @@ -0,0 +1,64 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.fal_ai.image_generation.recraft_v3_transformation import FalAIRecraftV3Config +from litellm.types.utils import ImageResponse + + +def _transform(payload: object) -> ImageResponse: + return FalAIRecraftV3Config().transform_image_generation_response( + model="fal-ai/recraft/v3/text-to-image", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_file_objects_and_bare_urls(): + response = _transform( + {"images": [{"url": "https://fal.media/a.png", "content_type": "image/png"}, "https://fal.media/b.png", 7]} + ) + + assert [image.url for image in response.data] == ["https://fal.media/a.png", "https://fal.media/b.png"] + assert [image.b64_json for image in response.data] == [None, None] + + +@pytest.mark.parametrize("payload", [{}, {"images": None}, {"images": "https://fal.media/a.png"}]) +def test_transform_image_generation_response_without_an_image_list_is_empty(payload: dict[str, object]): + assert _transform(payload).data == [] + + +@pytest.mark.parametrize("payload", [7, "https://fal.media/a.png", [{"url": "https://fal.media/a.png"}]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "fal.media" not in str(exc_info.value) + + +@pytest.mark.parametrize( + ("size", "expected"), + [ + ("1024x1024", "square_hd"), + ("576x1024", "portrait_16_9"), + ("640x480", {"width": 640, "height": 480}), + ("wide", "square_hd"), + ("axb", "square_hd"), + ], +) +def test_map_openai_params_translates_size_to_image_size(size: str, expected: object): + optional_params = FalAIRecraftV3Config().map_openai_params( + non_default_params={"size": size, "n": 2, "response_format": "url"}, + optional_params={}, + model="fal-ai/recraft/v3/text-to-image", + drop_params=False, + ) + + assert optional_params == {"image_size": expected} diff --git a/tests/unit/llms/fal_ai/image_generation/test_stable_diffusion_transformation.py b/tests/unit/llms/fal_ai/image_generation/test_stable_diffusion_transformation.py new file mode 100644 index 00000000000..8f112374efd --- /dev/null +++ b/tests/unit/llms/fal_ai/image_generation/test_stable_diffusion_transformation.py @@ -0,0 +1,76 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.fal_ai.image_generation.stable_diffusion_transformation import FalAIStableDiffusionConfig +from litellm.types.utils import ImageResponse + + +def _transform(payload: object) -> ImageResponse: + return FalAIStableDiffusionConfig().transform_image_generation_response( + model="fal-ai/stable-diffusion-v35-medium", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_images_and_metadata(): + response = _transform( + { + "images": [{"url": "https://fal.media/a.png", "width": 1024}, "https://fal.media/b.png", 7], + "seed": 42, + "timings": {"inference": 2.5}, + "has_nsfw_concepts": [False, False], + } + ) + + assert [image.url for image in response.data] == ["https://fal.media/a.png", "https://fal.media/b.png"] + assert response._hidden_params["seed"] == 42 + assert response._hidden_params["timings"] == {"inference": 2.5} + assert response._hidden_params["has_nsfw_concepts"] == [False, False] + + +@pytest.mark.parametrize("payload", [{}, {"images": None}, {"images": "https://fal.media/a.png"}]) +def test_transform_image_generation_response_without_an_image_list_is_empty(payload: dict[str, object]): + response = _transform(payload) + + assert response.data == [] + assert "seed" not in response._hidden_params + assert "timings" not in response._hidden_params + assert "has_nsfw_concepts" not in response._hidden_params + + +@pytest.mark.parametrize("payload", [7, "https://fal.media/a.png", [{"url": "https://fal.media/a.png"}]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "fal.media" not in str(exc_info.value) + + +@pytest.mark.parametrize( + ("size", "expected"), + [ + ("1024x1024", "square_hd"), + ("1024x576", "landscape_16_9"), + ("640x480", {"width": 640, "height": 480}), + ("wide", "landscape_4_3"), + ("axb", "landscape_4_3"), + ], +) +def test_map_openai_params_translates_size_to_image_size(size: str, expected: object): + optional_params = FalAIStableDiffusionConfig().map_openai_params( + non_default_params={"size": size, "n": 2, "response_format": "b64_json"}, + optional_params={}, + model="fal-ai/stable-diffusion-v35-medium", + drop_params=False, + ) + + assert optional_params == {"image_size": expected, "num_images": 2, "output_format": "jpeg"} diff --git a/tests/unit/llms/fireworks_ai/rerank/test_fireworks_ai_rerank_transformation.py b/tests/unit/llms/fireworks_ai/rerank/test_fireworks_ai_rerank_transformation.py index a03b7708238..a3d1879f97c 100644 --- a/tests/unit/llms/fireworks_ai/rerank/test_fireworks_ai_rerank_transformation.py +++ b/tests/unit/llms/fireworks_ai/rerank/test_fireworks_ai_rerank_transformation.py @@ -8,6 +8,7 @@ from unittest.mock import MagicMock import httpx import pytest +from pydantic import ValidationError from litellm.llms.fireworks_ai.rerank.transformation import FireworksAIRerankConfig from litellm.types.rerank import RerankResponse @@ -341,3 +342,68 @@ class TestFireworksAIRerankTransform: assert headers["Authorization"] == "Bearer test-api-key" assert headers["Content-Type"] == "application/json" + + +def _transform(payload: object) -> RerankResponse: + return FireworksAIRerankConfig().transform_rerank_response( + model="fireworks_ai/fireworks/qwen3-reranker-8b", + raw_response=httpx.Response(200, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + ) + + +@pytest.mark.parametrize( + ("result", "expected"), + [ + ({"index": "1", "relevance_score": "0.5"}, {"index": 1, "relevance_score": 0.5}), + ({"index": 2, "relevance_score": 1}, {"index": 2, "relevance_score": 1.0}), + ], +) +def test_transform_rerank_response_converts_index_and_score_with_int_and_float( + result: dict[str, object], expected: dict[str, object] +): + assert _transform({"id": "rerank-1", "data": [result]}).results == [expected] + + +@pytest.mark.parametrize("usage_fields", [{}, {"usage": {}}]) +def test_transform_rerank_response_without_usage_counters_reports_zero_tokens(usage_fields: dict[str, object]): + response = _transform({"id": "rerank-1", "data": [{"index": 0, "relevance_score": 0.5}], **usage_fields}) + + assert response.meta == {"billed_units": {"search_units": 0}, "tokens": {"input_tokens": 0, "output_tokens": 0}} + + +def test_transform_rerank_response_keeps_provider_id_and_falls_back_to_results_key(): + response = _transform({"id": "rerank-1", "data": [], "results": [{"index": 3, "relevance_score": 0.25}]}) + + assert response.id == "rerank-1" + assert response.results == [{"index": 3, "relevance_score": 0.25}] + + +def test_transform_rerank_response_empty_results_list_yields_no_results(): + assert _transform({"id": "rerank-1", "results": []}).results == [] + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"data": [{"index": 0, "relevance_score": 0.5}], "usage": None}, + {"data": [{"index": 0, "relevance_score": 0.5}], "usage": {"total_tokens": 1.5}}, + {"data": 7}, + {"data": ["not an object"]}, + {"data": [{"index": None, "relevance_score": 0.5}]}, + {"data": [{"index": 0, "relevance_score": [0.5]}]}, + {"id": 7, "data": [{"index": 0, "relevance_score": 0.5}]}, + ], +) +def test_transform_rerank_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError): + _transform(payload) + + +def test_transform_rerank_response_shape_errors_do_not_echo_the_payload(): + with pytest.raises(ValidationError) as exc_info: + _transform({"data": [{"index": {"leaked": "document text"}, "relevance_score": 0.5}]}) + + assert "document text" not in str(exc_info.value) diff --git a/tests/unit/llms/gemini/test_gemini_image_generation_transformation.py b/tests/unit/llms/gemini/test_gemini_image_generation_transformation.py index d9509856759..2a69926414a 100644 --- a/tests/unit/llms/gemini/test_gemini_image_generation_transformation.py +++ b/tests/unit/llms/gemini/test_gemini_image_generation_transformation.py @@ -1,4 +1,8 @@ +from unittest.mock import Mock + import httpx +import pytest +from pydantic import ValidationError from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.llms.gemini.image_generation.transformation import GoogleImageGenConfig @@ -419,3 +423,47 @@ def test_gemini_image_generation_response_without_grounding_has_no_web_search_re ) assert getattr(result.usage, "web_search_requests", None) is None + + +def _transform_response(model: str, payload: object) -> ImageResponse: + return GoogleImageGenConfig().transform_image_generation_response( + model=model, + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(data=[]), + logging_obj=Mock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize("predictions", [[], "", {}]) +def test_imagen_generation_response_with_empty_predictions_has_no_images(predictions: object): + assert _transform_response("gemini/imagen-4.0-generate-001", {"predictions": predictions}).data == [] + + +def test_imagen_generation_response_keeps_predictions_without_image_bytes(): + result = _transform_response( + "gemini/imagen-4.0-generate-001", + {"predictions": [{"bytesBase64Encoded": "first-image"}, {"mimeType": "image/png"}]}, + ) + + assert [image.b64_json for image in result.data or []] == ["first-image", None] + + +@pytest.mark.parametrize( + ("model", "payload"), + [ + ("gemini/imagen-4.0-generate-001", {"predictions": None}), + ("gemini/imagen-4.0-generate-001", {"predictions": "not a list"}), + ("gemini/imagen-4.0-generate-001", {"predictions": [{"bytesBase64Encoded": "a"}, "not an object"]}), + ("gemini-3.1-flash-image-preview", {"usageMetadata": "not an object"}), + ("gemini-3.1-flash-image-preview", {"usageMetadata": ["not an object"]}), + ], +) +def test_image_generation_response_rejects_malformed_payloads_without_echoing_them(model: str, payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform_response(model, payload) + + assert "input_value" not in str(exc_info.value) diff --git a/tests/unit/llms/gigachat/embedding/test_gigachat_embedding_transformation.py b/tests/unit/llms/gigachat/embedding/test_gigachat_embedding_transformation.py index 01fe66ca4c7..e9c65cbb591 100644 --- a/tests/unit/llms/gigachat/embedding/test_gigachat_embedding_transformation.py +++ b/tests/unit/llms/gigachat/embedding/test_gigachat_embedding_transformation.py @@ -12,6 +12,7 @@ from unittest.mock import MagicMock, patch import httpx import pytest +from pydantic import ValidationError from litellm import LlmProviders from litellm.llms.gigachat.embedding.transformation import ( @@ -339,4 +340,64 @@ class TestGetErrorClass: ) assert isinstance(error, GigaChatEmbeddingError) assert error.status_code == 400 - assert error.message == "embedding failed" \ No newline at end of file + assert error.message == "embedding failed" + + +def _transform(raw_response: httpx.Response) -> EmbeddingResponse: + return GigaChatEmbeddingConfig().transform_embedding_response( + model="Embeddings", + raw_response=raw_response, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="test-key", + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +def test_transform_embedding_response_moves_per_item_usage_into_the_response_usage(): + response = _transform( + httpx.Response( + 200, + json={ + "object": "list", + "model": "Embeddings", + "data": [ + {"object": "embedding", "index": 0, "embedding": [0.1, 2], "usage": {"prompt_tokens": 3}}, + {"object": "embedding", "index": 1, "embedding": [0.5], "usage": {"prompt_tokens": 4}}, + ], + "unknown": "ignored", + }, + ) + ) + + assert response.model_dump() == { + "model": "Embeddings", + "data": [ + {"object": "embedding", "index": 0, "embedding": [0.1, 2]}, + {"object": "embedding", "index": 1, "embedding": [0.5]}, + ], + "object": "list", + "usage": { + "completion_tokens": 0, + "prompt_tokens": 7, + "total_tokens": 7, + "completion_tokens_details": None, + "prompt_tokens_details": None, + }, + } + + +@pytest.mark.parametrize("body", [b"null", b"[]", b"7"]) +def test_transform_embedding_response_body_that_is_not_an_object_raises_type_error(body: bytes): + with pytest.raises(TypeError): + _transform(httpx.Response(200, content=body)) + + +def test_transform_embedding_response_invalid_envelope_field_is_reported_by_name(): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, json={"data": [], "model": 5})) + + assert exc_info.value.title == "EmbeddingResponse" + assert [error["loc"] for error in exc_info.value.errors()] == [("model",)] diff --git a/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py b/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py index f43e2e4d1cb..83c0cb2fec2 100644 --- a/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py +++ b/tests/unit/llms/github_copilot/embedding/test_github_copilot_embedding_transformation.py @@ -1,6 +1,8 @@ from unittest.mock import MagicMock, patch +import httpx import pytest +from pydantic import ValidationError from litellm.exceptions import AuthenticationError @@ -8,6 +10,7 @@ from litellm.llms.github_copilot.embedding.transformation import ( GithubCopilotEmbeddingConfig, ) from litellm.llms.github_copilot.common_utils import GetAPIKeyError +from litellm.types.utils import EmbeddingResponse def test_github_copilot_embedding_config_validate_environment(): @@ -201,3 +204,57 @@ def test_github_copilot_embedding_config_transform_response(): assert len(response.data) == 1 assert response.data[0]["embedding"] == [0.1, 0.2, 0.3] assert response.model == "text-embedding-3-small" + + +def _transform(raw_response: httpx.Response) -> EmbeddingResponse: + return GithubCopilotEmbeddingConfig().transform_embedding_response( + model="github_copilot/text-embedding-3-small", + raw_response=raw_response, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="test-key", + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +def test_transform_embedding_response_keeps_the_openai_envelope(): + response = _transform( + httpx.Response( + 200, + json={ + "object": "list", + "data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 5, "total_tokens": 5}, + "unknown": "ignored", + }, + ) + ) + + assert response.model_dump() == { + "model": "text-embedding-3-small", + "data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}], + "object": "list", + "usage": { + "completion_tokens": 0, + "prompt_tokens": 5, + "total_tokens": 5, + "completion_tokens_details": None, + "prompt_tokens_details": None, + }, + } + + +@pytest.mark.parametrize("body", [b"null", b"[]", b"7", b'"leaked payload text"', b'[{"data": "leaked payload text"}]']) +def test_transform_embedding_response_rejects_a_body_that_is_not_an_object(body: bytes): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, content=body)) + + assert "leaked payload text" not in str(exc_info.value) + + +def test_transform_embedding_response_object_without_data_is_an_invalid_response_object(): + with pytest.raises(Exception, match="Invalid response object"): + _transform(httpx.Response(200, json={"model": "text-embedding-3-small"})) diff --git a/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py b/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py index 38355f32da1..44e8145ff33 100644 --- a/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py +++ b/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py @@ -1,9 +1,12 @@ from unittest.mock import MagicMock +import httpx import pytest +from pydantic import ValidationError from litellm.llms.jina_ai.embedding.transformation import JinaAIEmbeddingConfig +from litellm.types.utils import EmbeddingResponse JINA_KEY_ENV_NAMES = ("JINA_AI_API_KEY", "JINA_API_KEY", "JINA_AI_TOKEN") @@ -131,3 +134,61 @@ class TestJinaAIEmbeddingTransform: "input": expected_input, } assert result == expected_result + + +def _transform(raw_response: httpx.Response) -> EmbeddingResponse: + return JinaAIEmbeddingConfig().transform_embedding_response( + model="jina-embeddings-v3", + raw_response=raw_response, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="test-key", + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +def test_transform_embedding_response_builds_the_response_from_the_body(): + response = _transform( + httpx.Response( + 200, + json={ + "model": "jina-embeddings-v3", + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 2]}], + "usage": {"prompt_tokens": 3, "total_tokens": 5}, + "unknown": "ignored", + }, + ) + ) + + assert response.model_dump() == { + "model": "jina-embeddings-v3", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 2]}], + "object": "list", + "usage": { + "completion_tokens": 0, + "prompt_tokens": 3, + "total_tokens": 5, + "completion_tokens_details": None, + "prompt_tokens_details": None, + }, + } + + +@pytest.mark.parametrize("body", [b"null", b"7", b'["leaked payload text"]', b'"leaked payload text"']) +def test_transform_embedding_response_rejects_a_body_that_is_not_an_object(body: bytes): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, content=body)) + + assert "leaked payload text" not in str(exc_info.value) + + +@pytest.mark.parametrize("field", ["model", "data", "usage"]) +def test_transform_embedding_response_invalid_field_is_reported_by_name(field: str): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, json={field: 5})) + + assert exc_info.value.title == "EmbeddingResponse" + assert [error["loc"] for error in exc_info.value.errors()] == [(field,)] diff --git a/tests/unit/llms/jina_ai/rerank/__init__.py b/tests/unit/llms/jina_ai/rerank/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/jina_ai/rerank/test_transformation.py b/tests/unit/llms/jina_ai/rerank/test_transformation.py new file mode 100644 index 00000000000..9ec02a4548b --- /dev/null +++ b/tests/unit/llms/jina_ai/rerank/test_transformation.py @@ -0,0 +1,111 @@ +from unittest.mock import MagicMock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.jina_ai.rerank.transformation import JinaAIRerankConfig +from litellm.types.rerank import RerankResponse + + +def _transform(payload: object, status_code: int = 200) -> RerankResponse: + return JinaAIRerankConfig().transform_rerank_response( + model="jina-reranker-v2-base-multilingual", + raw_response=httpx.Response(status_code, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + ) + + +def test_transform_rerank_response_maps_results_and_usage(): + response = _transform( + { + "id": "rerank-1", + "results": [ + {"index": 1, "relevance_score": 0.72, "document": "hello"}, + {"index": 0, "relevance_score": 0.25, "document": {"text": "world", "extra": 1}}, + {"index": 2, "relevance_score": 0, "extra": True}, + ], + "usage": {"total_tokens": 21, "prompt_tokens": 21}, + } + ) + + assert response.id == "rerank-1" + assert response.results == [ + {"index": 1, "relevance_score": 0.72, "document": {"text": "hello"}}, + {"index": 0, "relevance_score": 0.25, "document": {"text": "world"}}, + {"index": 2, "relevance_score": 0.0}, + ] + assert response.meta == {"billed_units": {"total_tokens": 21}, "tokens": {}} + + +@pytest.mark.parametrize( + ("payload", "expected_meta"), + [ + ({"results": []}, {"billed_units": {}, "tokens": {}}), + ({"results": [], "usage": {}}, {"billed_units": {}, "tokens": {}}), + ( + {"results": [], "usage": {"total_tokens": 12, "input_tokens": 7, "output_tokens": 3, "unknown": 1}}, + {"billed_units": {"total_tokens": 12}, "tokens": {"input_tokens": 7, "output_tokens": 3}}, + ), + ], +) +def test_transform_rerank_response_keeps_only_known_usage_counters( + payload: dict[str, object], expected_meta: dict[str, object] +): + assert _transform(payload).meta == expected_meta + + +def test_transform_rerank_response_generates_an_id_when_the_provider_sends_none(): + response = _transform({"id": None, "results": [{"index": 0, "relevance_score": 0.5}]}) + + assert isinstance(response.id, str) + assert response.id != "" + + +def test_transform_rerank_response_empty_results_list_yields_no_results(): + assert _transform({"id": "rerank-1", "results": []}).results == [] + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"results": [], "usage": None}, + {"results": [], "usage": {"total_tokens": 1.5}}, + {"results": [], "usage": {"input_tokens": "many"}}, + {"results": 7}, + {"results": ["not an object"]}, + {"results": [{"index": 0, "relevance_score": 0.5}], "id": 7}, + {"results": [{"index": None, "relevance_score": 0.5}]}, + ], +) +def test_transform_rerank_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError): + _transform(payload) + + +def test_transform_rerank_response_without_results_raises_value_error_naming_the_body(): + with pytest.raises(ValueError, match="No results found") as exc_info: + _transform({"id": "rerank-1"}) + + assert str(exc_info.value) == "No results found in the response={'id': 'rerank-1'}" + + +def test_transform_rerank_response_result_without_score_raises_key_error(): + with pytest.raises(KeyError, match="relevance_score"): + _transform({"results": [{"index": 0}]}) + + +def test_transform_rerank_response_non_200_raises_with_the_response_text(): + with pytest.raises(Exception, match="quota exceeded") as exc_info: + _transform({"detail": "quota exceeded"}, status_code=429) + + assert type(exc_info.value) is Exception + + +def test_transform_rerank_response_shape_errors_do_not_echo_the_payload(): + with pytest.raises(ValidationError) as exc_info: + _transform({"results": [["leaked document text"]]}) + + assert "leaked document text" not in str(exc_info.value) diff --git a/tests/unit/llms/manus/files/__init__.py b/tests/unit/llms/manus/files/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/manus/files/test_transformation.py b/tests/unit/llms/manus/files/test_transformation.py new file mode 100644 index 00000000000..a47e1c8cb9b --- /dev/null +++ b/tests/unit/llms/manus/files/test_transformation.py @@ -0,0 +1,128 @@ +import time +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.manus.files.transformation import ManusFilesConfig +from litellm.types.llms.openai import OpenAIFileObject + + +def _list_files(body: object) -> list[OpenAIFileObject]: + return ManusFilesConfig().transform_list_files_response( + raw_response=httpx.Response(200, json=body), logging_obj=Mock(), litellm_params={} + ) + + +def test_delete_file_response_is_read_into_a_file_deleted_object(): + deleted = ManusFilesConfig().transform_delete_file_response( + raw_response=httpx.Response(200, json={"id": "file-1", "deleted": True, "object": "file", "region": "eu"}), + logging_obj=Mock(), + litellm_params={}, + ) + + assert deleted.model_dump() == {"id": "file-1", "deleted": True, "object": "file", "region": "eu"} + + +@pytest.mark.parametrize("body", [b'["secret-file"]', b'"secret-file"', b"7", b"null"]) +def test_delete_file_response_rejects_a_body_that_is_not_an_object_without_echoing_it(body: bytes): + with pytest.raises(ValidationError) as exc_info: + ManusFilesConfig().transform_delete_file_response( + raw_response=httpx.Response(200, content=body), logging_obj=Mock(), litellm_params={} + ) + + assert "secret-file" not in str(exc_info.value) + + +def test_delete_file_response_requires_the_file_deleted_fields(): + with pytest.raises(ValidationError, match="FileDeleted"): + ManusFilesConfig().transform_delete_file_response( + raw_response=httpx.Response(200, json={"id": "file-1"}), logging_obj=Mock(), litellm_params={} + ) + + +def test_list_files_response_maps_every_listed_file(): + files = _list_files( + { + "object": "list", + "data": [ + { + "id": "file-1", + "bytes": 12, + "filename": "a.pdf", + "purpose": "batch", + "status": "processed", + "status_details": "done", + "created_at": "2024-01-02T03:04:05Z", + }, + {"id": "file-2", "created_at": "2024-01-02T03:04:05.123456+00:00"}, + ], + } + ) + + created_at = int(time.mktime(time.strptime("2024-01-02T03:04:05", "%Y-%m-%dT%H:%M:%S"))) + assert [file.model_dump() for file in files] == [ + { + "id": "file-1", + "bytes": 12, + "created_at": created_at, + "filename": "a.pdf", + "object": "file", + "purpose": "batch", + "status": "processed", + "expires_at": None, + "status_details": "done", + }, + { + "id": "file-2", + "bytes": 0, + "created_at": created_at, + "filename": "", + "object": "file", + "purpose": "assistants", + "status": "uploaded", + "expires_at": None, + "status_details": None, + }, + ] + + +@pytest.mark.parametrize("created_at", ["not-a-date", "", None, 0]) +def test_list_files_response_falls_back_to_now_for_an_unusable_created_at(created_at: object): + before = int(time.time()) + + (file,) = _list_files({"data": [{"id": "file-1", "created_at": created_at}]}) + + assert before <= file.created_at <= int(time.time()) + + +@pytest.mark.parametrize("body", [{}, {"data": []}]) +def test_list_files_response_is_empty_without_listed_files(body: dict[str, object]): + assert _list_files(body) == [] + + +@pytest.mark.parametrize( + "body", + [ + ["secret-file"], + "secret-file", + {"data": None}, + {"data": 7}, + {"data": "secret-file"}, + {"data": ["secret-file"]}, + {"data": [{"id": "file-1"}, ["secret-file"]]}, + {"data": [{"id": "file-1", "created_at": ["secret-file"]}]}, + {"data": [{"id": "file-1", "created_at": 1700000000}]}, + ], +) +def test_list_files_response_rejects_malformed_listings_without_echoing_them(body: object): + with pytest.raises(ValidationError) as exc_info: + _list_files(body) + + assert "secret-file" not in str(exc_info.value) + + +def test_list_files_response_rejects_a_file_the_file_object_cannot_hold(): + with pytest.raises(ValidationError, match="OpenAIFileObject"): + _list_files({"data": [{"id": "file-1", "purpose": "not-a-purpose"}]}) diff --git a/tests/unit/llms/minimax/text_to_speech/__init__.py b/tests/unit/llms/minimax/text_to_speech/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/minimax/text_to_speech/test_transformation.py b/tests/unit/llms/minimax/text_to_speech/test_transformation.py new file mode 100644 index 00000000000..a183c314a61 --- /dev/null +++ b/tests/unit/llms/minimax/text_to_speech/test_transformation.py @@ -0,0 +1,122 @@ +import base64 +from typing import Final +from unittest.mock import Mock + +import httpx +import pytest + +from litellm.llms.minimax.text_to_speech.transformation import MinimaxException, MinimaxTextToSpeechConfig +from litellm.types.llms.openai import HttpxBinaryResponseContent + +_AUDIO: Final = b"ID3\x04minimax-audio" +_REQUEST: Final = httpx.Request("POST", "https://api.minimax.io/v1/t2a_v2") + + +def _transform(payload: object, status_code: int = 200) -> HttpxBinaryResponseContent: + return MinimaxTextToSpeechConfig().transform_text_to_speech_response( + model="speech-02-hd", + raw_response=httpx.Response( + status_code, + json=payload, + request=_REQUEST, + headers={"content-encoding": "identity", "x-trace": "abc"}, + ), + logging_obj=Mock(), + ) + + +@pytest.mark.parametrize( + "payload", + [ + {"data": {"audio": _AUDIO.hex()}, "status": 0, "extra_info": {"audio_length": 5}}, + {"data": {"audio": _AUDIO.hex(), "audio_url": ""}}, + {"data": {"audio": ""}, "audio_file": _AUDIO.hex()}, + {"data": {"audio": None}, "audio_file": base64.b64encode(_AUDIO).decode()}, + {"base_resp": {"status_code": 0}, "audio_file": base64.b64encode(_AUDIO).decode()}, + ], +) +def test_transform_text_to_speech_response_decodes_hex_or_base64_audio(payload: dict[str, object]): + response = _transform(payload).response + + assert response.status_code == 200 + assert response.content == _AUDIO + assert response.headers["content-length"] == str(len(_AUDIO)) + assert response.headers["x-trace"] == "abc" + assert "content-encoding" not in response.headers + + +@pytest.mark.parametrize( + ("payload", "detail"), + [ + ({"status": 2, "ced": "invalid api key"}, "invalid api key"), + ({"status": 2}, "Unknown error"), + ({"status": 1004, "ced": ""}, "API returned status 1004"), + ({"status": "failed", "ced": None, "data": {"audio": _AUDIO.hex()}}, "API returned status failed"), + ], +) +def test_transform_text_to_speech_response_reports_api_status_errors(payload: dict[str, object], detail: str): + with pytest.raises(MinimaxException) as exc_info: + _transform(payload, status_code=401) + + assert exc_info.value.message == f"MiniMax TTS error: {detail}" + assert exc_info.value.status_code == 401 + + +def test_transform_text_to_speech_response_refuses_url_output(): + with pytest.raises(MinimaxException) as exc_info: + _transform({"data": {"audio_url": "https://cdn.example/a.mp3", "audio": _AUDIO.hex()}}) + + assert exc_info.value.message == ( + "URL output format is not yet supported. Use 'hex' format or fetch from URL: https://cdn.example/a.mp3" + ) + assert exc_info.value.status_code == 500 + + +@pytest.mark.parametrize( + ("payload", "keys"), + [ + ({}, []), + ({"data": {}, "status": 0}, ["data", "status"]), + ({"data": {"audio": ""}, "audio_file": ""}, ["data", "audio_file"]), + ({"data": {"audio": None}, "audio_file": None}, ["data", "audio_file"]), + ({"audio_file": 0, "data": {"audio": []}}, ["audio_file", "data"]), + ], +) +def test_transform_text_to_speech_response_without_audio_lists_the_response_keys( + payload: dict[str, object], keys: list[str] +): + with pytest.raises(MinimaxException) as exc_info: + _transform(payload) + + assert exc_info.value.message == f"No audio data in MiniMax response. Response keys: {keys}" + assert exc_info.value.status_code == 500 + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + "not an object", + {"data": None}, + {"data": "not an object"}, + {"data": ["not", "an", "object"]}, + {"data": {"audio": 7}}, + {"data": {"audio": ["49", "44"]}}, + {"audio_file": {"hex": "4944"}}, + ], +) +def test_transform_text_to_speech_response_wraps_malformed_payloads_without_echoing_them(payload: object): + with pytest.raises(MinimaxException) as exc_info: + _transform(payload) + + assert exc_info.value.message.startswith("Error processing MiniMax response: ") + assert "input_value" not in exc_info.value.message + assert exc_info.value.status_code == 500 + + +def test_transform_text_to_speech_response_reports_undecodable_audio(): + with pytest.raises(MinimaxException) as exc_info: + _transform({"data": {"audio": "zzz"}}) + + assert exc_info.value.message.startswith("Failed to decode audio data: ") + assert exc_info.value.status_code == 500 diff --git a/tests/unit/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py b/tests/unit/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py index 2ffe7c3e686..5b0ca09ed13 100644 --- a/tests/unit/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py +++ b/tests/unit/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py @@ -7,9 +7,12 @@ transformation between OpenAI-compatible format and ModelScope API format. from unittest.mock import MagicMock, patch +import httpx import pytest +from pydantic import ValidationError +import litellm from litellm.llms.modelscope.image_generation.transformation import ( ModelScopeImageGenerationConfig, ) @@ -449,3 +452,83 @@ class TestModelScopeImageGenerationTransformation: ) assert isinstance(error, BadRequestError) + + +def _transform_generation_response(payload: object, status_code: int = 200) -> ImageResponse: + return ModelScopeImageGenerationConfig().transform_image_generation_response( + model="Qwen/Qwen-Image", + raw_response=httpx.Response(status_code, json=payload), + model_response=ImageResponse(data=[]), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize( + ("item", "expected"), + [ + ({"url": "https://a.example/i.png"}, ("https://a.example/i.png", None, None)), + ({"b64_json": "aGVsbG8="}, (None, "aGVsbG8=", None)), + ( + {"url": "https://a.example/i.png", "revised_prompt": "a calmer cat", "seed": 7}, + ("https://a.example/i.png", None, "a calmer cat"), + ), + ({}, (None, None, None)), + ], +) +def test_transform_image_generation_response_maps_one_image( + item: dict[str, object], expected: tuple[str | None, str | None, str | None] +) -> None: + response = _transform_generation_response({"created": 1, "data": [item]}) + + assert [(image.url, image.b64_json, image.revised_prompt) for image in response.data] == [expected] + + +@pytest.mark.parametrize("payload", [{}, {"data": []}, {"data": ""}, {"data": {}}, {"created": 1}]) +def test_transform_image_generation_response_without_images_is_empty(payload: dict[str, object]) -> None: + assert _transform_generation_response(payload).data == [] + + +@pytest.mark.parametrize( + ("error", "status_code", "expected_class", "expected_message"), + [ + ({"message": "Invalid prompt"}, 400, litellm.BadRequestError, "ModelScope error: Invalid prompt"), + ({"message": "Bad key"}, 401, litellm.AuthenticationError, "ModelScope error: Bad key"), + ({"code": "overloaded"}, 503, litellm.InternalServerError, "ModelScope error: {'code': 'overloaded'}"), + ({}, 200, litellm.BadRequestError, "ModelScope error: {}"), + ], +) +def test_transform_image_generation_response_reports_api_error_bodies( + error: dict[str, object], status_code: int, expected_class: type[Exception], expected_message: str +) -> None: + with pytest.raises(expected_class) as exc_info: + _transform_generation_response({"error": error, "data": [{"url": "ignored"}]}, status_code) + + assert str(exc_info.value).endswith(expected_message) + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + "error", + 7, + {"error": "plain text error"}, + {"error": None}, + {"error": ["not", "an", "object"]}, + {"data": 7}, + {"data": None}, + {"data": ["not an object"]}, + {"data": [{"url": "https://a.example/i.png"}, None]}, + ], +) +def test_transform_image_generation_response_rejects_malformed_payloads_without_echoing_them( + payload: object, +) -> None: + with pytest.raises(ValidationError) as exc_info: + _transform_generation_response(payload) + + assert "input_value" not in str(exc_info.value) diff --git a/tests/unit/llms/nvidia_nim/rerank/test_nvidia_nim_rerank_transformation.py b/tests/unit/llms/nvidia_nim/rerank/test_nvidia_nim_rerank_transformation.py index 2b03b2d807b..39ba396d8d2 100644 --- a/tests/unit/llms/nvidia_nim/rerank/test_nvidia_nim_rerank_transformation.py +++ b/tests/unit/llms/nvidia_nim/rerank/test_nvidia_nim_rerank_transformation.py @@ -10,7 +10,9 @@ truncate. Two defects are covered here: import json from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest +from pydantic import ValidationError import litellm from litellm.llms.nvidia_nim.rerank.ranking_transformation import ( @@ -239,3 +241,87 @@ class TestNvidiaNimRetrievalRerankRequestTransform: doc = {"title": "no supported fields here"} request_data = self._build_request([doc]) assert request_data["passages"] == [{"text": json.dumps(doc)}] + + +def _transform_retrieval_response(payload: object) -> RerankResponse: + return NvidiaNimRerankConfig().transform_rerank_response( + model="nvidia/llama-3_2-nv-rerankqa-1b-v2", + raw_response=httpx.Response(200, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + request_data={"passages": [{"text": "first"}, {"text": "second"}]}, + ) + + +@pytest.mark.parametrize( + ("usage", "expected_total_tokens"), + [ + ({"total_tokens": 42}, 42), + ({"total_tokens": 0}, 2), + ({}, 2), + ], +) +def test_transform_rerank_response_bills_reported_tokens_or_falls_back_to_result_count( + usage: dict[str, object], expected_total_tokens: int +): + response = _transform_retrieval_response( + {"rankings": [{"index": 1, "logit": 2.5}, {"index": 0, "logit": -1}], "usage": usage} + ) + + assert response.meta == {"billed_units": {"total_tokens": expected_total_tokens}} + assert response.results == [ + {"index": 1, "relevance_score": 2.5, "document": {"text": "second"}}, + {"index": 0, "relevance_score": -1.0, "document": {"text": "first"}}, + ] + + +def test_transform_rerank_response_keeps_provider_id_and_bills_one_token_per_result_without_usage(): + response = _transform_retrieval_response({"id": "rank-1", "rankings": [{"index": 0, "logit": 0.5}]}) + + assert response.id == "rank-1" + assert response.meta == {"billed_units": {"total_tokens": 1}} + + +@pytest.mark.parametrize( + "payload", + [ + {"rankings": [], "usage": None}, + {"rankings": [], "usage": {"total_tokens": None}}, + {"rankings": [], "usage": {"total_tokens": "42"}}, + {"rankings": [], "usage": {"total_tokens": 1.5}}, + {"rankings": [], "id": 7}, + ], +) +def test_transform_rerank_response_rejects_malformed_usage_and_id(payload: dict[str, object]): + with pytest.raises(ValidationError): + _transform_retrieval_response(payload) + + +def test_transform_rerank_response_non_object_body_raises_attribute_error(): + with pytest.raises(AttributeError): + _transform_retrieval_response(["not", "an", "object"]) + + +def test_transform_rerank_response_usage_shape_errors_do_not_echo_the_payload(): + with pytest.raises(ValidationError) as exc_info: + _transform_retrieval_response({"rankings": [], "usage": {"total_tokens": "leaked usage text"}}) + + assert "leaked usage text" not in str(exc_info.value) + + +def test_map_cohere_rerank_params_passes_provider_params_through_and_maps_top_n(): + params = NvidiaNimRerankConfig().map_cohere_rerank_params( + non_default_params={"truncate": "END"}, + model="nvidia/llama-3_2-nv-rerankqa-1b-v2", + drop_params=False, + query="which passage shows a cat?", + documents=["a", {"text": "b"}], + top_n=1, + ) + + assert params == { + "query": "which passage shows a cat?", + "documents": ["a", {"text": "b"}], + "top_k": 1, + "truncate": "END", + } diff --git a/tests/unit/llms/oci/embed/test_oci_embed_transformation.py b/tests/unit/llms/oci/embed/test_oci_embed_transformation.py index 4ffd79ff147..01be7c7b904 100644 --- a/tests/unit/llms/oci/embed/test_oci_embed_transformation.py +++ b/tests/unit/llms/oci/embed/test_oci_embed_transformation.py @@ -379,3 +379,78 @@ class TestOCIEmbedConfig: litellm_params={}, ) assert "eu-frankfurt-1" in url + + +def _transform(raw_response: httpx.Response) -> EmbeddingResponse: + return OCIEmbedConfig().transform_embedding_response( + model="cohere.embed-v3.0", + raw_response=raw_response, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key=None, + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +@pytest.mark.parametrize( + ("token_fields", "expected_prompt_tokens", "expected_total_tokens"), + [ + ({"inputTextTokenCounts": [5, 6]}, 11, 11), + ({"usage": {"promptTokens": 3, "totalTokens": 4}}, 3, 4), + ({"inputTextTokenCounts": [5, 6], "usage": {"promptTokens": 3, "totalTokens": 4}}, 11, 11), + ({}, 0, 0), + ], +) +def test_transform_embedding_response_reads_usage_from_whichever_token_field_is_present( + token_fields: dict[str, object], expected_prompt_tokens: int, expected_total_tokens: int +): + response = _transform( + httpx.Response( + 200, + json={ + "embeddings": [[0.1, 1], [0.5]], + "modelId": "cohere.embed-v3.0", + "modelVersion": "3.0", + "unknown": "ignored", + **token_fields, + }, + ) + ) + + assert response.model_dump() == { + "model": "cohere.embed-v3.0", + "data": [ + {"object": "embedding", "index": 0, "embedding": [0.1, 1.0]}, + {"object": "embedding", "index": 1, "embedding": [0.5]}, + ], + "object": "list", + "usage": { + "completion_tokens": 0, + "prompt_tokens": expected_prompt_tokens, + "total_tokens": expected_total_tokens, + "completion_tokens_details": None, + "prompt_tokens_details": None, + }, + } + + +@pytest.mark.parametrize("body", [b"null", b"7", b'["leaked payload text"]', b'"leaked payload text"']) +def test_transform_embedding_response_body_that_is_not_an_object_is_a_schema_error(body: bytes): + with pytest.raises(OCIError) as exc_info: + _transform(httpx.Response(200, content=body)) + + assert exc_info.value.status_code == 500 + assert exc_info.value.message.startswith("OCI embed response does not match expected schema: ") + assert "leaked payload text" not in exc_info.value.message + + +def test_transform_embedding_response_object_missing_required_fields_names_them(): + with pytest.raises(OCIError) as exc_info: + _transform(httpx.Response(200, json={"embeddings": [[0.1]]})) + + assert exc_info.value.status_code == 500 + assert exc_info.value.message.startswith("OCI embed response does not match expected schema: ") + assert "modelId" in exc_info.value.message + assert "modelVersion" in exc_info.value.message diff --git a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py index 5c85faa5e13..5d71cbe9afc 100644 --- a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py +++ b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py @@ -2336,3 +2336,32 @@ class TestStreamingScanKey: handler = OpenAIChatCompletionsHandler() key = handler.get_streaming_scan_key([self._chunk("hi"), b"data: [DONE]"]) assert key.texts == ("hi",) + + def test_released_stream_as_ended_finishes_only_the_choice_whose_tool_call_was_in_flight(self): + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + released = ( + self._chunk("hi", index=0), + ModelResponseStream(choices=[StreamingChoices(index=1, delta=Delta(tool_calls=[tool_call]))]), + ) + handler = OpenAIChatCompletionsHandler() + ended = handler.released_stream_as_ended(released) + assert all(a is b for a, b in zip(ended[:-1], released, strict=True)) + assert [(choice.index, choice.finish_reason) for choice in ended[-1].choices] == [(1, "tool_calls")] + ended_key = handler.get_streaming_scan_key(ended) + assert ended_key.stream_ended is True and len(ended_key.tool_calls) == 1, ended_key + + def test_released_stream_as_ended_leaves_a_stream_with_no_tool_call_in_flight_as_released(self): + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + text_only = (self._chunk("hi", index=0), self._chunk(" there", index=1)) + finished_tool_call = ( + ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[tool_call]))]), + self._chunk(None, finish_reason="tool_calls", index=0), + ) + handler = OpenAIChatCompletionsHandler() + for released in (text_only, finished_tool_call): + ended = handler.released_stream_as_ended(released) + assert len(ended) == len(released) and all(a is b for a, b in zip(ended, released, strict=True)) diff --git a/tests/unit/llms/openai/responses/count_tokens/__init__.py b/tests/unit/llms/openai/responses/count_tokens/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/openai/responses/count_tokens/test_handler.py b/tests/unit/llms/openai/responses/count_tokens/test_handler.py new file mode 100644 index 00000000000..2427b8d272f --- /dev/null +++ b/tests/unit/llms/openai/responses/count_tokens/test_handler.py @@ -0,0 +1,26 @@ +import pytest + +from litellm.llms.openai.common_utils import OpenAIError +from litellm.llms.openai.responses.count_tokens.handler import OpenAICountTokensHandler + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model", "input_value", "expected_message"), + [ + ("", "hello", "CountTokens processing error: model parameter is required"), + ("gpt-4o", "", "CountTokens processing error: input parameter is required"), + ("gpt-4o", [], "CountTokens processing error: input parameter is required"), + ], +) +async def test_request_without_model_or_input_is_rejected_before_calling_openai( + model: str, input_value: str | list[object], expected_message: str +) -> None: + with pytest.raises(OpenAIError) as rejected: + await OpenAICountTokensHandler().handle_count_tokens_request( + model=model, + input=input_value, + api_key="sk-test", + ) + + assert (rejected.value.status_code, rejected.value.message) == (500, expected_message) diff --git a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py index 87980a47f87..8bd242308f7 100644 --- a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -6,6 +6,7 @@ with guardrail transformations. """ import copy +import json from collections.abc import Callable from typing import Any, Final, List, Literal, Optional, Tuple from unittest.mock import AsyncMock, MagicMock, patch @@ -3654,3 +3655,85 @@ class TestOpenAIResponsesHandlerStreamingScanKey: ended_key = handler.get_streaming_scan_key([self._delta(0, "hi"), added, self._completed(3, [function_call])]) assert ended_key.tool_calls_in_flight is False assert len(ended_key.tool_calls) == 1 + + def test_released_stream_as_ended_keys_the_tool_call_the_client_already_received(self): + handler = OpenAIResponsesHandler() + added = { + "type": "response.output_item.added", + "sequence_number": 1, + "item": {"type": "function_call", "id": "fc_1", "call_id": "call_1", "name": "get_weather", "arguments": ""}, + } + arguments_delta = { + "type": "response.function_call_arguments.delta", + "sequence_number": 2, + "item_id": "fc_1", + "delta": '{"city": "Paris"', + } + ended_key = handler.get_streaming_scan_key( + handler.released_stream_as_ended([self._delta(0, "hi"), added, arguments_delta]) + ) + assert ended_key.stream_ended is True + assert ended_key.texts == ("hi",) + assert len(ended_key.tool_calls) == 1 and "Paris" in ended_key.tool_calls[0], ended_key + + def test_released_stream_as_ended_leaves_a_text_only_stream_as_released(self): + released = (self._delta(0, "hi"), self._delta(1, " there")) + ended = OpenAIResponsesHandler().released_stream_as_ended(released) + assert ended == released and all(a is b for a, b in zip(ended, released, strict=True)) + + @staticmethod + def _finished_function_call(sequence_number: int, item_id: str, city: str) -> tuple[dict[str, object], ...]: + arguments = json.dumps({"city": city}) + pending = {"type": "function_call", "id": item_id, "call_id": "call_" + item_id, "name": "get_weather"} + return ( + {"type": "response.output_item.added", "sequence_number": sequence_number, "item": {**pending, "arguments": ""}}, + { + "type": "response.function_call_arguments.delta", + "sequence_number": sequence_number + 1, + "item_id": item_id, + "delta": arguments, + }, + { + "type": "response.output_item.done", + "sequence_number": sequence_number + 2, + "item": {**pending, "arguments": arguments, "status": "completed"}, + }, + ) + + def test_released_stream_as_ended_keys_every_tool_call_finished_before_the_disconnect(self): + handler = OpenAIResponsesHandler() + released = ( + self._delta(0, "hi"), + *self._finished_function_call(1, "fc_1", "Paris"), + *self._finished_function_call(4, "fc_2", "Rome"), + ) + ended_key = handler.get_streaming_scan_key(handler.released_stream_as_ended(released)) + assert ended_key.stream_ended is True + assert ended_key.texts == ("hi",) + cities = tuple(city for fingerprint in ended_key.tool_calls for city in ("Paris", "Rome") if city in fingerprint) + assert cities == ("Paris", "Rome"), ended_key + + def test_released_stream_as_ended_keys_a_message_whose_item_already_finished(self): + handler = OpenAIResponsesHandler() + message_done = { + "type": "response.output_item.done", + "sequence_number": 1, + "item": {"type": "message", "id": "msg_1", "content": [{"type": "output_text", "text": "hi"}]}, + } + ended_key = handler.get_streaming_scan_key(handler.released_stream_as_ended((self._delta(0, "hi"), message_done))) + assert ended_key.stream_ended is True + assert ended_key.texts == ("hi",) + + @pytest.mark.asyncio + async def test_scan_of_a_stream_released_through_a_finished_tool_call_covers_its_text_too(self): + handler = OpenAIResponsesHandler() + guardrail = MockRecordingGuardrail(guardrail_name="test") + released = (self._delta(0, "hi"), *self._finished_function_call(1, "fc_1", "Paris")) + await handler.process_output_streaming_response( + responses_so_far=list(handler.released_stream_as_ended(released)), + guardrail_to_apply=guardrail, + request_data={}, + ) + assert [inputs.get("texts") for inputs in guardrail.seen_inputs] == [["hi"]], guardrail.seen_inputs + tool_calls = guardrail.seen_inputs[0].get("tool_calls") or [] + assert [call["function"]["arguments"] for call in tool_calls] == ['{"city": "Paris"}'], tool_calls diff --git a/tests/unit/llms/openrouter/image_generation/test_openrouter_image_gen_transformation.py b/tests/unit/llms/openrouter/image_generation/test_openrouter_image_gen_transformation.py index e45270fb5e3..edd6168c6c3 100644 --- a/tests/unit/llms/openrouter/image_generation/test_openrouter_image_gen_transformation.py +++ b/tests/unit/llms/openrouter/image_generation/test_openrouter_image_gen_transformation.py @@ -580,3 +580,86 @@ class TestOpenRouterImageGenerationTransformation: assert isinstance(error, OpenRouterException) assert "Test error" in str(error) assert error.status_code == 400 + + +def _transform_generation_response(payload: object) -> ImageResponse: + return OpenRouterImageGenerationConfig().transform_image_generation_response( + model="google/gemini-2.5-flash-image", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(data=[]), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize( + "payload", + [ + {}, + {"choices": []}, + {"choices": ""}, + {"choices": {}}, + {"choices": [{}]}, + {"choices": [{"message": {}}]}, + {"choices": [{"message": {"images": ""}}]}, + {"choices": [{"message": {"images": [{}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {}}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {"url": None}}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {"url": ""}}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {"url": 0}}]}}]}, + ], +) +def test_transform_image_generation_response_without_usable_image_url_has_no_images( + payload: dict[str, object], +): + assert _transform_generation_response(payload).data == [] + + +@pytest.mark.parametrize( + ("url", "expected"), + [ + ("data:image/png;base64,aGVsbG8=", ("aGVsbG8=", None)), + ("data:image/png;base64,a,b", ("a,b", None)), + ("data:no-comma", (None, None)), + ("https://example.com/a.png", (None, "https://example.com/a.png")), + ], +) +def test_transform_image_generation_response_maps_one_image_url( + url: str, expected: tuple[str | None, str | None] +): + response = _transform_generation_response( + {"choices": [{"message": {"images": [{"image_url": {"url": url}}]}}]} + ) + + assert [(image.b64_json, image.url) for image in response.data] == [expected] + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"choices": 7}, + {"choices": ["not an object"]}, + {"choices": [{"message": "not an object"}]}, + {"choices": [{"message": {"images": 7}}]}, + {"choices": [{"message": {"images": ["not an object"]}}]}, + {"choices": [{"message": {"images": [{"image_url": "not an object"}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {"url": ["a"]}}]}}]}, + {"choices": [{"message": {"images": [{"image_url": {"url": 7}}]}}]}, + ], +) +def test_transform_image_generation_response_wraps_malformed_payloads_without_echoing_them( + payload: object, +): + with pytest.raises(OpenRouterException) as exc_info: + _transform_generation_response(payload) + + message = str(exc_info.value) + assert message.startswith( + "Error transforming OpenRouter image generation response: " + ) + assert "input_value" not in message + assert exc_info.value.status_code == 500 diff --git a/tests/unit/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py b/tests/unit/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py index 87e54dfba9b..422658932e6 100644 --- a/tests/unit/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py +++ b/tests/unit/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py @@ -1,4 +1,9 @@ +import httpx +import pytest +from pydantic import ValidationError +from litellm.llms.ovhcloud.audio_transcription.transformation import OVHCloudAudioTranscriptionConfig +from litellm.types.utils import TranscriptionResponse @@ -56,3 +61,40 @@ class TestOVHCloudDurationFieldMigration: mock_response.json.return_value = {"text": "silence", "seconds": 0.0} result = config.transform_audio_transcription_response(mock_response) assert result._hidden_params["duration"] == 0.0 + + +def _transform(payload: object) -> TranscriptionResponse: + return OVHCloudAudioTranscriptionConfig().transform_audio_transcription_response(httpx.Response(200, json=payload)) + + +@pytest.mark.parametrize( + ("payload", "expected_text", "expected_hidden_params"), + [ + ( + {"text": "hello", "seconds": 3.5, "duration": 9}, + "hello", + {"text": "hello", "seconds": 3.5, "duration": 3.5}, + ), + ( + {"transcript": "from transcript", "seconds": None, "duration": 4}, + "from transcript", + {"transcript": "from transcript", "seconds": None, "duration": 4}, + ), + ({"language": "en"}, "", {"language": "en"}), + ], +) +def test_transform_audio_transcription_response_normalizes_text_and_duration( + payload: dict[str, object], expected_text: str, expected_hidden_params: dict[str, object] +): + response = _transform(payload) + + assert response.text == expected_text + assert response._hidden_params == expected_hidden_params + + +@pytest.mark.parametrize("payload", [7, "spoken secret", [{"text": "spoken secret"}]]) +def test_transform_audio_transcription_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "spoken secret" not in str(exc_info.value) diff --git a/tests/unit/llms/recraft/image_generation/test_recraft_image_gen_transformation.py b/tests/unit/llms/recraft/image_generation/test_recraft_image_gen_transformation.py index 2dfe33b828c..4b677675694 100644 --- a/tests/unit/llms/recraft/image_generation/test_recraft_image_gen_transformation.py +++ b/tests/unit/llms/recraft/image_generation/test_recraft_image_gen_transformation.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock, patch import httpx import pytest +from pydantic import ValidationError from litellm.llms.recraft.image_generation.transformation import ( @@ -256,3 +257,56 @@ class TestRecraftImageGenerationTransformation: ) assert "Error transforming image generation response" in str(exc_info.value) + + +def _transform(payload: object) -> ImageResponse: + return RecraftImageGenerationConfig().transform_image_generation_response( + model="recraftv3", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_maps_each_image_object(): + response = _transform({"data": [{"url": "https://img.recraft.ai/a.png"}, {"b64_json": "QUJD"}, {}]}) + + assert [(image.url, image.b64_json) for image in response.data] == [ + ("https://img.recraft.ai/a.png", None), + (None, "QUJD"), + (None, None), + ] + + +@pytest.mark.parametrize("data", [[], "", {}]) +def test_transform_image_generation_response_with_empty_data_has_no_images(data: object): + assert _transform({"data": data}).data == [] + + +@pytest.mark.parametrize( + "payload", + [ + 7, + "https://img.recraft.ai/a.png", + [{"url": "https://img.recraft.ai/a.png"}], + {"data": None}, + {"data": "https://img.recraft.ai/a.png"}, + {"data": {"url": "https://img.recraft.ai/a.png"}}, + {"data": ["https://img.recraft.ai/a.png"]}, + {"data": [{"url": "https://img.recraft.ai/a.png"}, 7]}, + ], +) +def test_transform_image_generation_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "img.recraft.ai" not in str(exc_info.value) + + +def test_transform_image_generation_response_requires_data(): + with pytest.raises(KeyError): + _transform({"created": 1}) diff --git a/tests/unit/llms/replicate/__init__.py b/tests/unit/llms/replicate/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/replicate/chat/__init__.py b/tests/unit/llms/replicate/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/replicate/chat/test_transformation.py b/tests/unit/llms/replicate/chat/test_transformation.py new file mode 100644 index 00000000000..4c2b1840664 --- /dev/null +++ b/tests/unit/llms/replicate/chat/test_transformation.py @@ -0,0 +1,54 @@ +from unittest.mock import Mock + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.llms.replicate.chat.transformation import ReplicateConfig +from litellm.types.utils import ModelResponse + + +def _transform(raw_response: httpx.Response) -> ModelResponse: + return ReplicateConfig().transform_response( + model="acme/echo-model", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=Mock(), + request_data={"input": {"prompt": "Hello"}}, + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize( + ("output", "content"), + [ + (["Hello", ", ", "world"], "Hello, world"), + ("Hello", "Hello"), + ([], " "), + ("", " "), + ([""], " "), + ], +) +def test_transform_response_joins_the_prediction_output_into_the_message_content(output: object, content: str): + response = _transform(httpx.Response(200, json={"status": "succeeded", "output": output})) + + assert response.choices[0].message.content == content + assert response.model == "replicate/acme/echo-model" + assert response.usage.total_tokens == response.usage.prompt_tokens + response.usage.completion_tokens + + +def test_transform_response_uses_a_blank_message_when_the_prediction_has_no_output(): + response = _transform(httpx.Response(200, json={"status": "succeeded"})) + + assert response.choices[0].message.content == " " + + +@pytest.mark.parametrize("output", [None, 7, True, ["secret-output", 7], ["secret-output", None], [["secret-output"]]]) +def test_transform_response_rejects_an_output_that_is_not_made_of_strings_without_echoing_it(output: object): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, json={"status": "succeeded", "output": output})) + + assert "secret-output" not in str(exc_info.value) diff --git a/tests/unit/llms/sagemaker/test_sagemaker_common_utils.py b/tests/unit/llms/sagemaker/test_sagemaker_common_utils.py index e2dd3bca74f..51a2ccb5640 100644 --- a/tests/unit/llms/sagemaker/test_sagemaker_common_utils.py +++ b/tests/unit/llms/sagemaker/test_sagemaker_common_utils.py @@ -70,6 +70,19 @@ def test_sagemaker_response_stream_shape_load_failure_returns_none(): assert shape is None +@pytest.mark.parametrize("service_model", [["shapes"], None]) +def test_sagemaker_response_stream_shape_is_none_for_a_service_model_that_is_not_a_mapping( + service_model: object, +): + pytest.importorskip("botocore") + from unittest.mock import patch + + import litellm.llms.sagemaker.common_utils as mod + + with patch("botocore.loaders.Loader.load_service_model", return_value=service_model): + assert mod._load_sagemaker_response_stream_shape() is None + + def test_sagemaker_response_stream_shape_is_structure_shape(): """ The loaded shape should be the botocore StructureShape for diff --git a/tests/unit/llms/snowflake/embedding/test_snowflake_embedding.py b/tests/unit/llms/snowflake/embedding/test_snowflake_embedding.py index 8f4551bc79b..ff2af8a8d7b 100644 --- a/tests/unit/llms/snowflake/embedding/test_snowflake_embedding.py +++ b/tests/unit/llms/snowflake/embedding/test_snowflake_embedding.py @@ -2,9 +2,15 @@ import os import json import copy -from unittest.mock import patch +from unittest.mock import MagicMock, patch + +import httpx +import pytest +from pydantic import ValidationError import litellm +from litellm.llms.snowflake.embedding.transformation import SnowflakeEmbeddingConfig +from litellm.types.utils import EmbeddingResponse model_name = "snowflake-arctic-embed" @@ -94,3 +100,59 @@ def test_snowflake_env(mock_post): os.environ.pop("SNOWFLAKE_ACCOUNT_ID", None) os.environ.pop("SNOWFLAKE_JWT", None) + + +def _transform(raw_response: httpx.Response) -> EmbeddingResponse: + return SnowflakeEmbeddingConfig().transform_embedding_response( + model="snowflake-arctic-embed-m", + raw_response=raw_response, + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="test-key", + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +def test_transform_embedding_response_flattens_each_vector_and_prefixes_the_model(): + response = _transform( + httpx.Response( + 200, + json={ + "object": "list", + "model": "snowflake-arctic-embed-m", + "data": [{"object": "embedding", "index": 0, "embedding": [[0.1, 2]]}], + "usage": {"total_tokens": 5}, + "unknown": "ignored", + }, + ) + ) + + assert response.model_dump() == { + "model": "snowflake/snowflake-arctic-embed-m", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 2]}], + "object": "list", + "usage": { + "completion_tokens": 0, + "prompt_tokens": 0, + "total_tokens": 5, + "completion_tokens_details": None, + "prompt_tokens_details": None, + }, + } + assert response._hidden_params["model"] == "snowflake-arctic-embed-m" + + +@pytest.mark.parametrize("body", [b"null", b"[]", b"7"]) +def test_transform_embedding_response_body_that_is_not_an_object_raises_type_error(body: bytes): + with pytest.raises(TypeError): + _transform(httpx.Response(200, content=body)) + + +def test_transform_embedding_response_invalid_envelope_field_is_reported_by_name(): + with pytest.raises(ValidationError) as exc_info: + _transform(httpx.Response(200, json={"data": [], "model": 5})) + + assert exc_info.value.title == "EmbeddingResponse" + assert [error["loc"] for error in exc_info.value.errors()] == [("model",)] diff --git a/tests/unit/llms/stability/image_generation/test_stability_image_generation.py b/tests/unit/llms/stability/image_generation/test_stability_image_generation.py index c5a78603f9c..42eb1a2365b 100644 --- a/tests/unit/llms/stability/image_generation/test_stability_image_generation.py +++ b/tests/unit/llms/stability/image_generation/test_stability_image_generation.py @@ -9,6 +9,7 @@ from unittest.mock import MagicMock import httpx import pytest +from pydantic import ValidationError from litellm.llms.stability.image_generation import StabilityImageGenerationConfig from litellm.types.llms.stability import ( @@ -304,3 +305,35 @@ class TestStabilityGenerationModels: STABILITY_GENERATION_MODELS["stable-image-core"] == "/v2beta/stable-image/generate/core" ) + + +def _transform(payload: object) -> ImageResponse: + return StabilityImageGenerationConfig().transform_image_generation_response( + model="sd3", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_transform_image_generation_response_returns_the_base64_image(): + response = _transform({"image": "QUJD", "finish_reason": "SUCCESS", "seed": 7}) + + assert [(image.b64_json, image.url) for image in response.data] == [("QUJD", None)] + + +@pytest.mark.parametrize("payload", [{}, {"image": ""}, {"image": None, "finish_reason": None}]) +def test_transform_image_generation_response_without_an_image_has_no_data(payload: dict[str, object]): + assert _transform(payload).data == [] + + +@pytest.mark.parametrize("payload", ["base64encodedimage==", ["base64encodedimage=="]]) +def test_transform_image_generation_response_rejects_non_object_bodies(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "base64encodedimage" not in str(exc_info.value) diff --git a/tests/unit/llms/tinyfish/test_tinyfish_search.py b/tests/unit/llms/tinyfish/test_tinyfish_search.py index 69afbb416aa..e7020238ba6 100644 --- a/tests/unit/llms/tinyfish/test_tinyfish_search.py +++ b/tests/unit/llms/tinyfish/test_tinyfish_search.py @@ -7,6 +7,7 @@ from unittest.mock import MagicMock, patch import httpx import pytest +from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.tinyfish.search.transformation import ( TinyfishSearchConfig, _append_domain_filters, @@ -821,3 +822,56 @@ class TestDefaultMissingResultFields: assert raw_json["results"][0] == "string item" assert raw_json["results"][1] == 42 assert raw_json["results"][2] == {"title": "ok", "url": "", "snippet": ""} + + +def test_transform_search_response_defaults_missing_fields_of_a_decoded_http_body(): + payload = { + "query": "q", + "results": [ + {"title": "T", "url": "https://a.example", "snippet": "S", "position": 1}, + {"title": None, "position": 2}, + ], + } + + response = TinyfishSearchConfig().transform_search_response( + raw_response=httpx.Response(200, json=payload, headers={"x-request-id": "req-1"}), + logging_obj=None, + ) + + assert [result.model_dump(exclude_none=True) for result in response.results] == [ + {"title": "T", "url": "https://a.example", "snippet": "S", "position": 1}, + {"title": "", "url": "", "snippet": "", "position": 2}, + ] + assert response.model_dump()["query"] == "q" + assert response._hidden_params["headers"]["x-request-id"] == "req-1" + + +@pytest.mark.parametrize("payload", [[], "text", 5, {}, {"results": "text"}, {"results": [5]}]) +def test_transform_search_response_wraps_a_decoded_body_of_the_wrong_shape(payload: object): + with pytest.raises(BaseLLMException, match="TinyFish Search: Response shape does not match") as exc_info: + TinyfishSearchConfig().transform_search_response( + raw_response=httpx.Response(200, json=payload), logging_obj=None + ) + + assert exc_info.value.status_code == 200 + + +@pytest.mark.parametrize( + ("body", "expected"), + [ + ('{"error": {"code": "INVALID_INPUT", "message": "query is required"}}', "query is required"), + ('{"error": {"message": ""}}', '{"error": {"message": ""}}'), + ('{"error": {"message": 5}}', '{"error": {"message": 5}}'), + ('{"error": "flat"}', '{"error": "flat"}'), + ('["not", "an", "envelope"]', '["not", "an", "envelope"]'), + ("Bad Gateway", "Bad Gateway"), + ], +) +def test_transform_search_response_unwraps_only_the_tinyfish_error_envelope(body: str, expected: str): + with pytest.raises(BaseLLMException) as exc_info: + TinyfishSearchConfig().transform_search_response(raw_response=httpx.Response(400, text=body), logging_obj=None) + + assert ( + exc_info.value.message == f"TinyFish Search: {expected}. See https://docs.tinyfish.ai/search-api for details." + ) + assert exc_info.value.status_code == 400 diff --git a/tests/unit/llms/together_ai/rerank/__init__.py b/tests/unit/llms/together_ai/rerank/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/together_ai/rerank/test_handler.py b/tests/unit/llms/together_ai/rerank/test_handler.py new file mode 100644 index 00000000000..bc2f30364d2 --- /dev/null +++ b/tests/unit/llms/together_ai/rerank/test_handler.py @@ -0,0 +1,96 @@ +import json +from collections.abc import Iterator + +import httpx +import pytest +import respx +from pydantic import ValidationError + +import litellm +from litellm.llms.together_ai.rerank.handler import TogetherAIRerank + +API_BASE = "https://api.together.example/v1" +RERANKED = { + "id": "rr-1", + "results": [ + {"index": 1, "relevance_score": 0.9, "document": {"text": "Paris"}}, + {"index": 0, "relevance_score": 0.1}, + ], + "usage": {"total_tokens": 7}, +} +NON_OBJECT_BODIES = [7, "sensitive-document", [{"id": "sensitive-document"}]] + + +@pytest.fixture +def rerank_route(monkeypatch: pytest.MonkeyPatch) -> Iterator[respx.Route]: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + with respx.mock as router: + yield router.post(f"{API_BASE}/rerank") + + +def _rerank(*, is_async: bool): + return TogetherAIRerank().rerank( + model="Salesforce/Llama-Rank-V1", + api_key="sk-test", + api_base=API_BASE, + query="capital of france", + documents=["Berlin", {"text": "Paris"}], + top_n=1, + _is_async=is_async, + ) + + +def _assert_is_the_reranking(response: litellm.RerankResponse, route: respx.Route) -> None: + assert response.id == "rr-1" + assert response.results == [ + {"index": 1, "relevance_score": 0.9, "document": {"text": "Paris"}}, + {"index": 0, "relevance_score": 0.1}, + ] + assert response.meta == {"billed_units": {"total_tokens": 7}, "tokens": {}} + assert route.calls.last.request.headers["authorization"] == "Bearer sk-test" + assert json.loads(route.calls.last.request.content) == { + "model": "Salesforce/Llama-Rank-V1", + "query": "capital of france", + "top_n": 1, + "documents": ["Berlin", {"text": "Paris"}], + "return_documents": True, + } + + +def test_rerank_returns_the_upstream_ranking(rerank_route: respx.Route): + rerank_route.mock(return_value=httpx.Response(200, json=RERANKED)) + + _assert_is_the_reranking(_rerank(is_async=False), rerank_route) + + +async def test_async_rerank_returns_the_upstream_ranking(rerank_route: respx.Route): + rerank_route.mock(return_value=httpx.Response(200, json=RERANKED)) + + _assert_is_the_reranking(await _rerank(is_async=True), rerank_route) + + +@pytest.mark.parametrize("payload", NON_OBJECT_BODIES) +def test_rerank_rejects_non_object_bodies(rerank_route: respx.Route, payload: object): + rerank_route.mock(return_value=httpx.Response(200, json=payload)) + + with pytest.raises(ValidationError) as exc_info: + _rerank(is_async=False) + + assert "sensitive-document" not in str(exc_info.value) + + +@pytest.mark.parametrize("payload", NON_OBJECT_BODIES) +async def test_async_rerank_rejects_non_object_bodies(rerank_route: respx.Route, payload: object): + rerank_route.mock(return_value=httpx.Response(200, json=payload)) + + with pytest.raises(ValidationError) as exc_info: + await _rerank(is_async=True) + + assert "sensitive-document" not in str(exc_info.value) + + +def test_rerank_without_results_is_a_value_error(rerank_route: respx.Route): + rerank_route.mock(return_value=httpx.Response(200, json={"id": "rr-1"})) + + with pytest.raises(ValueError, match="No results found"): + _rerank(is_async=False) diff --git a/tests/unit/llms/vertex_ai/image_generation/test_image_generation_handler.py b/tests/unit/llms/vertex_ai/image_generation/test_image_generation_handler.py new file mode 100644 index 00000000000..ac5e14c0388 --- /dev/null +++ b/tests/unit/llms/vertex_ai/image_generation/test_image_generation_handler.py @@ -0,0 +1,43 @@ +import pytest +from pydantic import ValidationError + +from litellm.llms.vertex_ai.image_generation.image_generation_handler import VertexImageGeneration +from litellm.types.utils import ImageResponse + + +def _process(json_response: dict[str, object]) -> ImageResponse: + return VertexImageGeneration().process_image_generation_response( + json_response, ImageResponse(), "imagegeneration@006" + ) + + +@pytest.mark.parametrize( + ("predictions", "expected"), + [ + ([{"bytesBase64Encoded": "QUJD", "mimeType": "image/png"}, {"bytesBase64Encoded": "REVG"}], ["QUJD", "REVG"]), + ([{"bytesBase64Encoded": None}], [None]), + ([], []), + ], +) +def test_process_image_generation_response_maps_each_prediction( + predictions: list[dict[str, object]], expected: list[str | None] +): + response = _process({"predictions": predictions, "deployedModelId": "1"}) + + assert [image.b64_json for image in response.data] == expected + + +@pytest.mark.parametrize( + "predictions", + [None, 7, ["QUJD"], [{"bytesBase64Encoded": "QUJD"}, None], [{"bytesBase64Encoded": 7}]], +) +def test_process_image_generation_response_rejects_malformed_predictions(predictions: object): + with pytest.raises(ValidationError) as exc_info: + _process({"predictions": predictions}) + + assert "QUJD" not in str(exc_info.value) + + +def test_process_image_generation_response_requires_the_encoded_bytes_key(): + with pytest.raises(KeyError): + _process({"predictions": [{"mimeType": "image/png"}]}) diff --git a/tests/unit/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py b/tests/unit/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py index a72a570c2a2..0461c739395 100644 --- a/tests/unit/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py +++ b/tests/unit/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py @@ -1,6 +1,8 @@ from unittest.mock import MagicMock, patch import httpx +import pytest +from pydantic import ValidationError from litellm.llms.vertex_ai.image_generation import ( @@ -12,6 +14,7 @@ from litellm.llms.vertex_ai.image_generation.vertex_gemini_transformation import from litellm.llms.vertex_ai.image_generation.vertex_imagen_transformation import ( VertexAIImagenImageGenerationConfig, ) +from litellm.types.utils import ImageResponse class TestVertexAIGeminiImageGenerationConfig: @@ -635,3 +638,64 @@ class TestVertexAIImageGenerationIntegration: assert "us-central1" in url assert "imagegeneration@006" in url assert "predict" in url + + +def _transform_gemini_response(payload: object) -> ImageResponse: + return VertexAIGeminiImageGenerationConfig().transform_image_generation_response( + model="gemini-2.5-flash-image", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_gemini_image_generation_response_maps_usage_by_modality(): + response = _transform_gemini_response( + { + "candidates": [{"content": {"parts": [{"inlineData": {"data": "aGVsbG8="}}]}}], + "usageMetadata": { + "promptTokenCount": 5, + "candidatesTokenCount": 7, + "totalTokenCount": 12, + "promptTokensDetails": [ + {"modality": "TEXT", "tokenCount": 3}, + {"modality": "IMAGE", "tokenCount": 2}, + ], + }, + } + ) + + assert [image.b64_json for image in response.data] == ["aGVsbG8="] + assert response.usage.input_tokens == 5 + assert response.usage.output_tokens == 7 + assert response.usage.total_tokens == 12 + assert response.usage.input_tokens_details.text_tokens == 3 + assert response.usage.input_tokens_details.image_tokens == 2 + + +@pytest.mark.parametrize("usage_metadata", [None, {}, [], "", 0]) +def test_gemini_image_generation_response_with_falsy_usage_metadata_keeps_zeroed_usage(usage_metadata: object): + response = _transform_gemini_response( + { + "candidates": [{"content": {"parts": []}, "groundingMetadata": {"webSearchQueries": ["a"]}}], + "usageMetadata": usage_metadata, + } + ) + + assert response.data == [] + assert (response.usage.input_tokens, response.usage.output_tokens, response.usage.total_tokens) == (0, 0, 0) + assert response.usage.web_search_requests == 1 + + +@pytest.mark.parametrize("usage_metadata", ["not an object", ["not", "an", "object"], 7, True]) +def test_gemini_image_generation_response_rejects_non_object_usage_metadata_without_echoing_it( + usage_metadata: object, +): + with pytest.raises(ValidationError) as exc_info: + _transform_gemini_response({"candidates": [], "usageMetadata": usage_metadata}) + + assert "input_value" not in str(exc_info.value) 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 fd667e8f425..8c73b72a65a 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 @@ -4,6 +4,7 @@ from unittest.mock import MagicMock, Mock, patch import httpx import pytest +from pydantic import ValidationError import litellm from litellm.llms.vertex_ai.text_to_speech.transformation import ( @@ -633,3 +634,32 @@ def test_litellm_speech_vertex_ai_chirp(mock_get_token, mock_ensure_token, mock_ assert "headers" in call_kwargs assert "Authorization" in call_kwargs["headers"] assert call_kwargs["headers"]["Authorization"] == "Bearer mock-token" + + +@pytest.mark.parametrize("payload", [{}, {"audioContent": ""}, {"audioContent": None}, {"audioContent": []}]) +def test_transform_text_to_speech_response_without_audio_content_reports_it_missing(payload: dict[str, object]): + with pytest.raises(ValueError, match="No audioContent in Vertex AI TTS response"): + VertexAITextToSpeechConfig().transform_text_to_speech_response( + model="vertex_ai/chirp", + raw_response=httpx.Response(200, json=payload), + logging_obj=MagicMock(), + ) + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"audioContent": 7}, + {"audioContent": ["UklGRiQAAABXQVZFZm10IA=="]}, + ], +) +def test_transform_text_to_speech_response_rejects_malformed_payloads_without_echoing_them(payload: object): + with pytest.raises(ValidationError) as exc_info: + VertexAITextToSpeechConfig().transform_text_to_speech_response( + model="vertex_ai/chirp", + raw_response=httpx.Response(200, json=payload), + logging_obj=MagicMock(), + ) + + assert "input_value" not in str(exc_info.value) 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 5eb4bf31845..777e52f4e2a 100644 --- a/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py +++ b/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py @@ -8,6 +8,7 @@ from unittest.mock import MagicMock, patch import httpx import pytest +from pydantic import ValidationError from litellm.llms.voyage.rerank.transformation import VoyageRerankConfig from litellm.types.rerank import RerankResponse @@ -343,3 +344,81 @@ class TestVoyageRerankTransform: assert prompt_cost == 0.0 assert completion_cost == 0.0 + + +def _transform(payload: object) -> RerankResponse: + return VoyageRerankConfig().transform_rerank_response( + model="rerank-2.5", + raw_response=httpx.Response(200, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + ) + + +def test_transform_rerank_response_keeps_provider_id_and_drops_unknown_result_fields(): + response = _transform( + { + "id": "rerank-1", + "data": [ + {"index": 2, "relevance_score": 0.25, "document": {"text": "doc", "extra": 1}, "extra": 2}, + {"index": 0, "relevance_score": 1, "document": "plain"}, + ], + "usage": {"total_tokens": 9}, + } + ) + + assert response.id == "rerank-1" + assert response.results == [ + {"index": 2, "relevance_score": 0.25, "document": {"text": "doc"}}, + {"index": 0, "relevance_score": 1.0, "document": {"text": "plain"}}, + ] + + +@pytest.mark.parametrize( + ("payload", "expected_total_tokens"), + [ + ({"data": []}, 0), + ({"data": [], "usage": {}}, 0), + ({"data": [], "usage": {"total_tokens": 79}}, 79), + ({"data": [], "usage": {"total_tokens": None}}, None), + ], +) +def test_transform_rerank_response_reads_total_tokens_from_usage( + payload: dict[str, object], expected_total_tokens: int | None +): + assert _transform(payload).meta == { + "billed_units": {"total_tokens": expected_total_tokens}, + "tokens": {"input_tokens": expected_total_tokens, "output_tokens": 0}, + } + + +def test_transform_rerank_response_empty_data_list_yields_no_results(): + assert _transform({"id": "rerank-1", "data": []}).results == [] + + +@pytest.mark.parametrize( + "payload", + [ + ["not", "an", "object"], + {"data": 7}, + {"data": ["not an object"]}, + {"data": [{"index": 0, "relevance_score": 0.5}], "usage": None}, + {"data": [{"index": 0, "relevance_score": 0.5}], "usage": {"total_tokens": 1.5}}, + {"data": [{"index": 0, "relevance_score": 0.5}], "id": 7}, + ], +) +def test_transform_rerank_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError): + _transform(payload) + + +def test_transform_rerank_response_result_without_index_raises_key_error(): + with pytest.raises(KeyError, match="index"): + _transform({"data": [{"relevance_score": 0.5}]}) + + +def test_transform_rerank_response_shape_errors_do_not_echo_the_payload(): + with pytest.raises(ValidationError) as exc_info: + _transform({"data": [["leaked document text"]]}) + + assert "leaked document text" not in str(exc_info.value) diff --git a/tests/unit/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py b/tests/unit/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py index efe592f515e..b6e5007c456 100644 --- a/tests/unit/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py +++ b/tests/unit/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py @@ -6,6 +6,10 @@ Validates the WatsonX transcription response transformation. from unittest.mock import MagicMock +import httpx +import pytest +from pydantic import ValidationError + from litellm.llms.watsonx.audio_transcription.transformation import ( IBMWatsonXAudioTranscriptionConfig, ) @@ -83,3 +87,57 @@ class TestWatsonXAudioTranscription: # Verify duration is set via dictionary assignment assert result["duration"] == 5.5 + + +def _transform_transcription_response(payload: object) -> TranscriptionResponse: + return IBMWatsonXAudioTranscriptionConfig().transform_audio_transcription_response( + httpx.Response(200, json=payload) + ) + + +@pytest.mark.parametrize( + ("payload", "expected_extras"), + [ + ({"text": "hello"}, {}), + ({"text": "hello", "model": "whisper-large-v3-turbo"}, {}), + ( + {"text": "hello", "duration": 1.5, "language": "en", "task": "transcribe"}, + {"duration": 1.5, "language": "en", "task": "transcribe"}, + ), + ( + {"text": "hello", "segments": [{"id": 0, "text": "hello"}], "words": None}, + {"segments": [{"id": 0, "text": "hello"}], "words": None}, + ), + ], +) +def test_transform_audio_transcription_response_copies_every_field_except_model( + payload: dict[str, object], expected_extras: dict[str, object] +) -> None: + response = _transform_transcription_response(payload) + + assert response.text == "hello" + assert not hasattr(response, "model") + assert {key: response[key] for key in expected_extras} == expected_extras + + +@pytest.mark.parametrize("payload", [{}, {"model": "whisper"}, {"duration": 2.0}, {"text": None, "usage": None}]) +def test_transform_audio_transcription_response_without_text_or_usage_reports_the_body( + payload: dict[str, object], +) -> None: + with pytest.raises(ValueError, match="Invalid response format") as exc_info: + _transform_transcription_response(payload) + + assert exc_info.value.args == ( + "Invalid response format. Received response does not match the expected format. Got: ", + payload, + ) + + +@pytest.mark.parametrize("payload", [["not", "an", "object"], "plain text", 7, True]) +def test_transform_audio_transcription_response_rejects_non_object_bodies_without_echoing_them( + payload: object, +) -> None: + with pytest.raises(ValidationError) as exc_info: + _transform_transcription_response(payload) + + assert "input_value" not in str(exc_info.value) diff --git a/tests/unit/llms/watsonx/embed/test_watsonx_embedding_transformation.py b/tests/unit/llms/watsonx/embed/test_watsonx_embedding_transformation.py index 58f6bb23498..afcebcf0990 100644 --- a/tests/unit/llms/watsonx/embed/test_watsonx_embedding_transformation.py +++ b/tests/unit/llms/watsonx/embed/test_watsonx_embedding_transformation.py @@ -1,8 +1,13 @@ +from unittest.mock import MagicMock + +import httpx import pytest +from pydantic import ValidationError from litellm.llms.watsonx.embed.transformation import IBMWatsonXEmbeddingConfig +from litellm.types.utils import EmbeddingResponse class TestIBMWatsonXEmbeddingConfig: @@ -42,3 +47,88 @@ class TestIBMWatsonXEmbeddingConfig: optional_params={"project_id": "test-project-id"}, headers={}, ) + + +def _transform(payload: object) -> EmbeddingResponse: + return IBMWatsonXEmbeddingConfig().transform_embedding_response( + model="ibm/slate-125m-english-rtrvr", + raw_response=httpx.Response(200, json=payload), + model_response=EmbeddingResponse(), + logging_obj=MagicMock(), + api_key="test-key", + request_data={}, + optional_params={}, + litellm_params={}, + ) + + +def test_transform_embedding_response_numbers_each_result_and_bills_the_input_tokens(): + response = _transform( + { + "model_id": "ibm/slate-125m-english-rtrvr", + "results": [{"embedding": [0.1, 2], "input": "ignored"}, {"embedding": [0.5]}], + "input_token_count": 7, + } + ) + + assert response.object == "list" + assert response.data == [ + {"object": "embedding", "index": 0, "embedding": [0.1, 2]}, + {"object": "embedding", "index": 1, "embedding": [0.5]}, + ] + assert (response.usage.prompt_tokens, response.usage.completion_tokens, response.usage.total_tokens) == (7, 0, 7) + + +@pytest.mark.parametrize( + ("payload", "expected_tokens"), + [ + ({"input_token_count": None}, 0), + ({"results": []}, 0), + ({}, 0), + ], +) +def test_transform_embedding_response_defaults_missing_results_and_token_count( + payload: dict[str, object], expected_tokens: int +): + response = _transform(payload) + + assert response.data == [] + assert (response.usage.prompt_tokens, response.usage.total_tokens) == (expected_tokens, expected_tokens) + + +@pytest.mark.parametrize( + "payload", + [ + {"results": None}, + {"results": 7}, + {"results": ["not an object"]}, + {"results": [{"embedding": [0.1]}, None]}, + {"input_token_count": 1.5}, + {"input_token_count": "many"}, + ], +) +def test_transform_embedding_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError): + _transform(payload) + + +@pytest.mark.parametrize("payload", [["not", "an", "object"], "not an object"]) +def test_transform_embedding_response_non_object_body_raises_attribute_error(payload: object): + with pytest.raises(AttributeError): + _transform(payload) + + +def test_transform_embedding_response_result_without_embedding_raises_key_error(): + with pytest.raises(KeyError, match="embedding"): + _transform({"results": [{"input": "x"}]}) + + +@pytest.mark.parametrize( + "payload", + [{"results": ["leaked payload text"]}, {"results": {"leaked payload text": 1}}], +) +def test_transform_embedding_response_shape_errors_do_not_echo_the_payload(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "leaked payload text" not in str(exc_info.value) diff --git a/tests/unit/llms/watsonx/rerank/test_watsonx_rerank.py b/tests/unit/llms/watsonx/rerank/test_watsonx_rerank.py index ccbd318959f..fbdbddacd5f 100644 --- a/tests/unit/llms/watsonx/rerank/test_watsonx_rerank.py +++ b/tests/unit/llms/watsonx/rerank/test_watsonx_rerank.py @@ -9,6 +9,7 @@ from unittest.mock import MagicMock import httpx import pytest +from pydantic import ValidationError from litellm.llms.watsonx.common_utils import ( WatsonXAIError, @@ -260,3 +261,64 @@ class TestIBMWatsonXRerankTransform: assert "return_documents" in supported_params assert "max_tokens_per_doc" in supported_params assert len(supported_params) == 5 + + +def _transform(payload: object) -> RerankResponse: + return IBMWatsonXRerankConfig().transform_rerank_response( + model="watsonx/cross-encoder/ms-marco-minilm-l-12-v2", + raw_response=httpx.Response(200, json=payload), + model_response=RerankResponse(), + logging_obj=MagicMock(), + ) + + +def test_transform_rerank_response_keeps_the_upstream_id_documents_and_token_count(): + response = _transform( + { + "id": "rerank-1", + "results": [ + {"index": 1, "score": 0.5, "input": "plain text"}, + {"index": 0, "score": 0.25, "input": {"text": "object text"}}, + ], + "input_token_count": 62, + } + ) + + assert response.id == "rerank-1" + assert response.results == [ + {"index": 1, "relevance_score": 0.5, "document": {"text": "plain text"}}, + {"index": 0, "relevance_score": 0.25, "document": {"text": "object text"}}, + ] + assert response.meta == {"tokens": {"input_tokens": 62}} + + +def test_transform_rerank_response_without_a_token_count_reports_zero_tokens(): + response = _transform({"results": []}) + + assert response.results == [] + assert response.meta == {"tokens": {"input_tokens": 0}} + + +@pytest.mark.parametrize( + "payload", + [ + ["sensitive-document"], + "sensitive-document", + {"results": "sensitive-document"}, + {"results": 7}, + {"results": ["sensitive-document"]}, + {"results": [{"index": 0, "score": 0.5}, None]}, + {"results": [{"index": 0, "score": 0.5}], "id": 7}, + {"results": [{"index": 0, "score": 0.5}], "input_token_count": "many"}, + ], +) +def test_transform_rerank_response_rejects_malformed_payloads(payload: object): + with pytest.raises(ValidationError) as exc_info: + _transform(payload) + + assert "sensitive-document" not in str(exc_info.value) + + +def test_transform_rerank_response_requires_index_and_score_on_each_result(): + with pytest.raises(KeyError): + _transform({"results": [{"index": 0}]}) 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 cfcff73b857..2e255ddf853 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py @@ -173,6 +173,49 @@ def test_mcp_oauth_token_identity_detects_change_under_encryption(): assert mcp_oauth_token_identity(unchanged) != mcp_oauth_token_identity(changed) +def test_mcp_oauth_token_identity_reads_client_and_scopes_from_json_string_credentials(): + from litellm.proxy._experimental.mcp_server.db import mcp_oauth_token_identity + + credentials: Final = json.dumps( + {"client_id": "cid", "client_secret": "csec", "scopes": ["a"], "upstream_resource": "api://audience"} + ) + + assert mcp_oauth_token_identity(_identity_server(credentials=credentials)) == ( + "https://up.example.com/mcp", + None, + "oauth2", + "authorization_code", + None, + "https://idp.example.com/authorize", + "https://idp.example.com/token", + "https://idp.example.com/register", + "cid", + "csec", + ["a"], + "api://audience", + ) + + +@pytest.mark.parametrize("credentials", ["[]", '["client_id"]', '"cid"', "5", "null", "not json"]) +def test_mcp_oauth_token_identity_treats_stored_credentials_that_are_not_a_json_object_as_empty(credentials): + from litellm.proxy._experimental.mcp_server.db import mcp_oauth_token_identity + + assert mcp_oauth_token_identity(_identity_server(credentials=credentials)) == ( + "https://up.example.com/mcp", + None, + "oauth2", + "authorization_code", + None, + "https://idp.example.com/authorize", + "https://idp.example.com/token", + "https://idp.example.com/register", + None, + None, + None, + None, + ) + + def _oauth_row(user_id: str, server_id: str = "srv-1"): """A stored per-user OAuth token row (payload tagged type=oauth2, legacy plain-base64 encoding).""" row = _legacy_row(json.dumps({"type": "oauth2", "access_token": "tok-" + user_id})) 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 4e27ec134d4..6d71256ff18 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -11,6 +11,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException +from litellm.proxy._types import LiteLLM_MCPServerTable from litellm.types.mcp import MCPAuth if TYPE_CHECKING: @@ -1365,6 +1366,8 @@ async def test_register_client_persists_dcr_client_identity(): mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) + prisma = MagicMock() + prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None) mock_update = AsyncMock(return_value=MagicMock()) mock_update_server = AsyncMock() @@ -1373,7 +1376,7 @@ async def test_register_client_persists_dcr_client_identity(): "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", return_value=mock_async_client, ), - patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()), + patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=prisma), patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=mock_update), patch.object(global_mcp_server_manager, "update_server", new=mock_update_server), ): @@ -1390,7 +1393,12 @@ async def test_register_client_persists_dcr_client_identity(): import json assert response.status_code == 200 - assert json.loads(response.body.decode("utf-8")) == mock_response.json.return_value + assert json.loads(response.body.decode("utf-8")) == { + **mock_response.json.return_value, + "dcr_issuer": None, + "dcr_server_url": None, + "dcr_redirect_uris": ["https://proxy.litellm.example/callback"], + } mock_update.assert_called_once() update_data = mock_update.call_args.kwargs["data"] @@ -1452,6 +1460,8 @@ async def _register_persistence_attempted_for_auth_type(auth_type: MCPAuth) -> b mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) + prisma = MagicMock() + prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None) mock_update = AsyncMock(return_value=MagicMock()) with ( @@ -1459,7 +1469,7 @@ async def _register_persistence_attempted_for_auth_type(auth_type: MCPAuth) -> b "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", return_value=mock_async_client, ), - patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()), + patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=prisma), patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=mock_update), patch.object(global_mcp_server_manager, "update_server", new=AsyncMock()), ): @@ -1473,7 +1483,12 @@ async def _register_persistence_attempted_for_auth_type(auth_type: MCPAuth) -> b persist_credentials=True, ) - assert json.loads(response.body.decode("utf-8")) == mock_response.json.return_value + expected_binding = { + "dcr_issuer": None, + "dcr_server_url": None, + "dcr_redirect_uris": ["https://proxy.litellm.example/callback"], + } if auth_type == MCPAuth.oauth2 else {} + assert json.loads(response.body.decode("utf-8")) == {**mock_response.json.return_value, **expected_binding} return mock_update.await_count > 0 @@ -1529,7 +1544,7 @@ async def test_register_client_persists_only_to_its_own_row_when_another_server_ ) sibling_row_with_client = MagicMock(server_id="server-a", url=shared_url) sibling_row_with_client.credentials = {"client_id": "client-a-do-not-adopt"} - own_row_without_client = MagicMock(server_id="server-b", url=shared_url) + own_row_without_client = LiteLLM_MCPServerTable.model_validate(fresh_server.model_dump(exclude_none=True)) own_row_without_client.credentials = {} rows_by_server_id = {"server-a": sibling_row_with_client, "server-b": own_row_without_client} @@ -1622,6 +1637,8 @@ async def test_register_client_does_not_clobber_token_url_when_absent(): mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) + prisma = MagicMock() + prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None) mock_update = AsyncMock(return_value=MagicMock()) mock_update_server = AsyncMock() @@ -1630,7 +1647,7 @@ async def test_register_client_does_not_clobber_token_url_when_absent(): "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", return_value=mock_async_client, ), - patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()), + patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=prisma), patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=mock_update), patch.object(global_mcp_server_manager, "update_server", new=mock_update_server), ): @@ -1685,7 +1702,7 @@ async def test_register_client_reuses_persisted_client_id_for_non_admin_when_reg mock_request.base_url = "https://proxy.litellm.example/" mock_request.headers = {} - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = {"client_id": "persisted-client"} mock_get_mcp_server = AsyncMock(return_value=persisted_server) mock_update_mcp_server = AsyncMock() @@ -1761,7 +1778,7 @@ async def test_register_client_reuse_refreshes_request_server_when_manager_updat mock_request.base_url = "https://proxy.litellm.example/" mock_request.headers = {} - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = { "client_id": "persisted-client", "client_secret": "persisted-secret", @@ -1803,7 +1820,8 @@ async def test_register_client_reuse_refreshes_request_server_when_manager_updat @pytest.mark.asyncio -async def test_register_client_returns_reused_client_when_concurrent_persist_wins(): +@pytest.mark.parametrize("late_winner", [False, True]) +async def test_register_client_returns_reused_client_when_concurrent_persist_wins(late_winner): try: from fastapi import Request @@ -1843,10 +1861,11 @@ async def test_register_client_returns_reused_client_when_concurrent_persist_win mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = {"client_id": "persisted-client"} - mock_get_mcp_server = AsyncMock(side_effect=[None, persisted_server]) - mock_update_mcp_server = AsyncMock() + empty_server = persisted_server.model_copy(update={"credentials": None}) + mock_get_mcp_server = AsyncMock(side_effect=[empty_server, empty_server, persisted_server] if late_winner else [empty_server, persisted_server]) + mock_update_mcp_server = AsyncMock(return_value=None) mock_update_server = AsyncMock() with ( @@ -1878,7 +1897,10 @@ async def test_register_client_returns_reused_client_when_concurrent_persist_win assert response["client_id"] == "remote_server" assert oauth2_server.client_id == "persisted-client" mock_async_client.post.assert_called_once() - mock_update_mcp_server.assert_not_called() + if late_winner: + mock_update_mcp_server.assert_awaited_once() + else: + mock_update_mcp_server.assert_not_called() mock_update_server.assert_called_once_with(persisted_server) @@ -1937,7 +1959,7 @@ async def test_register_client_re_registers_when_persisted_redirect_uri_no_longe mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = { "client_id": "stale-client", "client_secret": "stale-secret", @@ -2014,7 +2036,7 @@ async def test_register_client_grandfathers_persisted_client_without_recorded_re mock_async_client = MagicMock() mock_async_client.post = AsyncMock() - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = {"client_id": "legacy-client"} mock_get_mcp_server = AsyncMock(return_value=persisted_server) mock_update_mcp_server = AsyncMock() @@ -2072,7 +2094,7 @@ async def test_register_client_keeps_persisted_client_when_recorded_redirect_uri mock_async_client = MagicMock() mock_async_client.post = AsyncMock() - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = { "client_id": "kept-client", "redirect_uris": ["https://proxy.litellm.example/callback"], @@ -2137,7 +2159,7 @@ async def test_register_client_non_admin_reuses_persisted_client_despite_redirec mock_async_client = MagicMock() mock_async_client.post = AsyncMock() - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = { "client_id": "persisted-client", "redirect_uris": ["https://old-origin.example/callback"], @@ -4465,6 +4487,173 @@ async def test_callback_error_path_reads_cookie_and_clears_it(monkeypatch): assert cleared[cookie_name]["max-age"] == "0" +def _issuer_anchored_oauth_server(issuer: str | None = "https://idp.example.com"): + from litellm.proxy._types import MCPTransport + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + return MCPServer( + server_id="rfc9207_server", + name="rfc9207", + server_name="rfc9207", + alias="rfc9207", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="upstream-client-id", + issuer=issuer, + authorization_url="https://idp.example.com/oauth/authorize", + token_url="https://idp.example.com/oauth/token", + ) + + +async def _authorize_then_callback(server, iss, monkeypatch, expected_issuer_override=..., error=None): + """Run /authorize for ``server``, then feed the resulting flow back through /callback with the + RFC 9207 ``iss`` the authorization server supposedly returned. Returns (callback_response, + sealed_state_data).""" + from http.cookies import SimpleCookie + from urllib.parse import parse_qs, urlparse + + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _oauth_state_cookie_name, + authorize_with_server, + callback, + decode_state_hash, + encode_state_with_base_url, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-for-LIT-5940") + client_redirect_uri = "http://127.0.0.1:6274/oauth/callback/debug" + + authorize_request = MagicMock(spec=Request) + authorize_request.base_url = "https://proxy.example.com/" + authorize_request.headers = {} + + authorize_response = await authorize_with_server( + request=authorize_request, + mcp_server=server, + client_id="upstream-client-id", + redirect_uri=client_redirect_uri, + state="client-original-state-9207", + code_challenge="challenge", + code_challenge_method="S256", + ) + handle = parse_qs(urlparse(authorize_response.headers["location"]).query)["state"][0] + jar = SimpleCookie() + jar.load(authorize_response.headers["set-cookie"]) + cookie_name = _oauth_state_cookie_name(handle) + sealed_state = jar[cookie_name].value + if expected_issuer_override is not ...: + # A state minted before the issuer was sealed into it: same shape, key absent. + sealed_state = encode_state_with_base_url( + base_url=client_redirect_uri, + original_state="client-original-state-9207", + client_redirect_uri=client_redirect_uri, + expected_issuer=expected_issuer_override, + ) + + callback_request = MagicMock(spec=Request) + callback_request.base_url = "https://proxy.example.com/" + callback_request.headers = {} + callback_request.cookies = {cookie_name: sealed_state} + + response = await callback( + request=callback_request, + code="upstream-auth-code" if error is None else None, + error=error, + state=handle, + iss=iss, + ) + return response, decode_state_hash(sealed_state) + + +@pytest.mark.asyncio +async def test_authorize_seals_the_issuer_and_callback_accepts_a_matching_rfc9207_iss(monkeypatch): + """The callback accepts the exact issuer sealed during authorization.""" + from urllib.parse import parse_qs, urlparse + + response, state_data = await _authorize_then_callback( + _issuer_anchored_oauth_server(), + iss="https://idp.example.com", + monkeypatch=monkeypatch, + ) + + assert state_data["expected_issuer"] == "https://idp.example.com" + assert response.status_code == 302 + query = parse_qs(urlparse(response.headers["location"]).query) + assert query["code"] == ["upstream-auth-code"] + assert query["state"] == ["client-original-state-9207"] + + +@pytest.mark.asyncio +async def test_callback_rejects_authorization_response_from_a_different_issuer(monkeypatch): + """LIT-5940 / RFC 9207 ยง2.4: an ``iss`` naming an authorization server we never sent the user to + is a mix-up attack, so the code must not reach the client's redirect_uri.""" + response, _ = await _authorize_then_callback( + _issuer_anchored_oauth_server(), + iss="https://attacker-idp.example.com", + monkeypatch=monkeypatch, + ) + + assert response.status_code == 400 + assert "location" not in response.headers + assert b"upstream-auth-code" not in response.body + assert b"invalid_issuer" in response.body + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", [None, "access_denied"]) +async def test_callback_rejects_supplied_issuer_without_expected_identity(monkeypatch, error): + response, state_data = await _authorize_then_callback( + _issuer_anchored_oauth_server(issuer=None), + iss="https://unknown-idp.example.com", + monkeypatch=monkeypatch, + error=error, + ) + assert state_data["expected_issuer"] is None + assert response.status_code == 400 + assert "location" not in response.headers + assert b"upstream-auth-code" not in response.body + + +@pytest.mark.asyncio +@pytest.mark.parametrize("issuer", [None, "https://idp.example.com"]) +async def test_callback_accepts_missing_unadvertised_issuer(monkeypatch, issuer): + response, _ = await _authorize_then_callback( + _issuer_anchored_oauth_server(issuer), iss=None, monkeypatch=monkeypatch, + ) + assert response.status_code == 302 + + +@pytest.mark.asyncio +async def test_callback_rejects_an_issuer_differing_only_outside_the_path(monkeypatch): + """A tenant a deployment encoded in a query string is part of that issuer's identity, so the + comparison must not canonicalize it away and let another tenant's response through.""" + response, state_data = await _authorize_then_callback( + _issuer_anchored_oauth_server(issuer="https://idp.example.com/?tenant=a"), + iss="https://idp.example.com/?tenant=b", + monkeypatch=monkeypatch, + ) + + assert state_data["expected_issuer"] == "https://idp.example.com/?tenant=a" + assert response.status_code == 400 + assert "location" not in response.headers + assert b"upstream-auth-code" not in response.body + + +@pytest.mark.asyncio +@pytest.mark.parametrize("iss,expected_status", [(None, 302), ("https://some-idp.example.com", 400)]) +async def test_callback_legacy_state_requires_absent_issuer(monkeypatch, iss, expected_status): + response, _ = await _authorize_then_callback( + _issuer_anchored_oauth_server(), + iss=iss, + monkeypatch=monkeypatch, + expected_issuer_override=None, + ) + assert response.status_code == expected_status + + @pytest.mark.asyncio async def test_oauth_authorize_includes_scopes_from_server_config(): """Test that authorize endpoint includes scopes from server configuration.""" @@ -9070,7 +9259,8 @@ async def test_persist_dcr_client_for_config_server_uses_side_store(): @pytest.mark.asyncio -async def test_hydrate_config_server_applies_stored_dcr_client(monkeypatch): +@pytest.mark.parametrize("legacy", [False, True]) +async def test_hydrate_config_server_applies_stored_dcr_client(monkeypatch, legacy): """On restart a config server's in-memory object has no client_id; hydration overlays the persisted DCR client from the server-scoped store, decrypting the encrypted-at-rest blob, so the refresh_token grant can authenticate as the registered client instead of re-authenticating.""" @@ -9090,12 +9280,14 @@ async def test_hydrate_config_server_applies_stored_dcr_client(monkeypatch): transport=MCPTransport.http, auth_type=MCPAuth.oauth2, client_id=None, + url="https://resource.example/mcp", + issuer="https://idp.example", ) monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-hydrate-key") stored_blob = safe_dumps( encrypt_credentials( - credentials={ + credentials={**({} if legacy else {"dcr_issuer": "https://idp.example", "dcr_server_url": "https://resource.example/mcp"}), "client_id": "stored-client", "client_secret": "stored-secret", "token_endpoint_auth_method": "client_secret_basic", @@ -9145,12 +9337,14 @@ async def test_reuse_config_server_reads_store_with_real_crypto(monkeypatch): transport=MCPTransport.http, auth_type=MCPAuth.oauth2, client_id=None, + url="https://resource.example/mcp", + issuer="https://idp.example", ) monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-reuse-key") blob = safe_dumps( encrypt_credentials( - credentials={"client_id": "stored-client", "client_secret": "sec", "redirect_uris": ["https://x/callback"]}, + credentials={"dcr_issuer": "https://idp.example", "dcr_server_url": "https://resource.example/mcp","client_id": "stored-client", "client_secret": "sec", "redirect_uris": ["https://x/callback"]}, encryption_key="salt-reuse-key", ) ) @@ -12759,3 +12953,232 @@ def test_invalidating_an_idle_server_leaves_no_generation_behind(): finally: for server_id in server_ids: discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server_id, None) + + +@pytest.mark.asyncio +async def test_callback_rejects_missing_advertised_issuer(monkeypatch): + server = _issuer_anchored_oauth_server().model_copy( + update={"authorization_response_iss_parameter_supported": True} + ) + response, _ = await _authorize_then_callback(server, iss=None, monkeypatch=monkeypatch) + assert response.status_code == 400 + assert "location" not in response.headers + assert b"upstream-auth-code" not in response.body + + +@pytest.mark.parametrize("issuer,url", [ + ("https://other.example", "https://resource.example/mcp"), + (None, "https://other.example/mcp"), +]) +def test_persisted_client_is_not_applied_to_another_upstream(issuer, url): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _apply_persisted_dcr_credentials, _PersistedDcrCredentials, + ) + server = _issuer_anchored_oauth_server(issuer).model_copy(update={"url": url, "client_id": None}) + stored = _PersistedDcrCredentials.model_validate({ + "client_id": "old-client", "client_secret": "old-secret", + "dcr_issuer": "https://idp.example.com", "dcr_server_url": "https://resource.example/mcp", + }) + assert _apply_persisted_dcr_credentials(server, stored) is False + assert server.client_id is None + assert server.client_secret is None + + +@pytest.mark.asyncio +async def test_registration_does_not_write_client_after_upstream_edit(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "test-mcp-registration-salt") + from prisma import models + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _persist_dcr_client_registration + + server = _issuer_anchored_oauth_server().model_copy(update={"url": "https://resource.example/mcp", "client_id": None}) + current = models.LiteLLM_MCPServerTable.model_construct( + server_id=server.server_id, url="https://other.example/mcp", issuer="https://other.example", + auth_type="oauth2", transport="http", credentials=None, updated_at=datetime.now(timezone.utc), + ) + prisma = MagicMock() + table = AsyncMock() + table.find_unique.return_value = current + table.update.return_value = current + prisma.db.litellm_mcpservertable = table + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + result = await _persist_dcr_client_registration(server, {"client_id": "old-issuer-client"}, "https://proxy.example/callback") + assert result == "failed" + table.update.assert_not_called() + table.update_many.assert_not_called() + assert server.client_id is None + + +@pytest.mark.parametrize("response_issuer", ["https://IDP.example.com", "https://idp.example.com/", "https://idp.example.com:443"]) +def test_issuer_validation_uses_exact_identifier(response_issuer): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _authorization_response_issuer_is_trusted + assert _authorization_response_issuer_is_trusted(response_issuer, {"expected_issuer": "https://idp.example.com"}) is False + + +@pytest.mark.parametrize("issuer,url", [ + ("https://idp.example.com", "https://resource.example/mcp"), + ("https://idp.example.com", "https://other-resource.example/mcp"), +]) +def test_persisted_client_remains_usable_for_its_issuer(issuer, url, monkeypatch): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _apply_persisted_dcr_credentials, _PersistedDcrCredentials + monkeypatch.setenv("LITELLM_SALT_KEY", "test-registration-binding") + server = _issuer_anchored_oauth_server(issuer).model_copy(update={"url": url, "client_id": None}) + credentials = _PersistedDcrCredentials( + client_id="registered-client", client_secret="registered-secret", dcr_issuer=issuer, + dcr_server_url="https://resource.example/mcp", + ) + assert _apply_persisted_dcr_credentials(server, credentials) is True + assert server.client_id == "registered-client" + assert server.client_secret == "registered-secret" + assert server.dcr_issuer == issuer + + +@pytest.mark.asyncio +async def test_callback_does_not_forward_error_from_another_issuer(monkeypatch): + response, _ = await _authorize_then_callback( + _issuer_anchored_oauth_server(), iss="https://other.example", monkeypatch=monkeypatch, error="access_denied", + ) + assert response.status_code == 400 + assert "location" not in response.headers + assert b"invalid_issuer" in response.body + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["authorize", "token"]) +async def test_bound_client_cannot_be_sent_to_a_different_issuer(monkeypatch, operation): + from fastapi import Request + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="bound-client", name="bound-client", transport="http", auth_type="oauth2", + url="https://new.example/mcp", issuer="https://new.example", client_id="old-client", + dcr_issuer="https://old.example", dcr_server_url="https://old.example/mcp", + ) + monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", AsyncMock(return_value=server)) + http_client = MagicMock() + monkeypatch.setattr(endpoints, "get_async_httpx_client", http_client) + request = ( + endpoints.authorize_with_server( + mcp_server=server, request=MagicMock(spec=Request), redirect_uri="http://localhost/callback", client_id="old-client", + ) + if operation == "authorize" + else endpoints.exchange_token_with_server( + mcp_server=server, request=MagicMock(spec=Request), grant_type="authorization_code", + code="test-code", redirect_uri="http://localhost/callback", client_id="old-client", + client_secret=None, code_verifier="verifier", + ) + ) + with pytest.raises(HTTPException) as error: + await request + assert error.value.status_code == 400 + assert "different issuer" in error.value.detail + http_client.assert_not_called() + + +@pytest.mark.asyncio +async def test_failed_registration_preserves_cached_credentials(monkeypatch): + from fastapi import Request + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + server = _dcr_redirect_test_server("old-client").model_copy(update={ + "url": "https://new.example/mcp", "issuer": "https://new.example", + "dcr_issuer": "https://old.example", "client_secret": "old-secret", + }) + monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", AsyncMock(return_value=server)) + monkeypatch.setattr(endpoints, "_reuse_persisted_dcr_client_if_available", AsyncMock(return_value=False)) + response = MagicMock() + response.json.return_value = {"client_id": "new-client"} + register = AsyncMock(return_value=response) + monkeypatch.setattr(endpoints, "_post_dcr_registration", register) + persist = AsyncMock(return_value="failed") + monkeypatch.setattr(endpoints, "_persist_dcr_client_registration", persist) + request = MagicMock(spec=Request) + request.base_url = "https://gateway.example/" + request.headers = {} + with pytest.raises(HTTPException) as error: + await endpoints.register_client_with_server( + request=request, mcp_server=server, client_name="app", grant_types=["authorization_code"], + response_types=["code"], token_endpoint_auth_method="none", persist_credentials=True, + ) + assert error.value.status_code == 503 + assert "could not be saved" in error.value.detail + assert server.client_id == "old-client" + assert server.client_secret == "old-secret" + register.assert_awaited_once() + persist.assert_awaited_once_with(server, {"client_id": "new-client"}, "https://gateway.example/callback") + + +@pytest.mark.asyncio +async def test_registration_without_database_keeps_client_in_temporary_server(monkeypatch): + from fastapi import Request + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.setattr(proxy_server, "prisma_client", None) + server = _dcr_redirect_test_server(None).model_copy(update={"issuer": "https://idp.example", "url": "https://resource.example/mcp"}) + monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", AsyncMock(return_value=server)) + upstream = MagicMock() + upstream.json.return_value = {"client_id": "temporary-client", "client_secret": "temporary-secret"} + monkeypatch.setattr(endpoints, "_post_dcr_registration", AsyncMock(return_value=upstream)) + request = MagicMock(spec=Request) + request.base_url = "https://gateway.example/" + request.headers = {} + response = await endpoints.register_client_with_server( + request=request, mcp_server=server, client_name="app", grant_types=["authorization_code"], + response_types=["code"], token_endpoint_auth_method="none", persist_credentials=True, + ) + assert response.status_code == 200 + assert json.loads(response.body)["client_id"] == "temporary-client" + assert server.client_id == "temporary-client" + assert server.client_secret == "temporary-secret" + assert server.dcr_issuer == server.issuer + assert server.dcr_server_url == server.url + + +@pytest.mark.asyncio +@pytest.mark.parametrize("config_store", [False, True]) +async def test_registration_does_not_overwrite_credentials_after_failed_identity_read(monkeypatch, config_store): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints, db + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy import utils + + server = _dcr_redirect_test_server(None) + monkeypatch.setattr(utils, "get_prisma_client_or_throw", lambda _: MagicMock()) + monkeypatch.setattr(db, "get_mcp_server", AsyncMock(return_value=None, side_effect=None if config_store else RuntimeError("unavailable"))) + monkeypatch.setattr(db, "get_mcp_server_oauth_client_credentials", AsyncMock(side_effect=RuntimeError("unavailable"))) + monkeypatch.setattr(global_mcp_server_manager, "is_config_declared_server", lambda _: config_store) + update = AsyncMock() + upsert = AsyncMock() + monkeypatch.setattr(db, "update_mcp_server", update) + monkeypatch.setattr(db, "upsert_mcp_server_oauth_client_credentials", upsert) + result = await endpoints._persist_dcr_client_registration(server, {"client_id": "new-client"}, "https://gateway.example/callback") + assert result == "failed" + assert server.client_id is None + update.assert_not_awaited() + upsert.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("winner_available", [False, True]) +async def test_registration_losing_conditional_write_reuses_only_a_matching_winner(monkeypatch, winner_available): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints, db + from litellm.proxy._types import LiteLLM_MCPServerTable + from litellm.proxy import utils + + server = _dcr_redirect_test_server(None).model_copy(update={"url": "https://upstream.example/mcp"}) + row = LiteLLM_MCPServerTable.model_validate(server.model_dump(exclude_none=True)) + winner = row.model_copy(update={"credentials": { + "client_id": "winner-client", "redirect_uris": ["https://gateway.example/callback"], + "dcr_server_url": server.url, + }}) if winner_available else row + monkeypatch.setattr(utils, "get_prisma_client_or_throw", lambda _: MagicMock()) + read = AsyncMock(side_effect=[row, winner]) + monkeypatch.setattr(db, "get_mcp_server", read) + monkeypatch.setattr(endpoints, "_refresh_persisted_dcr_server", AsyncMock()) + update = AsyncMock(return_value=None) + monkeypatch.setattr(db, "update_mcp_server", update) + result = await endpoints._persist_dcr_client_registration(server, {"client_id": "losing-client"}, "https://gateway.example/callback") + assert result == ("reused" if winner_available else "failed") + assert update.await_args.kwargs["expected_updated_at"] == row.updated_at + assert server.client_id == ("winner-client" if winner_available else None) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py index 7a4ad76d86d..f8e1de677b0 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -949,6 +949,34 @@ async def test_get_user_env_vars_returns_empty_for_missing_row(): assert await get_user_env_vars(prisma, "alice", "srv-1") == {} +@pytest.mark.asyncio +async def test_get_user_env_vars_stringifies_non_string_json_values(env_vars_salt_key): + from types import SimpleNamespace + + from litellm.proxy._experimental.mcp_server.db import get_user_env_vars + + row: Final = SimpleNamespace(values_b64=_encrypted_user_env_blob({"PORT": 8080, "DEBUG": True, "EMPTY": None})) + + assert await get_user_env_vars(_mock_env_vars_prisma(row=row), "alice", "srv-1") == { + "PORT": "8080", + "DEBUG": "True", + "EMPTY": "None", + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stored_json", ["[]", '[{"TOKEN": "t"}]', '"TOKEN"', "5", "null", "true", "not json"]) +async def test_get_user_env_vars_treats_a_blob_that_is_not_a_json_object_as_unset(env_vars_salt_key, stored_json): + from types import SimpleNamespace + + from litellm.proxy._experimental.mcp_server.db import get_user_env_vars + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + + row: Final = SimpleNamespace(values_b64=encrypt_value_helper(stored_json)) + + assert await get_user_env_vars(_mock_env_vars_prisma(row=row), "alice", "srv-1") == {} + + @pytest.mark.asyncio async def test_decode_user_env_vars_warns_when_undecryptable(env_vars_salt_key, monkeypatch): """A stored blob encrypted under a previous salt key must surface a warning diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py index a3c52dc16b7..122b22f6273 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py @@ -17,6 +17,7 @@ from fastapi import HTTPException from litellm.proxy._experimental.mcp_server.db import ( create_mcp_server, + decrypt_credentials, set_mcp_server_pinned_tools, update_mcp_server, ) @@ -1181,3 +1182,186 @@ async def test_protocol_update_preserves_missing_server_without_writing(clear_al result = await update_mcp_server(prisma, payload, "admin") assert result is None table.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("changed,previous_issuer", [ + ({"url": "https://new.example/mcp"}, None), + ({"issuer": "https://new.example"}, "https://old.example"), + ({"url": "https://new.example/mcp", "auth_type": "api_key"}, "https://old.example"), +]) +async def test_upstream_identity_edit_drops_previous_oauth_client(changed, previous_issuer): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer=previous_issuer, + credentials=json.dumps({"client_id": "old-client", "client_secret": "old-secret", "upstream_resource": "api://resource"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + await update_mcp_server(prisma, UpdateMCPServerRequest(server_id="test-server", **changed), "test-user") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + assert "credentials" in written + credentials = json.loads(written["credentials"]) + assert "client_id" not in credentials + assert "client_secret" not in credentials + assert credentials["upstream_resource"] == "api://resource" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("issuer", ["https://old.example/", "https://OLD.example:443"]) +async def test_distinct_issuer_identifier_edit_discards_registered_client(issuer): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer="https://old.example", + credentials=json.dumps({"client_id": "old-client", "dcr_issuer": "https://old.example"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + await update_mcp_server(prisma, UpdateMCPServerRequest(server_id="test-server", issuer=issuer), "test-user") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + assert "client_id" not in json.loads(written["credentials"]) + assert written["issuer"] == issuer + + +@pytest.mark.asyncio +@pytest.mark.parametrize("previous_issuer,binding", [("https://idp.example", None), (None, "https://idp.example")]) +@pytest.mark.parametrize("submitted_tokens", [None, {"access_token": "old-token", "refresh_token": "old-refresh", "expires_in": 3600}, {"access_token": "fresh-token", "refresh_token": "fresh-refresh", "expires_in": 3600}]) +async def test_url_edit_preserves_client_bound_to_previous_known_issuer_without_old_tokens(previous_issuer, binding, submitted_tokens): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer=previous_issuer, + credentials=json.dumps({"client_id": "static-client", "client_secret": "static-secret", "dcr_issuer": binding, + "access_token": "old-token", "refresh_token": "old-refresh", "expires_in": 3600, + "token_endpoint_auth_method": "client_secret_basic"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + await update_mcp_server(prisma, UpdateMCPServerRequest(server_id="test-server", url=existing.url + "?v=2", **({"credentials": submitted_tokens} if submitted_tokens is not None else {})), "test-user") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + credentials = json.loads(written["credentials"]) + assert credentials["client_id"] == "static-client" + assert credentials["client_secret"] == "static-secret" + assert credentials["dcr_issuer"] == "https://idp.example" + assert credentials["dcr_server_url"] == existing.url + assert "access_token" not in credentials + assert "refresh_token" not in credentials + assert "expires_in" not in credentials + assert credentials["token_endpoint_auth_method"] == "client_secret_basic" + assert written["issuer"] is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("changed", [0, 1]) +async def test_registration_write_requires_unchanged_server_revision(changed): + from datetime import datetime, timezone + from litellm.proxy._experimental.mcp_server.db import _update_mcp_server_row + + prisma = _mock_prisma() + table = prisma.db.litellm_mcpservertable + revision = datetime.now(timezone.utc) + table.update_many.return_value = changed + result = await _update_mcp_server_row( + prisma, server_id="test-server", data_dict={"credentials": "{}"}, expected_updated_at=revision, + ) + table.update_many.assert_awaited_once_with( + where={"server_id": "test-server", "updated_at": revision}, data={"credentials": "{}"}, + ) + table.update.assert_not_awaited() + assert (result is not None) is bool(changed) + if changed: + assert result.server_id == "test-server" + else: + table.find_unique.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["create", "rotate"]) +async def test_explicit_oauth_client_write_binds_to_server_identity(operation): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer="https://idp.example", + credentials=json.dumps({"client_id": "old-client"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + if operation == "create": + await create_mcp_server(prisma, NewMCPServerRequest( + server_name="bound-client", transport="http", auth_type="oauth2", + url=existing.url, issuer=existing.issuer, credentials={"client_id": "new-client"}, + ), "admin") + written = prisma.db.litellm_mcpservertable.create.call_args.kwargs["data"] + else: + await update_mcp_server(prisma, UpdateMCPServerRequest( + server_id=existing.server_id, credentials={"client_id": "new-client"}, + ), "admin") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + credentials = json.loads(written["credentials"]) + assert credentials["dcr_issuer"] == existing.issuer + assert credentials["dcr_server_url"] == existing.url + + +@pytest.mark.asyncio +async def test_unrelated_edit_does_not_backfill_legacy_oauth_binding(): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer="https://idp.example", + credentials=json.dumps({"client_id": "legacy-client"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + await update_mcp_server(prisma, UpdateMCPServerRequest( + server_id=existing.server_id, credentials={"scopes": ["tools.read"]}, + ), "admin") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + credentials = json.loads(written["credentials"]) + assert credentials["client_id"] == "legacy-client" + assert "dcr_issuer" not in credentials + assert "dcr_server_url" not in credentials + + +@pytest.mark.asyncio +async def test_issuer_edit_does_not_rebind_resubmitted_saved_client(): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + credentials=json.dumps({"client_id": "saved-client", "client_secret": "saved-secret"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + await update_mcp_server(prisma, UpdateMCPServerRequest( + server_id=existing.server_id, url="https://new.example/mcp", issuer="https://new.example", + credentials={"client_id": "saved-client", "client_secret": "saved-secret"}, + ), "admin") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + credentials = json.loads(written["credentials"]) + assert "client_id" not in credentials + assert "client_secret" not in credentials + + +@pytest.mark.asyncio +@pytest.mark.parametrize("replacement", [ + {"client_secret": "replacement-secret"}, + {"client_secret": "old-secret", "token_endpoint_auth_method": "client_secret_basic"}, + {"client_secret": None}, + {"dcr_issuer": "https://new.example", "dcr_server_url": "https://new.example/mcp"}, +]) +async def test_issuer_edit_preserves_replacement_with_same_client_id(replacement): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + credentials=json.dumps({"client_id": "shared-client", "client_secret": "old-secret"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + submitted = {"client_id": "shared-client", **replacement} + await update_mcp_server(prisma, UpdateMCPServerRequest( + server_id=existing.server_id, url="https://new.example/mcp", issuer="https://new.example", + credentials=submitted, + ), "admin") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + credentials = decrypt_credentials(json.loads(written["credentials"])) + assert credentials["client_id"] == "shared-client" + assert credentials["dcr_issuer"] == "https://new.example" + assert credentials["dcr_server_url"] == "https://new.example/mcp" + assert credentials.get("client_secret") == replacement.get("client_secret") + assert credentials.get("token_endpoint_auth_method") == replacement.get("token_endpoint_auth_method") 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 cb017afbea5..123a5505953 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 @@ -7197,7 +7197,7 @@ class TestMCPServerManager: manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server" manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server" manager._create_prefixed_tools(listed_tools, server) - manager._record_listed_tools(server, listed_tools, caller) + manager.record_listed_tools(server, listed_tools, caller) mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) @@ -7281,8 +7281,8 @@ class TestMCPServerManager: def test_get_listed_tool_resolves_the_bare_name_from_the_latest_listing(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._record_listed_tools(server, [MCPTool(name="echo", description="v1", inputSchema={})], None) - manager._record_listed_tools(server, [MCPTool(name="echo", description="v2", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="v1", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="v2", inputSchema={})], None) latest = manager.get_listed_tool(server, "echo") assert latest is not None and latest.description == "v2" @@ -7294,7 +7294,7 @@ class TestMCPServerManager: manager = MCPServerManager() server = MCPServer(server_id="srv-id", name="srv", alias="srv", transport=MCPTransport.http, url="http://srv") caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice")) - manager._record_listed_tools( + manager.record_listed_tools( server, [ MCPTool(name="foo", description="Fetches foo records", inputSchema={"type": "object"}), @@ -7355,8 +7355,8 @@ class TestMCPServerManager: manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") other = MCPServer(server_id="other", name="other", transport=MCPTransport.http, url="http://other") - manager._record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None) - manager._record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None) + manager.record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None) manager._invalidate_server_definition_caches(server.server_id) @@ -7418,7 +7418,7 @@ class TestMCPServerManager: caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="lister")) async def register_while_a_listing_records(server: MCPServer, *, initialize_mapping: bool = True) -> None: - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="search", description="pre-save", inputSchema={})], caller, @@ -7463,7 +7463,7 @@ class TestMCPServerManager: return [Prompt(name="greet")] async def register_while_discovery_fills(server: MCPServer, *, initialize_mapping: bool = True) -> None: - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="search", description="pre-save", inputSchema={})], caller, @@ -7497,7 +7497,7 @@ class TestMCPServerManager: async def test_user_oauth_refresh_keeps_listed_tools(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None) await manager.invalidate_user_oauth_token_cache("alice", server.server_id) @@ -7517,12 +7517,12 @@ class TestMCPServerManager: bob = UserAPIKeyAuth(user_id="bob", token="hashed-bob") alice_schema = {"type": "object", "properties": {"path": {"type": "string"}}} bob_schema = {"type": "object", "properties": {"path": {"type": "string"}, "site": {"type": "string"}}} - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="read", description="alice view", inputSchema=alice_schema)], ListedToolsCaller(user_api_key_auth=alice), ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="read", description="bob view", inputSchema=bob_schema)], ListedToolsCaller(user_api_key_auth=bob), @@ -7539,7 +7539,7 @@ class TestMCPServerManager: assert manager.get_listed_tool(server, "read", carol) is None shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared") - manager._record_listed_tools( + manager.record_listed_tools( shared, [MCPTool(name="echo", description="everyone", inputSchema={})], ListedToolsCaller(user_api_key_auth=alice), @@ -7595,8 +7595,8 @@ class TestMCPServerManager: server = MCPServer( **{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs} ) - manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog A", inputSchema={})], caller_a) - manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog B", inputSchema={})], caller_b) + manager.record_listed_tools(server, [MCPTool(name="turn", description="Catalog A", inputSchema={})], caller_a) + manager.record_listed_tools(server, [MCPTool(name="turn", description="Catalog B", inputSchema={})], caller_b) for_a = manager.get_listed_tool(server, "turn", caller_a) for_b = manager.get_listed_tool(server, "turn", caller_b) @@ -7607,7 +7607,7 @@ class TestMCPServerManager: def test_shared_server_ignores_headers_it_never_forwards(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="turn", description="everyone", inputSchema={})], ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}), @@ -7840,7 +7840,7 @@ class TestMCPServerManager: "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", return_value=signer, ): - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="turn", description="alice view", inputSchema={})], alice ) assert manager.get_listed_tool(server, "turn", bob) is None @@ -7860,7 +7860,7 @@ class TestMCPServerManager: "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", return_value=MagicMock(), ): - manager._record_listed_tools(server, [MCPTool(name="turn", description="slot a", inputSchema={})], alice) + manager.record_listed_tools(server, [MCPTool(name="turn", description="slot a", inputSchema={})], alice) assert manager.get_listed_tool(server, "turn", bob) is None same_key = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) @@ -7878,7 +7878,7 @@ class TestMCPServerManager: team_two: Final = ListedToolsCaller( user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-two") ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], team_one ) @@ -7896,7 +7896,7 @@ class TestMCPServerManager: alice_in_two: Final = ListedToolsCaller( user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-two") ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], alice_in_one ) @@ -7917,7 +7917,7 @@ class TestMCPServerManager: user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"), raw_headers={"authorization": "Bearer jwt-bob"}, ) - manager._record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice) + manager.record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice) assert manager.get_listed_tool(server, "foo", bob) is None listed: Final = manager.get_listed_tool(server, "foo", alice) @@ -7976,7 +7976,7 @@ class TestMCPServerManager: user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"), raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-B"}, ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="lookup", description="Workspace A lookup FLAGWORD", inputSchema={})], caller_a ) @@ -8053,18 +8053,18 @@ class TestMCPServerManager: url="http://srv", auth_type=MCPAuth.oauth2_token_exchange, ) - manager._record_listed_tools(server, [MCPTool(name="read", description="shared", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="read", description="shared", inputSchema={})], None) callers = [ ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id=f"u{i}", api_key=f"k{i}")) for i in range(_LISTED_TOOLS_CALLERS_PER_SERVER + 1) ] for caller in callers: - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})], caller, ) - manager._record_listed_tools(server, [MCPTool(name="read", description="u1 again", inputSchema={})], callers[1]) + manager.record_listed_tools(server, [MCPTool(name="read", description="u1 again", inputSchema={})], callers[1]) assert manager.get_listed_tool(server, "read", callers[0]) is None second = manager.get_listed_tool(server, "read", callers[1]) @@ -9186,6 +9186,8 @@ class TestMCPServerTimestamps: authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", + issuer="https://idp.example.com", + authorization_response_iss_parameter_supported=True, ) blipped = MCPServer( @@ -9200,6 +9202,16 @@ class TestMCPServerTimestamps: assert blipped.authorization_url == "https://idp.example.com/authorize" assert blipped.token_url == "https://idp.example.com/token" assert blipped.registration_url == "https://idp.example.com/register" + assert blipped.issuer == "https://idp.example.com" + assert blipped.authorization_response_iss_parameter_supported is True + fallback = MCPServerManager._merge_discovered_oauth_metadata( + blipped, MCPOAuthMetadata(from_origin_fallback=True), + ) + assert fallback.authorization_response_iss_parameter_supported is True + refreshed = MCPServerManager._merge_discovered_oauth_metadata( + fallback, MCPOAuthMetadata(discovered_issuer="https://idp.example.com"), + ) + assert refreshed.authorization_response_iss_parameter_supported is False same_authorize = MCPServer( server_id="s1", @@ -9316,8 +9328,8 @@ class TestMCPServerTimestamps: from litellm.proxy._experimental.mcp_server.mcp_server_manager import _issuer_matches assert _issuer_matches("https://mcp.slack.com", "https://mcp.slack.com") - assert _issuer_matches("https://MCP.slack.com/", "https://mcp.slack.com") - assert _issuer_matches("https://mcp.slack.com:443", "https://mcp.slack.com") + assert not _issuer_matches("https://MCP.slack.com/", "https://mcp.slack.com") + assert not _issuer_matches("https://mcp.slack.com:443", "https://mcp.slack.com") assert _issuer_matches("https://login.example.com/tenant/v2.0", "https://login.example.com/tenant/v2.0") assert not _issuer_matches("https://attacker.example.com", "https://mcp.slack.com") assert not _issuer_matches("https://login.example.com/other/v2.0", "https://login.example.com/tenant/v2.0") diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 556fe0b2536..3b53d023af3 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -8115,7 +8115,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_listed_entry_and_nothing patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), ): never_listed_tool, never_listed_data = await call() - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="list_pets", description="ADMIN DESC", inputSchema=schema)], ListedToolsCaller(user_api_key_auth=alice), @@ -8156,7 +8156,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl ) manager = mcp_module.global_mcp_server_manager alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice") - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=pinned_schema)], ListedToolsCaller(user_api_key_auth=alice), @@ -8206,12 +8206,12 @@ async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entr manager = mcp_module.global_mcp_server_manager guarded = UserAPIKeyAuth(api_key="sk-guarded", user_id="alice") opted_out = UserAPIKeyAuth(api_key="sk-opted-out", user_id="bob") - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=schema)], ListedToolsCaller(user_api_key_auth=guarded), ) - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="getpetbyid", description="Find a SECRET pet", inputSchema=schema)], ListedToolsCaller(user_api_key_auth=opted_out), @@ -8305,7 +8305,7 @@ async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation ) manager = mcp_module.global_mcp_server_manager alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice") - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="get_pet", description="Fetches pet records. FLAGWORD", inputSchema={"type": "object"})], ListedToolsCaller(user_api_key_auth=alice), diff --git a/tests/unit/proxy/client/cli/test_auth_commands.py b/tests/unit/proxy/client/cli/test_auth_commands.py index 3a7792db1fe..b4a6aaf706e 100644 --- a/tests/unit/proxy/client/cli/test_auth_commands.py +++ b/tests/unit/proxy/client/cli/test_auth_commands.py @@ -7,6 +7,7 @@ from unittest.mock import Mock, patch import pytest +import responses from click.testing import CliRunner from litellm.constants import CLI_JWT_EXPIRATION_HOURS @@ -2044,3 +2045,33 @@ class TestGetStoredApiKeyRefresh: assert captured.out == "" assert captured.err == "Could not renew the key: token request failed with 503: temporarily_unavailable\n" save.assert_not_called() + + +@pytest.mark.parametrize( + "body, expected_detail", + [ + ('{"detail": "Too many CLI login attempts.", "retry": {"after": [30, null]}}', ": Too many CLI login attempts."), + ('{"detail": 5}', ""), + ('{"detail": ""}', ""), + ('{"detail": {"message": "nested"}}', ""), + ("{}", ""), + ('["detail"]', ""), + ('"detail"', ""), + ("7", ""), + ("true", ""), + ("null", ""), + ("gateway", ""), + ("", ""), + ], +) +@responses.activate +def test_login_shows_the_error_detail_only_when_the_proxy_answers_with_a_json_object(body, expected_detail): + responses.post("https://test.example.com/sso/cli/start", body=body, status=429) + + result = CliRunner().invoke(login, obj={"base_url": "https://test.example.com"}) + + assert result.exit_code == 0 + assert result.output == ( + "Authentication failed: Starting CLI login failed: HTTP 429 from https://test.example.com/sso/cli/start" + f"{expected_detail}\n" + ) diff --git a/tests/unit/proxy/client/test_chat.py b/tests/unit/proxy/client/test_chat.py index 8fe1bfcbb2f..9ab7150c288 100644 --- a/tests/unit/proxy/client/test_chat.py +++ b/tests/unit/proxy/client/test_chat.py @@ -7,6 +7,7 @@ import sys import pytest import requests +from pydantic import ValidationError from litellm.proxy.client.chat import ChatClient from litellm.proxy.client.exceptions import UnauthorizedError @@ -256,3 +257,53 @@ def test_completions_stream_gives_up_at_the_timeout_instead_of_hanging(hanging_s next(client.completions_stream(model="gpt-5.4", messages=[{"role": "user", "content": "hi"}])) assert time.monotonic() - started < 10 + + +@pytest.mark.parametrize( + "body, expected", + [ + ( + b'data: {"choices": [{"delta": {"content": "Hel"}}]}\n\n' + b'data: {"choices": [{"delta": {"content": "lo"}}]}\n\n' + b"data: [DONE]\n\n", + [{"choices": [{"delta": {"content": "Hel"}}]}, {"choices": [{"delta": {"content": "lo"}}]}], + ), + ( + b': keep-alive\n\nevent: ping\n\ndata: not json\n\ndata: {"id": "caf\xc3\xa9"}\n\ndata: 7\n\n', + [{"id": "cafรฉ"}, 7], + ), + (b'data: {"id": 1}\n\ndata: [DONE] \n\ndata: {"id": 2}\n\n', [{"id": 1}]), + (b"", []), + ], +) +@responses.activate +def test_completions_stream_yields_parsed_sse_chunks(client, base_url, sample_messages, body, expected): + responses.add(responses.POST, f"{base_url}/chat/completions", body=body, status=200) + + assert list(client.completions_stream(model="gpt-4", messages=sample_messages)) == expected + + +@responses.activate +def test_completions_stream_rejects_bytes_that_are_not_utf8(client, base_url, sample_messages): + responses.add(responses.POST, f"{base_url}/chat/completions", body=b"data: \xff\n\n", status=200) + + with pytest.raises(UnicodeDecodeError): + list(client.completions_stream(model="gpt-4", messages=sample_messages)) + + +@pytest.mark.parametrize("line", ['data: {"id": 1}', 5, memoryview(b'data: {"id": 1}')]) +@responses.activate +def test_completions_stream_rejects_lines_that_are_not_bytes(client, base_url, sample_messages, monkeypatch, line): + responses.add(responses.POST, f"{base_url}/chat/completions", body=b"", status=200) + monkeypatch.setattr(requests.Response, "iter_lines", lambda self: iter([line])) + + with pytest.raises(ValidationError): + list(client.completions_stream(model="gpt-4", messages=sample_messages)) + + +@responses.activate +def test_completions_stream_accepts_bytearray_lines(client, base_url, sample_messages, monkeypatch): + responses.add(responses.POST, f"{base_url}/chat/completions", body=b"", status=200) + monkeypatch.setattr(requests.Response, "iter_lines", lambda self: iter([bytearray(b'data: {"id": 1}')])) + + assert list(client.completions_stream(model="gpt-4", messages=sample_messages)) == [{"id": 1}] diff --git a/tests/unit/proxy/common_utils/test_fips.py b/tests/unit/proxy/common_utils/test_fips.py new file mode 100644 index 00000000000..0e8138856c0 --- /dev/null +++ b/tests/unit/proxy/common_utils/test_fips.py @@ -0,0 +1,106 @@ +import pytest + +from litellm.proxy.common_utils.fips import ( + FipsModeError, + FipsModeOff, + FipsModeOn, + MalformedFipsMode, + ProviderDoesNotEnforceFips, + TlsVerificationDisabled, + enforce_fips_boot_verdict, + fips_boot_verdict, + is_fips_mode, + parse_fips_mode, +) + + +def _verdict( + raw: str | None, + *, + enforcing: bool = True, + ssl_env: str | None = None, + ssl_setting: object = True, +): + return fips_boot_verdict( + raw_fips_mode=raw, + provider_enforces_fips=lambda: enforcing, + ssl_verify_environment=ssl_env, + ssl_verify_setting=ssl_setting, + ) + + +@pytest.mark.parametrize("raw", [None, "false", "0", "no", "off", "", " False "]) +def test_unset_and_false_spellings_leave_fips_mode_off(raw): + assert parse_fips_mode(raw) == FipsModeOff() + assert is_fips_mode({"LITELLM_FIPS_MODE": raw}.get) is False + + +@pytest.mark.parametrize("raw", ["true", "1", "yes", "on", " TRUE "]) +def test_true_spellings_turn_fips_mode_on(raw): + assert parse_fips_mode(raw) == FipsModeOn() + assert is_fips_mode({"LITELLM_FIPS_MODE": raw}.get) is True + + +@pytest.mark.parametrize("raw", ["enforced", "2", "strict", "yes please"]) +def test_anything_else_is_malformed_and_refused_with_the_offending_value(raw): + assert parse_fips_mode(raw) == MalformedFipsMode(value=raw) + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(_verdict(raw, enforcing=False), announce=lambda _: None) + assert f"LITELLM_FIPS_MODE={raw} is not a boolean" in str(refused.value) + assert "true or false" in str(refused.value) + + +def test_off_never_consults_the_provider_or_tls_settings(): + def explode() -> bool: + raise AssertionError("provider probe must not run while FIPS mode is off") + + verdict = fips_boot_verdict( + raw_fips_mode=None, provider_enforces_fips=explode, ssl_verify_environment="false", ssl_verify_setting=False + ) + assert verdict == FipsModeOff() + enforce_fips_boot_verdict(verdict, announce=lambda _: pytest.fail("nothing to announce when off")) + + +def test_on_with_an_enforcing_provider_and_verified_tls_boots(): + verdict = _verdict("true", enforcing=True, ssl_env="true", ssl_setting="/etc/ssl/certs/ca.pem") + assert verdict == FipsModeOn() + enforce_fips_boot_verdict(verdict, announce=lambda _: pytest.fail("nothing to announce when on")) + + +def test_on_with_a_non_enforcing_provider_is_refused_and_names_the_fix(): + announced = [] + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(_verdict("true", enforcing=False), announce=announced.append) + message = str(refused.value) + assert message.startswith("LiteLLM proxy refused to start") + assert "LITELLM_FIPS_MODE is on but this Python does not enforce FIPS" in message + assert "FIPS image" in message + assert announced == [f"\n{message}\n\n"] + + +@pytest.mark.parametrize( + "ssl_env, ssl_setting, sources", + [ + ("false", True, ("SSL_VERIFY",)), + (" FALSE ", True, ("SSL_VERIFY",)), + (None, False, ("litellm_settings.ssl_verify",)), + (None, "False", ("litellm_settings.ssl_verify",)), + ("false", False, ("SSL_VERIFY", "litellm_settings.ssl_verify")), + ], +) +def test_disabled_tls_verification_is_refused_naming_every_source(ssl_env, ssl_setting, sources): + verdict = _verdict("true", enforcing=True, ssl_env=ssl_env, ssl_setting=ssl_setting) + assert verdict == TlsVerificationDisabled(sources=sources) + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(verdict, announce=lambda _: None) + assert "TLS certificate verification is disabled by " + " and ".join(sources) in str(refused.value) + + +@pytest.mark.parametrize("ssl_setting", [True, "true", "/etc/ssl/certs/ca.pem", None, "", "0", "no"]) +def test_verified_or_custom_bundle_tls_settings_are_not_treated_as_disabled(ssl_setting): + assert _verdict("true", enforcing=True, ssl_setting=ssl_setting) == FipsModeOn() + + +def test_disabled_tls_is_reported_before_the_provider_so_operators_see_config_mistakes_first(): + assert _verdict("true", enforcing=False, ssl_env="false") == TlsVerificationDisabled(sources=("SSL_VERIFY",)) + assert _verdict("true", enforcing=False) == ProviderDoesNotEnforceFips() diff --git a/tests/unit/proxy/db/test_object_permission_repository.py b/tests/unit/proxy/db/test_object_permission_repository.py new file mode 100644 index 00000000000..e2d79592915 --- /dev/null +++ b/tests/unit/proxy/db/test_object_permission_repository.py @@ -0,0 +1,73 @@ +from collections.abc import Mapping +from types import SimpleNamespace +from typing import Final + +import pytest + +from litellm.repositories.object_permission_repository import ObjectPermissionRepository + + +class _RecordingPermissionTable: + def __init__(self, stored: Mapping[str, object]) -> None: + self.stored: Final = stored + self.created: Final[list[Mapping[str, object]]] = [] + self.updated: Final[list[tuple[Mapping[str, object], Mapping[str, object]]]] = [] + + async def create(self, data: Mapping[str, object]) -> Mapping[str, object]: + self.created.append(data) + return {**self.stored, **data} + + async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> Mapping[str, object]: + self.updated.append((where, data)) + return {**self.stored, **data} + + +def _repository(table: _RecordingPermissionTable) -> ObjectPermissionRepository: + return ObjectPermissionRepository(SimpleNamespace(db=SimpleNamespace(litellm_objectpermissiontable=table))) + + +@pytest.mark.asyncio +async def test_create_permission_writes_only_the_fields_it_was_given() -> None: + table: Final = _RecordingPermissionTable({"object_permission_id": "perm-1"}) + + permission: Final = await _repository(table).create_permission( + mcp_servers=["server-1"], mcp_tool_permissions={"server-1": ["search"]}, models=[] + ) + + assert table.created == [ + {"mcp_servers": ["server-1"], "mcp_tool_permissions": {"server-1": ["search"]}, "models": []} + ] + assert permission.model_dump(exclude_unset=True) == { + "object_permission_id": "perm-1", + "mcp_servers": ["server-1"], + "mcp_tool_permissions": {"server-1": ["search"]}, + "models": [], + } + + +@pytest.mark.asyncio +async def test_create_permission_without_fields_writes_an_empty_row() -> None: + table: Final = _RecordingPermissionTable({"object_permission_id": "perm-1"}) + + permission: Final = await _repository(table).create_permission() + + assert table.created == [{}] + assert permission.model_dump(exclude_unset=True) == {"object_permission_id": "perm-1"} + + +@pytest.mark.asyncio +async def test_update_permission_changes_only_the_fields_it_was_given() -> None: + table: Final = _RecordingPermissionTable( + {"object_permission_id": "perm-1", "models": ["gpt-4o"], "agents": ["agent-1"]} + ) + + permission: Final = await _repository(table).update_permission("perm-1", models=["gpt-4o-mini"], skills=[]) + + assert table.updated == [({"object_permission_id": "perm-1"}, {"models": ["gpt-4o-mini"], "skills": []})] + assert permission is not None + assert permission.model_dump(exclude_unset=True) == { + "object_permission_id": "perm-1", + "models": ["gpt-4o-mini"], + "agents": ["agent-1"], + "skills": [], + } diff --git a/tests/unit/proxy/fine_tuning_endpoints/test_endpoints.py b/tests/unit/proxy/fine_tuning_endpoints/test_endpoints.py index 7ed1a436cb6..d1b110f288a 100644 --- a/tests/unit/proxy/fine_tuning_endpoints/test_endpoints.py +++ b/tests/unit/proxy/fine_tuning_endpoints/test_endpoints.py @@ -13,9 +13,12 @@ seam stayed untouched, so a guard that raises after the provider call would stil import base64 from contextlib import ExitStack from dataclasses import dataclass +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest +import respx from fastapi import Response @@ -273,3 +276,68 @@ async def test_cancel__unified_job_id_allowed_when_managed_files_required(seams) await _cancel(_unified_job_id()) assert seams.router.acancel_fine_tuning_job.call_count == 1 + + +FINE_TUNING_API_BASE: Final = "https://fine-tuning.test/v1" +PROVIDER_JOBS_PAGE: Final = { + "object": "list", + "data": [ + { + "id": RAW_JOB_ID, + "created_at": 1234567890, + "fine_tuned_model": None, + "finished_at": None, + "hyperparameters": {"n_epochs": 1}, + "model": "gpt-4o-mini", + "object": "fine_tuning.job", + "organization_id": "org-test", + "result_files": [], + "seed": 0, + "status": "running", + "trained_tokens": None, + "training_file": RAW_FILE_ID, + "validation_file": None, + } + ], + "has_more": False, +} + + +async def _list(custom_llm_provider: str | None): + return await endpoints.list_fine_tuning_jobs( + request=FakeRequest(), + fastapi_response=Response(), + custom_llm_provider=custom_llm_provider, + target_model_names=None, + after=None, + limit=None, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + + +@pytest.mark.asyncio +async def test_list__rejected_when_neither_a_provider_nor_a_target_model_is_named(seams): + with pytest.raises(ProxyException) as exc: + await _list(custom_llm_provider=None) + + assert exc.value.code == "400" + assert exc.value.message == "Invalid request, No litellm managed file id or custom_llm_provider provided." + + +@pytest.mark.asyncio +@respx.mock +async def test_list__returns_the_jobs_the_configured_provider_lists(seams): + respx.get(f"{FINE_TUNING_API_BASE}/fine_tuning/jobs").mock( + return_value=httpx.Response(200, json=PROVIDER_JOBS_PAGE) + ) + provider_config: Final = { + "custom_llm_provider": "openai", + "api_key": "sk-provider", + "api_base": FINE_TUNING_API_BASE, + } + + with patch.object(endpoints, "fine_tuning_config", [provider_config]): + page: Final = await _list(custom_llm_provider="openai") + + assert [job.id for job in page.data] == [RAW_JOB_ID] + assert page.has_more is False 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 45ad336368c..0cae9344160 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -5707,6 +5707,41 @@ async def test_unbuffered_end_of_stream_hook_yields_chunks_before_scan(): assert len(chunk_events) == 3 +@pytest.mark.asyncio +async def test_unbuffered_end_of_stream_hook_scans_released_chunks_when_the_client_closes_early(): + guardrail = BedrockGuardrail( + guardrail_name="bedrock-audit-mode", + guardrailIdentifier="test-id", + guardrailVersion="DRAFT", + event_hook=GuardrailEventHooks.post_call, + default_on=True, + streaming_buffer_until_moderated=False, + streaming_end_of_stream_only=True, + ) + scans = [] + + async def record_scan(*args, **kwargs): + scans.append(kwargs["source"]) + return {"action": "NONE", "assessments": [], "outputs": []} + + async def mock_stream(): + yield _chat_chunk("Hello", None) + yield _chat_chunk(" world", None) + yield _chat_chunk("", "stop") + + with patch.object(guardrail, "make_bedrock_api_request", AsyncMock(side_effect=record_scan)): + stream = guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(), + response=mock_stream(), + request_data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}, + ) + first = await stream.__anext__() + await stream.aclose() + + assert first.choices[0].delta.content == "Hello" + assert scans == ["OUTPUT"] + + @pytest.mark.asyncio async def test_buffered_default_hook_scans_before_any_chunk(): guardrail = BedrockGuardrail( diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py new file mode 100644 index 00000000000..ed34c07022e --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_llm_shield_proxy.py @@ -0,0 +1,2003 @@ +import asyncio +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from httpx import Request, Response + +import litellm +from litellm.caching.caching_handler import _PENDING_CACHE_WRITES +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.exceptions import GuardrailRaisedException +from litellm.proxy.guardrails.guardrail_hooks.llm_shield_proxy.llm_shield_proxy import ( + GUARDRAIL_NAME, + LLMShieldProxyGuardrail, +) +from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.llms.openai import ( + FunctionCallArgumentsDeltaEvent, + OutputTextDeltaEvent, + OutputTextDoneEvent, + ResponsesAPIStreamEvents, +) +from litellm.types.utils import ( + Choices, + Delta, + Message, + ModelResponse, + ModelResponseStream, + StreamingChoices, + TextChoices, + TextCompletionResponse, + Usage, +) + + +def _guardrail(**overrides: object) -> LLMShieldProxyGuardrail: + params: dict[str, object] = { + "api_key": "test-key", + "api_base": "http://shield.test", + "guardrail_name": GUARDRAIL_NAME, + "event_hook": "pre_call", + "default_on": True, + } + params.update(overrides) + return LLMShieldProxyGuardrail(**params) + + +def _response(payload: dict, status_code: int = 200) -> Response: + return Response( + status_code=status_code, + json=payload, + request=Request("POST", "http://shield.test/v1/guard/redact"), + ) + + +def _mock_post(guardrail: LLMShieldProxyGuardrail, *payloads: dict) -> AsyncMock: + """Queues one shield response per expected call.""" + mock = AsyncMock(side_effect=[_response(p) for p in payloads]) + guardrail.async_handler.post = mock # type: ignore[method-assign] + return mock + + +def _chunk(content: str | None, finish_reason: str | None = None) -> ModelResponseStream: + return ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(content=content), finish_reason=finish_reason)] + ) + + +async def _drain(generator) -> list: + return [chunk async for chunk in generator] + + +def _tool_chunk(arguments: str, finish_reason: str | None = None) -> ModelResponseStream: + """One streamed fragment of tool call 0's arguments.""" + tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "send", "arguments": arguments}} + return ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[tool_call]), finish_reason=finish_reason)] + ) + + +def _field(holder: object, name: str) -> object: + """Reads a field from a dict or a model; the guardrail emits both shapes.""" + return holder.get(name) if isinstance(holder, dict) else getattr(holder, name) + + +class _FakeShield: + """The three guard endpoints over one fixed vault, placeholder -> original. + + The stream endpoint holds back a trailing `[` that has not closed yet, which is the + behaviour that makes a placeholder split across two chunks come out whole. + """ + + def __init__(self, vault: dict[str, str]) -> None: + self.vault = vault + self.urls: list[str] = [] + + def _restore(self, text: str) -> str: + for placeholder, original in self.vault.items(): + text = text.replace(placeholder, original) + return text + + async def post(self, url: str, headers: dict, json: dict, timeout: float) -> Response: + self.urls.append(url) + if url.endswith("/rehydrate/stream"): + text = self._restore(json["carry"] + json["text"]) + opening = text.rfind("[") + if json["final"] or opening == -1 or "]" in text[opening:]: + return _response({"text": text, "carry": ""}) + return _response({"text": text[:opening], "carry": text[opening:]}) + return _response({"texts": [self._restore(text) for text in json["texts"]]}) + + +def _shielded(vault: dict[str, str]) -> tuple[LLMShieldProxyGuardrail, _FakeShield]: + guardrail = _guardrail(event_hook="post_call") + shield = _FakeShield(vault) + guardrail.async_handler.post = shield.post # type: ignore[method-assign] + return guardrail, shield + + +def _sse(event: dict) -> bytes: + return f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() + + +def _sse_events(frames: list) -> list[dict]: + """Parses emitted SSE output, whatever its chunking, back into event payloads.""" + raw = b"".join(frame.encode() if isinstance(frame, str) else frame for frame in frames).decode() + return [ + json.loads(line[len("data:") :]) + for event in raw.split("\n\n") + for line in event.split("\n") + if line.startswith("data:") + ] + + +def _text_block_stream(*deltas: str) -> list[bytes]: + """An Anthropic /v1/messages stream with one text block made of `deltas`.""" + return [ + _sse({"type": "message_start", "message": {"id": "msg_1", "role": "assistant", "content": []}}), + _sse({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}), + *( + _sse({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": d}}) + for d in deltas + ), + _sse({"type": "content_block_stop", "index": 0}), + _sse({"type": "message_stop"}), + ] + + +async def _restore_stream(guardrail: LLMShieldProxyGuardrail, chunks: list) -> list: + async def stream(): + for chunk in chunks: + yield chunk + + return await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + +def test_llm_shield_guardrail_config(monkeypatch: pytest.MonkeyPatch): + """Should register through init_guardrails_v2 like any other provider.""" + monkeypatch.setattr(litellm, "guardrail_name_config_map", {}) + monkeypatch.setenv("LLM_SHIELD_PROXY_API_KEY", "test-key") + + init_guardrails_v2( + all_guardrails=[ + { + "guardrail_name": "llm_shield_proxy", + "litellm_params": {"guardrail": "llm_shield_proxy", "mode": "pre_call", "default_on": True}, + } + ], + config_file_path="", + ) + + registered = [cb for cb in litellm.callbacks if isinstance(cb, LLMShieldProxyGuardrail)] + assert len(registered) == 1 + assert registered[0].guardrail_name == "llm_shield_proxy" + + +class TestLLMShieldProxyInitialization: + def test_api_base_defaults_to_localhost(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("LLM_SHIELD_PROXY_API_BASE", raising=False) + assert _guardrail(api_base=None).api_base == "http://localhost:8000" + + def test_api_base_reads_environment(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LLM_SHIELD_PROXY_API_BASE", "http://shield.internal:9000") + assert _guardrail(api_base=None).api_base == "http://shield.internal:9000" + + def test_trailing_slash_is_stripped(self): + assert _guardrail(api_base="http://shield.test/").api_base == "http://shield.test" + + def test_both_modes_can_be_enabled_on_one_entry(self): + """Redaction and restoration are two halves of one config entry. + + A deployment that lists only pre_call would redact the request and then hand + the placeholders straight back to the end user. + """ + guardrail = _guardrail(event_hook=["pre_call", "post_call"]) + data: dict = {"messages": []} + + assert guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is True + assert guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.during_call) is False + + +class TestRedaction: + @pytest.mark.asyncio + async def test_string_content_is_redacted(self): + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["Email [EMAIL_1] about it"]}) + + data = {"messages": [{"role": "user", "content": "Email a@b.com about it"}]} + result = await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="completion" + ) + + assert result["messages"][0]["content"] == "Email [EMAIL_1] about it" + + @pytest.mark.asyncio + async def test_multimodal_text_parts_are_redacted(self): + """The list content shape is a historical bypass; text parts must be covered.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["call [PHONE_1]"]}) + + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "call 555-0100"}, + {"type": "image_url", "image_url": {"url": "http://x/y.png"}}, + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["content"][0]["text"] == "call [PHONE_1]" + assert data["messages"][0]["content"][1]["image_url"]["url"] == "http://x/y.png" + + @pytest.mark.asyncio + async def test_request_without_text_is_untouched(self): + """No text to redact means no call to LLM Shield Proxy. + + This deliberately uses a request with no caller text at all. An earlier + version used a Responses-API `input`, which asserted the very bypass that + let `input` reach the provider unredacted. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail) + data = {"model": "gpt-4o", "temperature": 0.2} + + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + mock.assert_not_called() + + @pytest.mark.asyncio + async def test_session_id_is_reused_across_hooks(self): + """Rehydration can only resolve tokens minted under the same session.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["a@b.com"]}) + + data = {"messages": [{"role": "user", "content": "a@b.com"}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + await guardrail._rehydrate(["[EMAIL_1]"], guardrail._session_id(data)) + + sessions = {call.kwargs["headers"]["X-Session-ID"] for call in mock.call_args_list} + assert len(sessions) == 1 + + +class TestRequestCoverage: + """Every request shape that carries caller text must be redacted. + + A shape missed here is not a cosmetic gap: the guardrail reports as enabled + while the raw value goes to the provider. + """ + + @pytest.mark.asyncio + async def test_responses_api_string_input_is_redacted(self): + """Measured against a live provider: `input` reached the model unredacted.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["Email [EMAIL_1] the invoice"]}) + + data = {"input": "Email jane.doe@example.com the invoice"} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["Email jane.doe@example.com the invoice"] + assert data["input"] == "Email [EMAIL_1] the invoice" + + @pytest.mark.asyncio + async def test_responses_api_list_input_is_redacted(self): + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = { + "input": [ + {"role": "user", "content": "jane.doe@example.com"}, + {"role": "user", "content": [{"type": "input_text", "text": "555-0100"}]}, + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["content"] == "[EMAIL_1]" + assert data["input"][1]["content"][0]["text"] == "[PHONE_1]" + + @pytest.mark.asyncio + async def test_tool_call_arguments_are_redacted(self): + """Tool arguments carry the values the user asked the model to act on.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}']}) + + data = { + "messages": [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "send", "arguments": '{"email": "jane.doe@example.com"}'}, + } + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}' + + @pytest.mark.asyncio + async def test_responses_api_instructions_are_redacted(self): + """`instructions` is provider-bound text that sits outside `messages`.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["contact [EMAIL_1]"]}) + + data = {"instructions": "contact jane.doe@example.com", "input": ""} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["contact jane.doe@example.com"] + assert data["instructions"] == "contact [EMAIL_1]" + + @pytest.mark.asyncio + async def test_legacy_function_call_arguments_are_redacted(self): + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}']}) + + data = { + "messages": [ + { + "role": "assistant", + "function_call": {"name": "send", "arguments": '{"email": "jane.doe@example.com"}'}, + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["function_call"]["arguments"] == '{"email": "[EMAIL_1]"}' + + @pytest.mark.asyncio + async def test_completions_prompt_is_redacted(self): + """/v1/completions puts its text in a top-level `prompt`, not in messages.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["Email [EMAIL_1]"]}) + + data = {"prompt": "Email jane.doe@example.com"} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["Email jane.doe@example.com"] + assert data["prompt"] == "Email [EMAIL_1]" + + @pytest.mark.asyncio + async def test_completions_prompt_array_is_redacted(self): + """`prompt` also accepts an array, and each entry is provider-bound.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = {"prompt": ["jane.doe@example.com", "555-0100"]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion") + + assert data["prompt"] == ["[EMAIL_1]", "[PHONE_1]"] + + @pytest.mark.asyncio + async def test_responses_function_call_items_are_redacted(self): + """Responses input items hold tool data in `arguments` and `output`.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}', "sent to [EMAIL_1]"]}) + + data = { + "input": [ + {"type": "function_call", "name": "send", "arguments": '{"email": "jane.doe@example.com"}'}, + {"type": "function_call_output", "call_id": "c1", "output": "sent to jane.doe@example.com"}, + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["arguments"] == '{"email": "[EMAIL_1]"}' + assert data["input"][1]["output"] == "sent to [EMAIL_1]" + + @pytest.mark.asyncio + async def test_responses_tool_output_parts_are_redacted(self): + """A function_call_output can carry its result as a list of input_text parts.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["sent to [EMAIL_1]"]}) + + data = { + "input": [ + { + "type": "function_call_output", + "call_id": "c1", + "output": [{"type": "input_text", "text": "sent to jane.doe@example.com"}], + }, + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["sent to jane.doe@example.com"] + assert data["input"][0]["output"][0]["text"] == "sent to [EMAIL_1]" + + @pytest.mark.asyncio + async def test_responses_custom_tool_call_input_is_redacted(self): + """A replayed custom_tool_call carries its payload in `input`, not `arguments`.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["email [EMAIL_1]"]}) + + data = { + "input": [ + {"type": "custom_tool_call", "call_id": "c1", "name": "mail", "input": "email jane.doe@example.com"}, + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["input"] == "email [EMAIL_1]" + assert data["input"][0]["name"] == "mail" + + @pytest.mark.asyncio + async def test_extra_body_overrides_are_redacted(self): + """LiteLLM merges `extra_body` over the request just before sending, so its fields win on the wire.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["safe", "Mail [EMAIL_1]", "[EMAIL_1]"]}) + + data = { + "model": "gpt-4o", + "input": "safe", + "extra_body": { + "input": "alice@example.com", + "messages": [{"role": "user", "content": "Mail alice@example.com"}], + "service_tier": "flex", + }, + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["safe", "Mail alice@example.com", "alice@example.com"] + assert data["extra_body"]["input"] == "[EMAIL_1]" + assert data["extra_body"]["messages"][0]["content"] == "Mail [EMAIL_1]" + assert data["extra_body"]["service_tier"] == "flex" + + def test_extra_body_system_text_is_privileged(self) -> None: + """An application-authored override is no more restorable than the field it replaces.""" + data = {"messages": [{"role": "user", "content": "U"}], "extra_body": {"system": "S", "instructions": "I"}} + + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert [text for text, _ in caller] == ["U"] + assert sorted(text for text, _ in privileged) == ["I", "S"] + + @pytest.mark.asyncio + async def test_anthropic_text_documents_are_redacted(self): + """A document block carries text inline, as `source.data` or as `source.content`.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[NAME_1] notes", "[EMAIL_1]", "Reach [EMAIL_2]"]}) + + data = { + "messages": [ + { + "role": "user", + "content": [ + { + "type": "document", + "title": "Jane Doe notes", + "source": {"type": "text", "media_type": "text/plain", "data": "alice@example.com"}, + }, + { + "type": "document", + "source": {"type": "content", "content": [{"type": "text", "text": "Reach bob@example.com"}]}, + }, + {"type": "document", "source": {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0x"}}, + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages") + + sent = mock.call_args_list[0].kwargs["json"]["texts"] + assert sent == ["Jane Doe notes", "alice@example.com", "Reach bob@example.com"] + blocks = data["messages"][0]["content"] + assert blocks[0]["source"]["data"] == "[EMAIL_1]" + assert blocks[1]["source"]["content"][0]["text"] == "Reach [EMAIL_2]" + assert blocks[2]["source"]["data"] == "JVBERi0x", "binary sources are not text" + + @pytest.mark.asyncio + async def test_responses_code_interpreter_code_is_redacted(self): + """A replayed code_interpreter_call carries the code the model wrote, which the reply side restores.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["send('[EMAIL_1]')"]}) + + data = {"input": [{"type": "code_interpreter_call", "id": "ci_1", "code": "send('jane.doe@example.com')"}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["code"] == "send('[EMAIL_1]')" + + @pytest.mark.asyncio + async def test_anthropic_system_prompt_is_redacted(self): + """/v1/messages carries its system prompt at the top level, not in messages.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["the user is [EMAIL_1]"]}) + + data = {"system": "the user is jane.doe@example.com", "messages": []} + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages" + ) + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["the user is jane.doe@example.com"] + assert data["system"] == "the user is [EMAIL_1]" + + @pytest.mark.asyncio + async def test_anthropic_system_blocks_are_redacted(self): + """`system` also accepts a list of text blocks.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"system": [{"type": "text", "text": "jane.doe@example.com"}], "messages": []} + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="anthropic_messages" + ) + + assert data["system"][0]["text"] == "[EMAIL_1]" + + @pytest.mark.asyncio + async def test_string_array_input_is_redacted(self): + """Embeddings and moderations send `input` as an array of bare strings.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = {"input": ["jane.doe@example.com", "555-0100"]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aembedding") + + assert data["input"] == ["[EMAIL_1]", "[PHONE_1]"] + + @pytest.mark.asyncio + async def test_participant_name_is_redacted(self): + """`name` on a user turn identifies a person.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["hi", "[PERSON_1]"]}) + + data = {"messages": [{"role": "user", "name": "Jane Doe", "content": "hi"}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["hi", "Jane Doe"] + assert data["messages"][0]["name"] == "[PERSON_1]" + + @pytest.mark.asyncio + async def test_tool_function_name_is_left_alone(self): + """On a tool turn the same field is the function name. + + Redacting it would stop the call routing, so this asserts it is never sent + to the shield at all. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["result"]}) + + data = {"messages": [{"role": "tool", "name": "get_weather", "content": "result"}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["name"] == "get_weather" + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["result"] + + @pytest.mark.asyncio + async def test_anthropic_tool_result_content_is_redacted(self): + """A tool_result nests its own content, as a string or as more blocks.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[EMAIL_2]"]}) + + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": "t1", "content": "found jane.doe@example.com"}, + { + "type": "tool_result", + "tool_use_id": "t2", + "content": [{"type": "text", "text": "also bob@example.com"}], + }, + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["content"][0]["content"] == "[EMAIL_1]" + assert data["messages"][0]["content"][1]["content"][0]["text"] == "[EMAIL_2]" + + @pytest.mark.asyncio + async def test_nesting_past_the_bound_blocks_the_request(self): + """Nesting is caller controlled, so the descent has to stop somewhere -- and where + it stops, the request must not go out. + + This test used to assert the opposite: that text past the bound was skipped. That + sent `past-the-bound@example.com` to the provider unredacted while the guardrail + reported as enabled. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail) + + deep: dict = {"type": "tool_result", "content": "past-the-bound@example.com"} + for _ in range(200): + deep = {"type": "tool_result", "content": [deep]} + data = {"messages": [{"role": "user", "content": [{"type": "text", "text": "shallow"}, deep]}]} + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + mock.assert_not_called() + + @pytest.mark.asyncio + async def test_deep_tool_input_blocks_the_request(self): + """A tool_use input past the JSON bound must not be forwarded half-redacted.""" + guardrail = _guardrail() + mock = _mock_post(guardrail) + + deep: dict = {"email": "past-the-bound@example.com"} + for _ in range(100): + deep = {"next": deep} + block = {"type": "tool_use", "id": "t1", "name": "f", "input": deep} + data = {"messages": [{"role": "assistant", "content": [block]}]} + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + mock.assert_not_called() + + @pytest.mark.asyncio + async def test_realistic_nesting_is_redacted_in_full(self): + """The bounds are far past real payloads: a tool input nested inside a tool result, + several JSON levels deep, is redacted whole rather than refused.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + tool_use = { + "type": "tool_use", + "id": "t1", + "name": "f", + "input": {"a": {"b": {"c": {"d": {"to": "x@example.com"}}}}}, + } + data = {"messages": [{"role": "user", "content": [{"type": "tool_result", "content": [tool_use]}]}]} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert tool_use["input"]["a"]["b"]["c"]["d"]["to"] == "[EMAIL_1]" + + @pytest.mark.asyncio + async def test_responses_prompt_object_variables_are_redacted(self): + """A PromptObject's variables are substituted into the prompt provider side. + + The id and version pick which stored prompt to run and have to arrive + unchanged; the variables are caller text. + """ + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"prompt": {"id": "pmpt_123", "version": "2", "variables": {"customer": "jane.doe@example.com"}}} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["prompt"]["variables"]["customer"] == "[EMAIL_1]" + assert data["prompt"]["id"] == "pmpt_123" + assert data["prompt"]["version"] == "2" + + @pytest.mark.asyncio + async def test_responses_typed_prompt_variables_are_redacted(self): + """A variable can be a typed input rather than a string; its `text` is caller text.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = { + "prompt": { + "id": "pmpt_123", + "variables": { + "customer": {"type": "input_text", "text": "jane.doe@example.com"}, + "logo": {"type": "input_image", "image_url": "https://example.com/logo.png"}, + }, + } + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["prompt"]["variables"]["customer"] == {"type": "input_text", "text": "[EMAIL_1]"} + assert data["prompt"]["variables"]["logo"]["image_url"] == "https://example.com/logo.png" + + @pytest.mark.asyncio + async def test_completions_suffix_is_redacted(self): + """LiteLLM forwards the legacy `suffix` to providers that support it.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["signed [EMAIL_1]", "write to [EMAIL_1]"]}) + + data = {"prompt": "write to jane.doe@example.com", "suffix": "signed jane.doe@example.com"} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="atext_completion") + + assert data["suffix"] == "signed [EMAIL_1]" + + @pytest.mark.asyncio + async def test_every_shape_in_one_request_is_redacted(self): + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["a", "b", "c", "d"]}) + + data = { + "messages": [ + {"role": "user", "content": "one"}, + {"role": "user", "content": [{"type": "text", "text": "two"}]}, + { + "role": "assistant", + "tool_calls": [{"function": {"name": "f", "arguments": "three"}}], + }, + ], + "input": "four", + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["one", "two", "three", "four"] + assert data["messages"][0]["content"] == "a" + assert data["messages"][1]["content"][0]["text"] == "b" + assert data["messages"][2]["tool_calls"][0]["function"]["arguments"] == "c" + assert data["input"] == "d" + + @pytest.mark.asyncio + async def test_anthropic_tool_use_input_is_redacted(self): + """A replayed tool_use block carries its arguments as a JSON object, not a string.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = { + "messages": [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "t1", + "name": "send", + "input": {"to": "jane.doe@example.com", "meta": {"phone": "555-0100"}}, + } + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["jane.doe@example.com", "555-0100"] + block = data["messages"][0]["content"][0] + assert block["input"] == {"to": "[EMAIL_1]", "meta": {"phone": "[PHONE_1]"}} + assert block["name"] == "send", "the tool name has to arrive unchanged for the call to route" + + @pytest.mark.asyncio + async def test_responses_reasoning_summary_is_redacted(self): + """A replayed reasoning item quotes the conversation in its summary parts.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["user asked about [EMAIL_1]"]}) + + data = { + "input": [ + { + "type": "reasoning", + "id": "rs_1", + "summary": [{"type": "summary_text", "text": "user asked about jane.doe@example.com"}], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["summary"][0]["text"] == "user asked about [EMAIL_1]" + + def test_tool_schemas_give_up_their_free_text_and_nothing_else(self): + """Every string is collected except what must reach the model verbatim: names, + types, formats, patterns and required lists.""" + data = { + "tools": [ + { + "type": "function", + "function": { + "name": "lookup", + "description": "top", + "parameters": { + "type": "object", + "title": "title", + "properties": { + # A property that is itself named "description". + "description": {"type": "string", "description": "named"}, + "kind": {"type": "string", "enum": ["a", "b"], "const": "a", "description": "enum"}, + "deep": {"type": "array", "items": {"type": "object", "description": "nested"}}, + "to": { + "type": "string", + "format": "email", + "pattern": "^.+@.+$", + "examples": ["example"], + "default": "default", + }, + "choice": {"anyOf": [{"type": "object", "default": {"type": "object-default"}}]}, + # Property names that collide with keywords are subschemas all the same. + "type": {"type": "string", "description": "named-type"}, + }, + "required": ["to"], + "$defs": {"shared": {"description": "defined"}}, + # Keywords nobody listed: scanned by default. + "dependencies": {"mode": {"description": "dependent"}}, + "$comment": "comment", + "x-note": "vendor", + }, + }, + } + ] + } + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert sorted(text for text, _ in caller) == ["a", "a", "b"], "enum and const go to the caller vault" + assert sorted(text for text, _ in privileged) == [ + "comment", + "default", + "defined", + "dependent", + "enum", + "example", + "named", + "named-type", + "nested", + "object-default", + "title", + "top", + "vendor", + ] + + @pytest.mark.asyncio + async def test_enum_values_are_redacted_and_restored_in_the_tool_call(self): + """An enum value holding PII is redacted, and the model's use of the stand-in is + restored in its tool arguments, so the call still carries a value the schema allows.""" + guardrail = _guardrail(event_hook=["pre_call", "post_call"]) + shield = _FakeShield({"[EMAIL_1]": "ops@example.com"}) + redact_mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + data = { + "messages": [], + "tools": [ + { + "type": "function", + "function": { + "name": "notify", + "parameters": {"properties": {"to": {"type": "string", "enum": ["ops@example.com"]}}}, + }, + } + ], + } + + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + assert redact_mock.call_args_list[0].kwargs["json"]["texts"] == ["ops@example.com"] + assert data["tools"][0]["function"]["parameters"]["properties"]["to"]["enum"] == ["[EMAIL_1]"] + + guardrail.async_handler.post = shield.post # type: ignore[method-assign] + call = SimpleNamespace(function=SimpleNamespace(name="notify", arguments='{"to": "[EMAIL_1]"}')) + reply = ModelResponse(choices=[Choices(message=Message(content=None, tool_calls=None))]) + reply.choices[0].message.tool_calls = [call] + restored = await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply) + + assert json.loads(restored.choices[0].message.tool_calls[0].function.arguments) == {"to": "ops@example.com"} + + def test_schema_nesting_past_the_bound_is_refused(self): + schema: dict = {"type": "object", "description": "past-the-bound@example.com"} + for _ in range(100): + schema = {"type": "object", "properties": {"next": schema}} + data = {"tools": [{"type": "function", "function": {"name": "f", "parameters": schema}}]} + + with pytest.raises(Exception, match="schema"): + LLMShieldProxyGuardrail._locate_request_texts(data) + + +class TestRestoration: + @pytest.mark.asyncio + async def test_openai_shape_is_restored(self): + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result.choices[0].message.content == "a@b.com" + + @pytest.mark.asyncio + async def test_responses_api_shape_is_restored(self): + """The Responses API reply carries output items, not choices. + + Measured against a live provider: once the request side was fixed the reply + came back still holding the placeholder, because this shape has no choices + to walk. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = SimpleNamespace(output=[SimpleNamespace(content=[{"type": "output_text", "text": "[EMAIL_1]"}])]) + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result.output[0].content[0]["text"] == "a@b.com" + + @pytest.mark.asyncio + async def test_responses_api_object_blocks_are_restored(self): + """Blocks arrive as objects too, depending on how far the reply is parsed.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + block = SimpleNamespace(text="[EMAIL_1]") + response = SimpleNamespace(output=[SimpleNamespace(content=[block])]) + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result.output[0].content[0].text == "a@b.com" + + @pytest.mark.asyncio + async def test_anthropic_message_shape_is_restored(self): + """The /v1/messages reply is a plain dict with no choices. + + Measured against a live provider: without its own branch the reply went + back to the caller still carrying the placeholder, even though the + request had been redacted correctly. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = { + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "[EMAIL_1]"}], + } + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result["content"][0]["text"] == "a@b.com" + + @pytest.mark.asyncio + async def test_anthropic_non_text_blocks_are_left_alone(self): + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = { + "type": "message", + "role": "assistant", + "content": [ + {"type": "text", "text": "[EMAIL_1]"}, + {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}}, + ], + } + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result["content"][0]["text"] == "a@b.com" + assert result["content"][1] == {"type": "tool_use", "id": "t1", "name": "lookup", "input": {}} + + +class TestVaultIsolation: + """The vault id must never be something a caller can choose. + + The vault holds the plaintext behind every placeholder. If a caller could name + the vault, they could send a placeholder, have the model echo it back, and get + another caller's value restored into their own reply. + """ + + @pytest.mark.asyncio + async def test_caller_supplied_session_id_is_not_used(self): + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = { + "messages": [{"role": "user", "content": "a@b.com"}], + "metadata": {"llm_shield_session_id": "victim-session"}, + "litellm_session_id": "victim-session", + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + used = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] + assert used != "victim-session" + assert data["litellm_metadata"]["llm_shield_session_id"] == used + + @pytest.mark.asyncio + async def test_session_id_is_not_forwarded_to_the_provider(self): + """The vault id is a capability, so it must stay out of provider-visible metadata. + + `metadata` is forwarded upstream on /v1/responses; `litellm_metadata` is not. A + provider holding both the placeholders and the session id could call the shield's + rehydrate endpoint and read back exactly what this guardrail withholds. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"messages": [{"role": "user", "content": "a@b.com"}], "metadata": {}} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + used = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] + assert "llm_shield_session_id" not in data["metadata"] + assert data["litellm_metadata"]["llm_shield_session_id"] == used + + @pytest.mark.asyncio + async def test_restore_ignores_a_foreign_session_id(self): + """A reply is left unrestored rather than resolved against another vault.""" + guardrail = _guardrail(event_hook="post_call") + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}) + + data = {"metadata": {"llm_shield_session_id": "victim-session"}} + response = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=response) + + assert mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] != "victim-session" + + @pytest.mark.asyncio + async def test_each_request_gets_its_own_vault(self): + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["[EMAIL_1]"]}) + + for _ in range(2): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + seen = {call.kwargs["headers"]["X-Session-ID"] for call in mock.call_args_list} + assert len(seen) == 2 + + + @pytest.mark.parametrize( + "data", + [ + pytest.param( + {"messages": [{"role": "system", "content": "S"}, {"role": "user", "content": "U"}]}, + id="system-turn", + ), + pytest.param( + {"messages": [{"role": "developer", "content": "S"}, {"role": "user", "content": "U"}]}, + id="developer-turn", + ), + pytest.param( + {"system": "S", "messages": [{"role": "user", "content": "U"}]}, + id="anthropic-top-level-system", + ), + pytest.param({"instructions": "S", "input": "U"}, id="responses-instructions"), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "tools": [{"type": "function", "function": {"name": "f", "description": "S"}}], + }, + id="chat-tool-description", + ), + pytest.param( + {"input": "U", "tools": [{"type": "function", "name": "f", "description": "S"}]}, + id="responses-tool-description", + ), + pytest.param( + {"messages": [{"role": "user", "content": "U"}], "functions": [{"name": "f", "description": "S"}]}, + id="legacy-function-description", + ), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "tools": [{"name": "f", "input_schema": {"properties": {"to": {"description": "S"}}}}], + }, + id="anthropic-schema-description", + ), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "response_format": { + "type": "json_schema", + "json_schema": {"name": "n", "description": "S", "schema": {"type": "object"}}, + }, + }, + id="chat-response-format", + ), + pytest.param( + { + "input": "U", + "text": { + "format": { + "type": "json_schema", + "name": "n", + "schema": {"properties": {"a": {"description": "S"}}}, + } + }, + }, + id="responses-text-format", + ), + pytest.param({"prediction": {"type": "content", "content": "U"}, "instructions": "S"}, id="prediction"), + pytest.param( + {"prediction": {"type": "content", "content": [{"type": "text", "text": "U"}]}, "instructions": "S"}, + id="prediction-parts", + ), + pytest.param( + { + "messages": [{"role": "user", "content": "U"}], + "web_search_options": {"user_location": {"type": "approximate", "approximate": {"city": "S"}}}, + }, + id="chat-web-search-location", + ), + pytest.param( + { + "input": "U", + "tools": [{"type": "web_search", "user_location": {"type": "approximate", "region": "S"}}], + }, + id="responses-web-search-location", + ), + pytest.param({"messages": [{"role": "user", "content": "U"}], "user": "S"}, id="end-user-id"), + pytest.param({"input": "U", "safety_identifier": "S"}, id="safety-identifier"), + ], + ) + def test_server_authored_text_is_split_from_the_callers(self, data: dict) -> None: + """Every request shape must sort its server-authored spans out of the caller's.""" + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert [text for text, _ in caller] == ["U"] + assert [text for text, _ in privileged] == ["S"] + + @pytest.mark.asyncio + async def test_a_system_prompt_gets_a_vault_of_its_own(self) -> None: + """The reply is restored against the caller's vault, so the two cannot be one.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["[EMAIL_2]"]}) + + data = { + "messages": [ + {"role": "system", "content": "escalate to admin@corp.internal"}, + {"role": "user", "content": "email a@b.com"}, + ] + } + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="completion" + ) + + privileged_id, caller_id = ( + call.kwargs["headers"]["X-Session-ID"] for call in mock.call_args_list + ) + assert privileged_id != caller_id + assert guardrail._session_id(data) == caller_id + + @pytest.mark.asyncio + async def test_the_system_prompt_vault_id_is_never_stored(self) -> None: + """Nothing can restore against the system vault later, because its id is not kept. + + This is what stops a caller from having the model echo a placeholder out of a + system prompt they cannot see and receiving the plaintext behind it. + """ + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["[EMAIL_1]"]}, {"texts": ["[EMAIL_2]"]}) + + data = { + "messages": [ + {"role": "system", "content": "escalate to admin@corp.internal"}, + {"role": "user", "content": "email a@b.com"}, + ] + } + await guardrail.async_pre_call_hook( + user_api_key_dict=None, cache=None, data=data, call_type="completion" + ) + + privileged_id = mock.call_args_list[0].kwargs["headers"]["X-Session-ID"] + assert privileged_id not in json.dumps(data, default=str) + + def test_responses_system_and_developer_items_are_privileged(self) -> None: + """Responses `input` carries system and developer turns as items, like Chat messages. + + In the caller's vault, a caller could have the model echo a placeholder out of a + system message they cannot see and receive the plaintext behind it. + """ + data = { + "input": [ + {"role": "system", "content": "S"}, + {"type": "message", "role": "developer", "content": [{"type": "input_text", "text": "D"}]}, + {"role": "user", "content": "U"}, + ] + } + caller, privileged = LLMShieldProxyGuardrail._locate_request_texts(data) + + assert [text for text, _ in caller] == ["U"] + assert [text for text, _ in privileged] == ["S", "D"] + + +class TestFailClosed: + @pytest.mark.asyncio + async def test_unreachable_shield_blocks_the_request(self): + """Failing open would send the PII upstream, defeating the guardrail.""" + guardrail = _guardrail() + guardrail.async_handler.post = AsyncMock(side_effect=ConnectionError("refused")) + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + @pytest.mark.asyncio + async def test_error_status_blocks_the_request(self): + guardrail = _guardrail() + guardrail.async_handler.post = AsyncMock(return_value=_response({"error": "nope"}, status_code=500)) + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + @pytest.mark.asyncio + async def test_short_payload_blocks_the_request(self): + """A response that loses an entry would silently misalign the write-back.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": []}) + + with pytest.raises(GuardrailRaisedException): + await guardrail.async_pre_call_hook( + user_api_key_dict=None, + cache=None, + data={"messages": [{"role": "user", "content": "a@b.com"}]}, + call_type="completion", + ) + + +class TestStreamingRehydration: + @pytest.mark.asyncio + async def test_split_placeholder_is_not_emitted_in_fragments(self): + """The window holds back a partial placeholder and releases it once complete.""" + guardrail = _guardrail(event_hook="post_call") + # Shield holds "[EMAIL" back, then releases the restored value. + _mock_post( + guardrail, + {"text": "Email ", "carry": "[EMAIL"}, + {"text": "a@b.com about it", "carry": ""}, + ) + + async def stream(): + yield _chunk("Email [EMAIL") + yield _chunk("_1] about it", finish_reason="stop") + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + emitted = [c.choices[0].delta.content for c in chunks] + assert emitted == ["Email ", "a@b.com about it"] + # No fragment of the placeholder ever reached the client. + assert not any("[EMAIL" in (text or "") for text in emitted) + + @pytest.mark.asyncio + async def test_carry_is_returned_to_the_next_call(self): + guardrail = _guardrail(event_hook="post_call") + mock = _mock_post( + guardrail, + {"text": "", "carry": "hold"}, + {"text": "held-and-more", "carry": ""}, + ) + + async def stream(): + yield _chunk("hold") + yield _chunk("-and-more", finish_reason="stop") + + await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert mock.call_args_list[0].kwargs["json"]["carry"] == "" + assert mock.call_args_list[1].kwargs["json"]["carry"] == "hold" + assert mock.call_args_list[1].kwargs["json"]["final"] is True + + @pytest.mark.asyncio + async def test_every_choice_is_restored(self): + """With n>1 a later choice must not be handed back still holding a placeholder.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": "first@example.com", "carry": ""}, + {"text": "second@example.com", "carry": ""}, + ) + + async def stream(): + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="[EMAIL_1]"), finish_reason="stop"), + StreamingChoices(index=1, delta=Delta(content="[EMAIL_2]"), finish_reason="stop"), + ] + ) + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + restored = [choice.delta.content for choice in chunks[0].choices] + assert restored == ["first@example.com", "second@example.com"] + + @pytest.mark.asyncio + async def test_choice_windows_do_not_cross_contaminate(self): + """Each choice is its own token stream, so each carries its own window. + + One shared window would send the characters held back for choice 0 up + against choice 1's next delta and splice the two streams together. + """ + guardrail = _guardrail(event_hook="post_call") + mock = _mock_post( + guardrail, + {"text": "", "carry": "A-held"}, + {"text": "", "carry": "B-held"}, + {"text": "a-done", "carry": ""}, + {"text": "b-done", "carry": ""}, + ) + + async def stream(): + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="a1")), + StreamingChoices(index=1, delta=Delta(content="b1")), + ] + ) + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="a2"), finish_reason="stop"), + StreamingChoices(index=1, delta=Delta(content="b2"), finish_reason="stop"), + ] + ) + + await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + sent = [call.kwargs["json"] for call in mock.call_args_list] + assert sent[2]["carry"] == "A-held", "choice 0 must get its own window back" + assert sent[3]["carry"] == "B-held", "choice 1 must get its own window back" + + @pytest.mark.asyncio + async def test_a_choice_missing_from_the_last_chunk_still_flushes(self): + """Held text must not be dropped because its choice ended earlier. + + Choice 1 finishes and stops appearing, then the stream ends without a + finish_reason for choice 0. Flushing only the terminal chunk's choices would + discard whatever choice 1 was still holding and truncate its answer. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": "", "carry": "held-0"}, + {"text": "", "carry": "held-1"}, + {"text": "zero-done", "carry": ""}, + {"text": "one-done", "carry": ""}, + ) + + async def stream(): + yield ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="a")), + StreamingChoices(index=1, delta=Delta(content="b")), + ] + ) + yield ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(content=None))]) + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + flushed = { + choice.index: choice.delta.content for chunk in chunks for choice in chunk.choices if choice.delta.content + } + assert flushed.get(1) == "one-done", "choice 1's held text was dropped" + assert flushed.get(0) == "zero-done" + + @pytest.mark.asyncio + async def test_held_tool_arguments_land_in_the_finishing_chunk(self): + """A client parses tool arguments on finish_reason, so the flush must ride that chunk. + + The finishing chunk also carries its own fragment for the same tool call. That + entry has to survive, with the held text appended after it as a continuation. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": '{"to": "', "carry": "[EMA"}, + {"text": "", "carry": '[EMAIL_1]"}'}, + {"text": 'a@example.com"}', "carry": ""}, + ) + + async def stream(): + yield _tool_chunk('{"to": "[EMA') + yield _tool_chunk('IL_1]"}', finish_reason="tool_calls") + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert len(chunks) == 2, "the flush must not arrive after the finish_reason chunk" + final_calls = chunks[1].choices[0].delta.tool_calls + assert len(final_calls) == 2, "the finishing chunk's own fragment was dropped" + assert _field(final_calls[1], "index") == 0 + arguments = "".join( + _field(_field(call, "function"), "arguments") or "" + for chunk in chunks + for call in chunk.choices[0].delta.tool_calls + ) + assert json.loads(arguments) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_held_tool_arguments_flush_when_the_stream_ends_unfinished(self): + """No finish_reason at all: a trailing chunk carries the held arguments alone.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": '{"to": "', "carry": "[EMA"}, + {"text": 'a@example.com"}', "carry": ""}, + ) + + async def stream(): + yield _tool_chunk('{"to": "[EMA') + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert len(chunks) == 2 + trailing = chunks[1].choices[0].delta.tool_calls + assert trailing == [{"index": 0, "function": {"arguments": 'a@example.com"}'}}], ( + "the copied chunk's own fragment was already delivered and must not repeat" + ) + + @pytest.mark.asyncio + async def test_chunks_are_forwarded_as_they_arrive(self): + """Restoration must not buffer the stream into a single terminal chunk.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post( + guardrail, + {"text": "one ", "carry": ""}, + {"text": "two ", "carry": ""}, + {"text": "three", "carry": ""}, + ) + + async def stream(): + yield _chunk("one ") + yield _chunk("two ") + yield _chunk("three", finish_reason="stop") + + chunks = await _drain( + guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=None, response=stream(), request_data={"messages": []} + ) + ) + + assert len(chunks) == 3 + assert [c.choices[0].delta.content for c in chunks] == ["one ", "two ", "three"] + + +class TestApplyGuardrailToolCalls: + """The unified entry point the UI's Test button and the translation handlers use.""" + + @pytest.mark.asyncio + async def test_response_tool_call_arguments_are_rehydrated(self): + """Regression: this path deep-copied tool calls with `copy` never imported. + + 47 tests passed with a guaranteed NameError here, because every tool-call test + covered the request side and this is the only path that reaches the copy. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["hi", '{"email": "a@b.com"}']}) + + data = {"litellm_metadata": {"llm_shield_session_id": "shield-abc"}} + inputs = { + "texts": ["hi"], + "tool_calls": [{"function": {"name": "send", "arguments": '{"email": "[EMAIL_1]"}'}}], + } + + merged = await guardrail.apply_guardrail(inputs=inputs, request_data=data, input_type="response") + + assert merged["tool_calls"][0]["function"]["arguments"] == '{"email": "a@b.com"}' + assert inputs["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}' + + +class TestAnthropicStreamRestoration: + """/v1/messages streams reach the hook as raw SSE frames, with no `choices` to walk.""" + + VAULT = {"[EMAIL_1]": "a@example.com"} + + @pytest.mark.asyncio + async def test_split_placeholder_is_restored_and_never_fragmented(self): + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, _text_block_stream("Mail [EMA", "IL_1] now")) + + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com now" + assert not any("[EMA" in d for d in deltas), "a placeholder fragment reached the client" + + @pytest.mark.asyncio + async def test_held_text_lands_before_its_block_stops(self): + """A trailing `[` that never became a placeholder is still part of the answer.""" + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, _text_block_stream("Mail [EMAIL_1], x = a[")) + + types = [e["type"] for e in _sse_events(out)] + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com, x = a[" + assert types.index("content_block_stop") > max(i for i, t in enumerate(types) if t == "content_block_delta") + + @pytest.mark.asyncio + async def test_events_split_across_network_chunks_are_restored(self): + """A chunk can end mid-event; the frame is parsed once it is whole.""" + guardrail, _ = _shielded(self.VAULT) + raw = b"".join(_text_block_stream("Mail [EMA", "IL_1] now")) + + out = await _restore_stream(guardrail, [raw[i : i + 7] for i in range(0, len(raw), 7)]) + + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com now" + + @pytest.mark.asyncio + async def test_str_frames_stay_str(self): + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, [frame.decode() for frame in _text_block_stream("[EMAIL_1]")]) + + assert all(isinstance(frame, str) for frame in out) + assert "a@example.com" in "".join(out) + + @pytest.mark.asyncio + async def test_tool_input_json_is_restored(self): + guardrail, _ = _shielded(self.VAULT) + frames = [ + _sse({"type": "content_block_start", "index": 1, "content_block": {"type": "tool_use", "input": {}}}), + *( + _sse( + { + "type": "content_block_delta", + "index": 1, + "delta": {"type": "input_json_delta", "partial_json": p}, + } + ) + for p in ('{"to": "[EMAI', 'L_1]"}') + ), + _sse({"type": "content_block_stop", "index": 1}), + ] + + out = await _restore_stream(guardrail, frames) + + partial = "".join(e["delta"]["partial_json"] for e in _sse_events(out) if e["type"] == "content_block_delta") + assert json.loads(partial) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_signed_thinking_and_foreign_frames_pass_through_byte_for_byte(self): + """Rewriting a signed thinking block breaks it; other frames are not ours to touch.""" + guardrail, shield = _shielded(self.VAULT) + thinking = {"type": "thinking_delta", "thinking": "[EMAIL_1]"} + frames = [ + _sse({"type": "content_block_delta", "index": 0, "delta": thinking}), + b'data: {"candidates": [{"content": {"parts": [{"text": "[EMAIL_1]"}]}}]}\n\n', + b"data: not json\n\n", + ] + + out = await _restore_stream(guardrail, frames) + + assert b"".join(out) == b"".join(frames) + assert shield.urls == [] + + @pytest.mark.asyncio + async def test_a_raw_stream_that_is_not_sse_is_never_buffered(self): + """Without event boundaries to wait for, buffering would hold the whole reply.""" + guardrail, _ = _shielded(self.VAULT) + chunks = [b'[{"candidates": []}', b', {"candidates": []}]'] + + out = await _restore_stream(guardrail, chunks) + + assert out == chunks + + @pytest.mark.asyncio + @pytest.mark.parametrize("cut", [1, 3, 5, 6]) + async def test_a_field_name_split_by_the_first_chunk_still_reads_as_sse(self, cut: int): + """`b"eve"` then `b"nt: ..."` is still SSE; deciding on the first chunk alone + would pass the whole stream through with its placeholders.""" + guardrail, _ = _shielded(self.VAULT) + raw = b"".join(_text_block_stream("Mail [EMAIL_1]")) + + out = await _restore_stream(guardrail, [raw[:cut], raw[cut:]]) + + deltas = [e["delta"]["text"] for e in _sse_events(out) if e["type"] == "content_block_delta"] + assert "".join(deltas) == "Mail a@example.com" + + +class TestResponsesStreamRestoration: + """/v1/responses streams are typed events, with no `choices` to walk.""" + + VAULT = {"[EMAIL_1]": "a@example.com"} + + @staticmethod + def _text_delta(delta: str, sequence_number: int, content_index: int = 0) -> OutputTextDeltaEvent: + return OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=content_index, + delta=delta, + sequence_number=sequence_number, + ) + + @pytest.mark.asyncio + async def test_deltas_and_done_text_are_restored(self): + guardrail, _ = _shielded(self.VAULT) + done = OutputTextDoneEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE, + item_id="msg_1", + output_index=0, + content_index=0, + text="Mail [EMAIL_1] x[", + ) + + out = await _restore_stream( + guardrail, [self._text_delta("Mail [EMA", 1), self._text_delta("IL_1] x[", 2), done] + ) + + deltas = [e.delta for e in out if isinstance(e, OutputTextDeltaEvent)] + assert not any("[EMA" in d for d in deltas), "a placeholder fragment reached the client" + assert "".join(deltas) == "Mail a@example.com x[" + assert out[-1].text == "Mail a@example.com x[", "the done event repeats the full, restored text" + assert isinstance(out[-2], OutputTextDeltaEvent), "held text must land before the done event" + + @pytest.mark.asyncio + async def test_function_call_arguments_are_restored(self): + guardrail, _ = _shielded(self.VAULT) + events = [ + FunctionCallArgumentsDeltaEvent( + type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, + item_id="fc_1", + output_index=1, + delta=part, + ) + for part in ('{"to": "[EMAI', 'L_1]"}') + ] + + out = await _restore_stream(guardrail, events) + + assert json.loads("".join(e.delta for e in out)) == {"to": "a@example.com"} + + @pytest.mark.asyncio + async def test_a_truncated_stream_still_flushes(self): + """No done event at all: whatever the window holds goes out at the end.""" + guardrail, _ = _shielded(self.VAULT) + + out = await _restore_stream(guardrail, [self._text_delta("see [EMAIL_1] a[", 1)]) + + assert "".join(e.delta for e in out) == "see a@example.com a[" + + @pytest.mark.asyncio + async def test_completed_response_is_restored(self): + """The terminal event repeats the whole reply, and clients read it as the answer.""" + guardrail, _ = _shielded(self.VAULT) + block = {"type": "output_text", "text": "Mail [EMAIL_1]"} + call = SimpleNamespace(type="function_call", arguments='{"to": "[EMAIL_1]"}') + completed = SimpleNamespace( + type="response.completed", + response=SimpleNamespace(output=[SimpleNamespace(content=[block]), call]), + ) + + (out,) = await _restore_stream(guardrail, [completed]) + + restored_block, restored_call = out.response.output[0].content[0], out.response.output[1] + assert restored_block["text"] == "Mail a@example.com" + assert restored_call.arguments == '{"to": "a@example.com"}' + + @pytest.mark.asyncio + async def test_streams_on_different_parts_do_not_share_a_window(self): + guardrail, _ = _shielded(self.VAULT) + events = [ + self._text_delta("one [EMA", 1), + self._text_delta("two", 2, content_index=1), + self._text_delta("IL_1]", 3), + ] + + out = await _restore_stream(guardrail, events) + + by_part: dict[int, str] = {} + for event in out: + by_part[event.content_index] = by_part.get(event.content_index, "") + event.delta + assert by_part == {0: "one a@example.com", 1: "two"} + + @pytest.mark.asyncio + async def test_reasoning_summary_part_done_is_restored(self): + """The summary part repeats the whole summary text after its deltas.""" + guardrail, _ = _shielded(self.VAULT) + part = SimpleNamespace(type="summary_text", text="asked about [EMAIL_1]") + event = SimpleNamespace( + type="response.reasoning_summary_part.done", item_id="rs_1", output_index=0, summary_index=0, part=part + ) + + (out,) = await _restore_stream(guardrail, [event]) + + assert out.part.text == "asked about a@example.com" + + @pytest.mark.asyncio + async def test_mcp_call_arguments_are_restored(self): + """A stream family outside the chat-era set: matched by shape, not by name.""" + guardrail, _ = _shielded(self.VAULT) + deltas = [ + {"type": "response.mcp_call_arguments.delta", "item_id": "mcp_1", "output_index": 0, "delta": d} + for d in ('{"to": "[EMAI', 'L_1]"}') + ] + done = { + "type": "response.mcp_call_arguments.done", + "item_id": "mcp_1", + "output_index": 0, + "arguments": '{"to": "[EMAIL_1]"}', + } + + out = await _restore_stream(guardrail, [*deltas, done]) + + assert json.loads("".join(e["delta"] for e in out[:-1])) == {"to": "a@example.com"} + assert json.loads(out[-1]["arguments"]) == {"to": "a@example.com"} + assert out[-1]["item_id"] == "mcp_1", "identifiers are not text and stay as sent" + + @pytest.mark.asyncio + async def test_audio_deltas_are_not_sent_to_the_shield(self): + """Audio arrives base64-encoded; restoring it would cost a round trip for nothing.""" + guardrail, shield = _shielded(self.VAULT) + audio = {"type": "response.audio.delta", "item_id": "a_1", "output_index": 0, "delta": "UklGRiQAAABXQVZF"} + + out = await _restore_stream(guardrail, [audio]) + + assert out == [audio] + assert shield.urls == [] + + +class TestResponseCacheIsolation: + """The reply LiteLLM caches must keep its placeholders. + + Placeholders are numbered per request, so two callers' redacted requests can be + identical and share a cache key. LiteLLM keeps the provider's reply object -- for a + native Anthropic dict or an in-memory cache, the object itself -- so restoring it in + place would hand one caller's values to the next caller who hits that key. + """ + + @pytest.mark.asyncio + async def test_a_cache_hit_is_restored_against_the_new_callers_vault(self): + cache = InMemoryCache() + provider_reply = {"type": "message", "content": [{"type": "text", "text": "Repeat [EMAIL_1]"}]} + cache.set_cache("redacted-request", provider_reply) + + alice, _ = _shielded({"[EMAIL_1]": "alice@example.com"}) + to_alice = await alice.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=cache.get_cache("redacted-request") + ) + bob, _ = _shielded({"[EMAIL_1]": "bob@example.com"}) + to_bob = await bob.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=cache.get_cache("redacted-request") + ) + + assert to_alice["content"][0]["text"] == "Repeat alice@example.com" + assert to_bob["content"][0]["text"] == "Repeat bob@example.com" + assert cache.get_cache("redacted-request")["content"][0]["text"] == "Repeat [EMAIL_1]" + + @pytest.mark.asyncio + async def test_a_model_response_is_restored_as_a_copy(self): + """A reply as `acompletion` returns it, hidden params and all.""" + reply = await litellm.acompletion( + model="gpt-4o-mini", messages=[{"role": "user", "content": "hi"}], mock_response="Mail [EMAIL_1]" + ) + guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"}) + + restored = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=reply + ) + + assert restored.choices[0].message.content == "Mail a@example.com" + assert reply.choices[0].message.content == "Mail [EMAIL_1]" + + @pytest.mark.asyncio + async def test_responses_reply_is_restored_as_a_copy(self): + guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"}) + block = {"type": "output_text", "text": "Mail [EMAIL_1]"} + reply = SimpleNamespace(output=[SimpleNamespace(content=[block])]) + + restored = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=reply + ) + + assert restored.output[0].content[0]["text"] == "Mail a@example.com" + assert block["text"] == "Mail [EMAIL_1]" + + @pytest.mark.asyncio + async def test_stream_chunks_are_restored_as_copies(self): + """LiteLLM assembles the reply it caches from the chunks it yielded.""" + guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"}) + chunk = _chunk("Mail [EMAIL_1]", finish_reason="stop") + event = OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=0, + delta="Mail [EMAIL_1]", + ) + + out = await _restore_stream(guardrail, [chunk, event]) + + assert out[0].choices[0].delta.content == "Mail a@example.com" + assert chunk.choices[0].delta.content == "Mail [EMAIL_1]" + assert "Mail [EMAIL_1]" == event.delta + + +class TestCompletionsRestoration: + """`/v1/completions` replies carry their text on the choice, with no message or delta.""" + + VAULT = {"[EMAIL_1]": "a@example.com"} + + @pytest.mark.asyncio + async def test_completions_reply_is_restored(self): + guardrail, _ = _shielded(self.VAULT) + reply = TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAIL_1]", finish_reason="stop")]) + + restored = await guardrail.async_post_call_success_hook( + data={"prompt": "x"}, user_api_key_dict=None, response=reply + ) + + assert restored.choices[0].text == "Mail a@example.com" + + @pytest.mark.asyncio + async def test_completions_stream_is_restored_across_chunks(self): + guardrail, _ = _shielded(self.VAULT) + chunks = [ + TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAI")]), + TextCompletionResponse(choices=[TextChoices(index=0, text="L_1] now", finish_reason="stop")]), + ] + + out = await _restore_stream(guardrail, chunks) + + assert "".join(chunk.choices[0].text for chunk in out) == "Mail a@example.com now" + + @pytest.mark.asyncio + async def test_completions_stream_without_finish_reason_is_flushed(self): + guardrail, _ = _shielded(self.VAULT) + chunks = [TextCompletionResponse(choices=[TextChoices(index=0, text="Mail [EMAIL_1")])] + + out = await _restore_stream(guardrail, chunks) + + assert "".join(chunk.choices[0].text or "" for chunk in out) == "Mail [EMAIL_1" + + +class TestProxyWiring: + def test_dashboard_config_model_is_exposed(self): + """The guardrail garden reads the provider's fields from `get_config_model`.""" + model = LLMShieldProxyGuardrail.get_config_model() + + assert model is not None + assert {"api_key", "api_base"} <= set(model.model_fields) + + @pytest.mark.asyncio + async def test_the_deployment_hook_leaves_a_proxy_reply_for_the_proxy_hook(self): + """Inside the proxy the deployment hook must not restore: LiteLLM caches what it returns. + + A proxy request was redacted by the proxy's pre-call hook, so it carries no + deployment-restore marker, and the proxy's post-call hook restores it after the cache + write. A caller-sent marker that does not match the minted vault id is ignored. + """ + guardrail, shield = _shielded({"[EMAIL_1]": "a@example.com"}) + reply = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + data = {"messages": [], "guardrails": [GUARDRAIL_NAME], "litellm_metadata": {}} + LLMShieldProxyGuardrail._mint_session_id(data) + data["litellm_metadata"]["llm_shield_restore_at_deployment"] = "caller-chosen" + + result = await guardrail.async_post_call_success_deployment_hook( + request_data=data, response=reply, call_type=None + ) + + assert result is None + assert reply.choices[0].message.content == "[EMAIL_1]" + assert shield.urls == [] + + @pytest.mark.asyncio + async def test_model_level_use_outside_the_proxy_is_restored_and_never_cached(self, monkeypatch): + """SDK use with model-level `guardrails`: the deployment hooks are the only redact and + restore steps, so the reply is restored there, and the request bypasses the cache -- + its key is built from the redacted request, and a cache hit would skip restoration. + """ + vault = {"[EMAIL_1]": "alice@example.com"} + + async def shield(url: str, headers: dict, json: dict, timeout: float) -> Response: + texts = json["texts"] + if url.endswith("/redact"): + return _response({"texts": [t.replace("alice@example.com", "[EMAIL_1]") for t in texts]}) + restored = [] + for text in texts: + for placeholder, original in vault.items(): + text = text.replace(placeholder, original) + restored.append(text) + return _response({"texts": restored}) + + guardrail = _guardrail(event_hook=["pre_call", "post_call"], default_on=False) + guardrail.async_handler.post = shield # type: ignore[method-assign] + cache = InMemoryCache() + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + monkeypatch.setattr(litellm, "cache", litellm.Cache(type="local")) + monkeypatch.setattr(litellm.cache, "cache", cache) + + reply = await litellm.acompletion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Repeat alice@example.com"}], + mock_response="Repeat [EMAIL_1]", + guardrails=[GUARDRAIL_NAME], + ) + + # LiteLLM writes the cache from background tasks; let them land before looking. + await asyncio.gather(*_PENDING_CACHE_WRITES) + + assert reply.choices[0].message.content == "Repeat alice@example.com" + assert cache.cache_dict == {}, "the redacted request's reply must not be cached" + + @pytest.mark.asyncio + async def test_model_level_streaming_outside_the_proxy_is_refused(self, monkeypatch): + """No hook restores an SDK stream, and its cache writer misses the bypass, so it fails closed.""" + guardrail = _guardrail(event_hook=["pre_call", "post_call"], default_on=False) + _mock_post(guardrail, {"texts": ["Repeat [EMAIL_1]"]}) + cache = InMemoryCache() + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + monkeypatch.setattr(litellm, "cache", litellm.Cache(type="local")) + monkeypatch.setattr(litellm.cache, "cache", cache) + + with pytest.raises(GuardrailRaisedException, match="cannot restore a streamed reply"): + await litellm.acompletion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Repeat alice@example.com"}], + mock_response="Repeat [EMAIL_1]", + stream=True, + guardrails=[GUARDRAIL_NAME], + ) + await asyncio.gather(*_PENDING_CACHE_WRITES) + + assert cache.cache_dict == {} + + @pytest.mark.asyncio + async def test_restored_values_are_not_recorded_as_guardrail_telemetry(self): + """Guardrail logging is exported to traces even with message logging off, so the + restored reply must not land in it.""" + guardrail, _ = _shielded({"[EMAIL_1]": "alice@example.com"}) + reply = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content="[EMAIL_1]"))]) + data = {"messages": [], "metadata": {}} + + restored = await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=None, response=reply) + + assert restored.choices[0].message.content == "alice@example.com" + assert "alice@example.com" not in json.dumps(data, default=str) + + +class TestStreamUsage: + @pytest.mark.asyncio + async def test_trailing_flush_does_not_repeat_usage(self): + """With n>=2 and include_usage, the last chunk carries usage; a flush copied from it must not. + + Any consumer that sums usage chunks would otherwise count the request twice. + """ + guardrail, _ = _shielded({"[EMAIL_1]": "a@example.com"}) + both = ModelResponseStream( + choices=[ + StreamingChoices(index=0, delta=Delta(content="Mail [EMAI")), + StreamingChoices(index=1, delta=Delta(content="Call [EMAI")), + ] + ) + usage_chunk = ModelResponseStream( + choices=[StreamingChoices(index=1, delta=Delta(content=None))], + usage=Usage(prompt_tokens=5, completion_tokens=7, total_tokens=12), + ) + + out = await _restore_stream(guardrail, [both, usage_chunk]) + + with_usage = [chunk for chunk in out if getattr(chunk, "usage", None) is not None] + assert with_usage == [out[1]], "only the provider's own usage chunk carries usage" + flushed = out[2:] + assert flushed, "the held-back text is flushed at end of stream" + assert sorted(chunk.choices[0].index for chunk in flushed) == [0, 1] + assert [chunk.choices[0].delta.content for chunk in flushed] == ["[EMAI", "[EMAI"] diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index e547575ef9c..e797304a1c6 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -1,16 +1,21 @@ """Tests for unified guardrail.""" +import asyncio +import contextlib import io import logging +from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator from types import SimpleNamespace from typing import TYPE_CHECKING, Final, Literal +import anyio import pytest import litellm from litellm.caching import DualCache from litellm.integrations.custom_guardrail import ( CustomGuardrail, + ModifyResponseException, log_guardrail_information, ) from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route @@ -42,7 +47,15 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai ) from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import ResponsesAPIResponse -from litellm.types.utils import CallTypes, Delta, GenericGuardrailAPIInputs, ModelResponseStream, StreamingChoices +from litellm.types.utils import ( + CallTypes, + Delta, + GenericGuardrailAPIInputs, + ModelResponse, + ModelResponseStream, + StandardLoggingGuardrailInformation, + StreamingChoices, +) if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -2164,6 +2177,554 @@ class _ScanCountingGuardrail(CustomGuardrail): return inputs +class _GatedScanGuardrail(_ScanCountingGuardrail): + """End-of-stream scan that holds until released, recording scans that finished.""" + + def __init__(self) -> None: + super().__init__(end_of_stream_only=True) + self.scan_started = anyio.Event() + self.scan_released = anyio.Event() + self.finished_scans = 0 + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + self.scan_started.set() + await self.scan_released.wait() + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + self.finished_scans += 1 + return recorded + + +class _FinishReasonRecordingGuardrail(_ScanCountingGuardrail): + """Records the finish reasons of the stream handed to each response-side scan""" + + def __init__(self) -> None: + super().__init__(end_of_stream_only=True) + self.finish_reasons: tuple[str | None, ...] = () + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + rebuilt = request_data.get("response") + released = request_data.get("responses") + chunks = ( + tuple(chunk for chunk in released if isinstance(chunk, ModelResponseStream)) + if isinstance(released, list) + else () + ) + choices = [ + *(rebuilt.choices if isinstance(rebuilt, ModelResponse) else ()), + *(choice for chunk in chunks for choice in chunk.choices), + ] + self.finish_reasons = (*self.finish_reasons, *(choice.finish_reason for choice in choices)) + return await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + + +class _GatedToolCallGuardrail(_StreamingTextGuardrail): + """Tool-call inspection that holds until released""" + + def __init__(self) -> None: + super().__init__() + self.inspection_started = anyio.Event() + self.inspection_released = anyio.Event() + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + if input_type == "response" and inputs.get("tool_calls"): + self.inspection_started.set() + await self.inspection_released.wait() + return await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + + +class _MarkerBlockingScanGuardrail(_ScanCountingGuardrail): + """Scan-counting guardrail that blocks any scan whose text contains BLOCKME""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + if any("BLOCKME" in text for text in recorded.get("texts") or []): + raise ModifyResponseException( + message="blocked", model="gpt-4", request_data=request_data, guardrail_name=self.guardrail_name + ) + return recorded + + +class _MarkerHttpErrorScanGuardrail(_ScanCountingGuardrail): + """Scan-counting guardrail that raises an HTTPException for any scan whose text contains BLOCKME""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + if any("BLOCKME" in text for text in recorded.get("texts") or []): + raise unified_module.HTTPException(status_code=400, detail={"error": "Violated guardrail policy"}) + return recorded + + +class _MarkerBlockingStreamingTextGuardrail(_StreamingTextGuardrail): + """incremental_diff guardrail that blocks any round whose text contains BLOCKME""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + transformed = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + if any("BLOCKME" in text for text in inputs.get("texts") or []): + raise ModifyResponseException( + message="blocked", model="gpt-4", request_data=request_data, guardrail_name=self.guardrail_name + ) + return transformed + + +class _DisconnectRewritingGuardrail(_ScanCountingGuardrail): + """End-of-stream guardrail that rewrites every scanned text""" + + def __init__(self) -> None: + super().__init__(end_of_stream_only=True) + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + return {**recorded, "texts": ["REWRITTEN" for _ in recorded.get("texts") or []]} + + +class _RecordedScanGuardrail(CustomGuardrail): + """End-of-stream scan recorded through log_guardrail_information, returning ``reply`` or raising ``error``""" + + def __init__(self, *, reply: GenericGuardrailAPIInputs | None = None, error: Exception | None = None) -> None: + super().__init__(guardrail_name="recorded-scan") + self.streaming_end_of_stream_only = True + self.streaming_buffer_until_moderated = False + self.guardrail_config = {} + self._reply = reply + self._error = error + + def should_run_guardrail(self, data: dict[str, object], event_type: GuardrailEventHooks) -> bool: + return True + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + if self._error is not None: + raise self._error + return inputs if self._reply is None else self._reply + + +def _recorded_guardrail_statuses(request_data: dict[str, object]) -> list[str]: + metadata = request_data["metadata"] + assert isinstance(metadata, dict), request_data + entries: list[StandardLoggingGuardrailInformation] = metadata.get("standard_logging_guardrail_information", []) + return [entry["guardrail_status"] for entry in entries] + + +class TestStreamingClientDisconnectScan: + """A client that reads streamed content and then disconnects must not skip + the end-of-stream scan of what it already received.""" + + @pytest.fixture(autouse=True) + def _use_real_mappings(self, monkeypatch: pytest.MonkeyPatch) -> None: + _patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings()) + + @staticmethod + def _guarded_stream(guardrail: CustomGuardrail, upstream: AsyncIterable[object]) -> AsyncGenerator[object, None]: + return UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/chat/completions"), + response=upstream, + request_data={"guardrail_to_apply": guardrail, "model": "gpt-4", "metadata": {}}, + ) + + @pytest.mark.asyncio + async def test_closing_after_released_content_still_scans_it(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert _delta_text(received) == "synthetic secret" + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_mid_text_stream_does_not_hand_the_scan_a_tool_calls_finish(self): + guardrail = _FinishReasonRecordingGuardrail() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + await stream.__anext__() + await stream.aclose() + + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + assert "tool_calls" not in guardrail.finish_reasons, guardrail.finish_reasons + + @pytest.mark.asyncio + async def test_upstream_cancellation_after_released_content_still_scans_it(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + raise asyncio.CancelledError() + + stream = self._guarded_stream(guardrail, upstream()) + received = await stream.__anext__() + with pytest.raises(asyncio.CancelledError): + await stream.__anext__() + + assert _delta_text(received) == "synthetic secret" + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + + @pytest.mark.asyncio + async def test_cancellation_while_the_disconnect_scan_is_in_flight_lets_it_finish(self): + guardrail = _GatedScanGuardrail() + first_chunk_received = anyio.Event() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + await anyio.sleep_forever() + yield _stream_chunk(" tail", finish_reason="stop") + + async def consume(scope_ready: list[anyio.CancelScope]) -> None: + with anyio.CancelScope() as scope: + scope_ready.append(scope) + async with contextlib.aclosing(self._guarded_stream(guardrail, upstream())) as stream: + async for _item in stream: + first_chunk_received.set() + + scopes = [] + with anyio.fail_after(5): + async with anyio.create_task_group() as task_group: + task_group.start_soon(consume, scopes) + await first_chunk_received.wait() + scopes[0].cancel() + await guardrail.scan_started.wait() + await anyio.sleep(0) + guardrail.scan_released.set() + + assert guardrail.finished_scans == 1 + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_after_a_sampled_scan_covered_everything_released_does_not_scan_again(self): + guardrail = _ScanCountingGuardrail(sampling_rate=1) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic") + yield _stream_chunk(" secret") + await anyio.sleep_forever() + + stream = self._guarded_stream(guardrail, upstream()) + released = [await stream.__anext__(), await stream.__anext__()] + await stream.aclose() + + scanned_texts = [scan["texts"] for scan in guardrail.scans] + assert "".join(_delta_text(chunk) for chunk in released) == "synthetic secret" + assert scanned_texts[-1] == ["synthetic secret"], scanned_texts + assert len(scanned_texts) == len({tuple(texts) for texts in scanned_texts}), scanned_texts + + @pytest.mark.asyncio + async def test_closing_before_any_content_is_released_does_not_scan(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True, buffer_until_moderated=True) + upstream_started = anyio.Event() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("withheld") + upstream_started.set() + await anyio.sleep_forever() + yield _stream_chunk("never", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + async with anyio.create_task_group() as task_group: + task_group.start_soon(stream.__anext__) + await upstream_started.wait() + task_group.cancel_scope.cancel() + + assert guardrail.scans == () + + @pytest.mark.asyncio + async def test_cancellation_during_end_of_stream_scan_lets_the_scan_finish(self): + guardrail = _GatedScanGuardrail() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + async def consume(scope_ready: list[anyio.CancelScope]) -> None: + with anyio.CancelScope() as scope: + scope_ready.append(scope) + async for _item in self._guarded_stream(guardrail, upstream()): + pass + + scopes = [] + async with anyio.create_task_group() as task_group: + task_group.start_soon(consume, scopes) + await guardrail.scan_started.wait() + scopes[0].cancel() + await anyio.sleep(0) + guardrail.scan_released.set() + + assert guardrail.finished_scans == 1 + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret tail"]], guardrail.scans + + @staticmethod + async def _close_after_first_chunk(guardrail: CustomGuardrail) -> dict[str, object]: + request_data: dict[str, object] = {"guardrail_to_apply": guardrail, "model": "gpt-4", "metadata": {}} + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/chat/completions"), + response=upstream(), + request_data=request_data, + ) + received = await stream.__anext__() + await stream.aclose() + assert _delta_text(received) == "synthetic secret" + return request_data + + @pytest.mark.asyncio + async def test_disconnect_scan_that_fails_after_the_verdict_records_the_failure(self): + request_data = await self._close_after_first_chunk( + _RecordedScanGuardrail(reply={"texts": ["synthetic secret", "unmatched extra text"]}) + ) + + assert _recorded_guardrail_statuses(request_data) == ["success", "guardrail_failed_to_respond"] + + @pytest.mark.asyncio + async def test_disconnect_scan_whose_guardrail_raises_records_one_failure(self): + request_data = await self._close_after_first_chunk(_RecordedScanGuardrail(error=RuntimeError("provider down"))) + + assert _recorded_guardrail_statuses(request_data) == ["guardrail_failed_to_respond"] + + @pytest.mark.asyncio + async def test_disconnect_scan_that_passes_records_only_the_verdict(self): + request_data = await self._close_after_first_chunk(_RecordedScanGuardrail()) + + assert _recorded_guardrail_statuses(request_data) == ["success"] + + @staticmethod + async def _tool_call_upstream() -> AsyncIterator[ModelResponseStream]: + from litellm.types.utils import ChatCompletionDeltaToolCall, Function + + tool_call = ChatCompletionDeltaToolCall( + id="call_1", index=0, type="function", function=Function(name="get_weather", arguments='{"city": "Paris"}') + ) + yield ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(content=None, tool_calls=[tool_call]))] + ) + yield _stream_chunk(None, finish_reason="tool_calls") + + @pytest.mark.asyncio + async def test_cancellation_during_incremental_diff_tool_call_inspection_lets_it_finish_once(self): + guardrail = _GatedToolCallGuardrail() + + async def consume(scope_ready: list[anyio.CancelScope]) -> None: + with anyio.CancelScope() as scope: + scope_ready.append(scope) + async with contextlib.aclosing(self._guarded_stream(guardrail, self._tool_call_upstream())) as stream: + async for _item in stream: + pass + + scopes = [] + async with anyio.create_task_group() as task_group: + task_group.start_soon(consume, scopes) + await guardrail.inspection_started.wait() + scopes[0].cancel() + await anyio.sleep(0) + guardrail.inspection_released.set() + + assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [ + ["get_weather"] + ], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_after_the_incremental_diff_tool_call_inspection_does_not_inspect_again(self): + guardrail = _StreamingTextGuardrail() + + stream = self._guarded_stream(guardrail, self._tool_call_upstream()) + released = [await stream.__anext__(), await stream.__anext__()] + await stream.aclose() + + assert released[-1].choices[0].finish_reason == "tool_calls" + assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [ + ["get_weather"] + ], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_after_a_released_tool_call_under_incremental_diff_still_inspects_it(self): + guardrail = _StreamingTextGuardrail() + + stream = self._guarded_stream(guardrail, self._tool_call_upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert [call.function.name for call in received.choices[0].delta.tool_calls] == ["get_weather"] + assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [ + ["get_weather"] + ], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_while_a_mid_stream_block_is_delivered_does_not_scan_the_blocked_content_again(self): + guardrail = _MarkerBlockingScanGuardrail(sampling_rate=2) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + for text in ("a", "b", "c", "BLOCKME"): + yield _stream_chunk(text) + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()) for _ in range(3)] + await stream.__anext__() + await stream.aclose() + + assert received == ["a", "b", "c"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab"], ["abcBLOCKME"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_while_a_mid_stream_guardrail_error_is_delivered_does_not_scan_the_content_again(self): + guardrail = _MarkerHttpErrorScanGuardrail(sampling_rate=2) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + for text in ("a", "b", "c", "BLOCKME"): + yield _stream_chunk(text) + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()) for _ in range(3)] + await stream.__anext__() + await stream.aclose() + + assert received == ["a", "b", "c"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab"], ["abcBLOCKME"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_while_the_scanned_final_chunk_is_delivered_does_not_scan_the_stream_again(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("a") + yield _stream_chunk("b") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()) for _ in range(3)] + await stream.aclose() + + assert received == ["a", "b", " tail"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab tail"]], guardrail.scans + + @pytest.mark.asyncio + async def test_cancellation_with_a_withheld_window_scans_only_the_released_chunks(self): + guardrail = _ScanCountingGuardrail(sampling_rate=2, buffer_until_moderated=True) + guardrail.streaming_buffer_release_on_scan = True + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("a") + yield _stream_chunk("b") + yield _stream_chunk("WITHHELD") + raise asyncio.CancelledError() + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()), _delta_text(await stream.__anext__())] + with pytest.raises(asyncio.CancelledError): + await stream.__anext__() + + assert received == ["a", "b"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab"]], guardrail.scans + + @pytest.mark.asyncio + async def test_disconnect_scan_that_rewrites_text_leaves_the_released_chunks_untouched(self): + stream = self._guarded_stream(_DisconnectRewritingGuardrail(), self._text_then_tail_upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert _delta_text(received) == "synthetic secret" + + @staticmethod + async def _text_then_tail_upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + @pytest.mark.asyncio + async def test_closing_while_an_incremental_diff_block_is_delivered_does_not_inspect_the_tool_call_again(self): + guardrail = _MarkerBlockingStreamingTextGuardrail() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + async for tool_chunk in self._tool_call_upstream(): + if tool_chunk.choices[0].finish_reason is None: + yield tool_chunk + yield _stream_chunk("BLOCKME") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + released = [await stream.__anext__(), await stream.__anext__()] + await stream.aclose() + + assert [call.function.name for call in released[0].choices[0].delta.tool_calls] == ["get_weather"] + assert guardrail.received_texts == [["BLOCKME"]], guardrail.received_texts + assert guardrail.received_tool_calls == [], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_after_a_released_tool_call_under_incremental_diff_scans_no_held_back_text(self): + guardrail = _StreamingTextGuardrail(holdback_schedule=[len("held secret")] * 2) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("held secret") + async for tool_chunk in self._tool_call_upstream(): + yield tool_chunk + + stream = self._guarded_stream(guardrail, upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert [call.function.name for call in received.choices[0].delta.tool_calls] == ["get_weather"] + assert guardrail.received_tool_calls, guardrail.received_texts + assert all("held secret" not in text for text in guardrail.received_texts[-1]), guardrail.received_texts + + def _responses_delta(sequence_number, text): return { "type": "response.output_text.delta", diff --git a/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py b/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py index 9c785d59830..31f1c054c59 100644 --- a/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py +++ b/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py @@ -7,6 +7,7 @@ Verifies that the hook: 3. Actually yields chunks from async generators """ +import logging from typing import AsyncGenerator, Any from unittest.mock import MagicMock, patch @@ -31,7 +32,7 @@ class MockStreamingCallback(CustomLogger): self, user_api_key_dict: UserAPIKeyAuth, response: AsyncGenerator[Any, None], - request_data: dict, + request_data: dict[str, object], ) -> AsyncGenerator[Any, None]: """Transform chunks by tracking and optionally prefixing.""" async for chunk in response: @@ -185,3 +186,340 @@ async def test_streaming_hook_propagates_callback_errors(): with pytest.raises(RuntimeError, match="Callback failed!"): async for _ in result: pass + + +class CleanupRecordingCallback(CustomLogger): + """Iterator hook whose cleanup marks when it ran.""" + + def __init__(self): + super().__init__() + self.cleaned_up = False + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> AsyncGenerator[Any, None]: + try: + async for chunk in response: + yield chunk + finally: + self.cleaned_up = True + + +@pytest.mark.asyncio +async def test_closing_the_stream_runs_every_callback_cleanup_before_returning(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callbacks = [CleanupRecordingCallback(), CleanupRecordingCallback()] + + with patch.object(litellm, "callbacks", callbacks): + ProxyLogging._callback_capabilities_cache.clear() + stream = proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + first = await stream.__anext__() + await stream.aclose() + ProxyLogging._callback_capabilities_cache.clear() + + assert first == {"choices": [{"delta": {"content": "Hello"}}]} + assert [callback.cleaned_up for callback in callbacks] == [True, True] + + +class RaisingCleanupCallback(CustomLogger): + """Iterator hook whose cleanup raises.""" + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> AsyncGenerator[Any, None]: + try: + async for chunk in response: + yield chunk + finally: + raise RuntimeError("cleanup failed") + + +@pytest.mark.asyncio +async def test_closing_the_stream_still_cleans_up_inner_callbacks_when_an_outer_cleanup_raises(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + inner = CleanupRecordingCallback() + + with patch.object(litellm, "callbacks", [inner, RaisingCleanupCallback()]): + ProxyLogging._callback_capabilities_cache.clear() + stream = proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + await stream.__anext__() + await stream.aclose() + ProxyLogging._callback_capabilities_cache.clear() + + assert inner.cleaned_up is True + + +class _PlainAsyncIterator: + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + self._response = response + + def __aiter__(self) -> "_PlainAsyncIterator": + return self + + async def __anext__(self) -> Any: + return await self._response.__anext__() + + +class PlainIteratorCallback(CustomLogger): + """Iterator hook that returns an async iterator with no aclose.""" + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a plain async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _PlainAsyncIterator: + return _PlainAsyncIterator(response) + + +@pytest.mark.asyncio +async def test_a_hook_returning_a_plain_async_iterator_streams_every_chunk(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + + with patch.object(litellm, "callbacks", [PlainIteratorCallback()]): + ProxyLogging._callback_capabilities_cache.clear() + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + + +class _ClosableAsyncIterator(_PlainAsyncIterator): + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + super().__init__(response) + self.closed = False + + async def aclose(self) -> None: + self.closed = True + + +class _SyncClosableAsyncIterator(_PlainAsyncIterator): + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + super().__init__(response) + self.closed = False + + def aclose(self) -> None: + self.closed = True + + +class ClosableIteratorCallback(CustomLogger): + """Iterator hook that returns a non-generator async iterator with its own aclose.""" + + def __init__( + self, + iterator_type: type[_ClosableAsyncIterator] | type[_SyncClosableAsyncIterator] = _ClosableAsyncIterator, + ) -> None: + super().__init__() + self.iterator_type = iterator_type + self.returned: tuple[_ClosableAsyncIterator | _SyncClosableAsyncIterator, ...] = () + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _ClosableAsyncIterator | _SyncClosableAsyncIterator: + iterator = self.iterator_type(response) + self.returned = (*self.returned, iterator) + return iterator + + +@pytest.mark.asyncio +async def test_closing_the_stream_closes_a_hook_iterator_that_is_not_a_generator(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = ClosableIteratorCallback() + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + stream = proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + await stream.__anext__() + await stream.aclose() + ProxyLogging._callback_capabilities_cache.clear() + + assert [iterator.closed for iterator in callback.returned] == [True] + + +@pytest.mark.asyncio +async def test_a_hook_iterator_with_a_synchronous_aclose_streams_everything_and_is_closed(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = ClosableIteratorCallback(iterator_type=_SyncClosableAsyncIterator) + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + + +class _RaisingAcloseIterator(_ClosableAsyncIterator): + """Non-generator async iterator whose asynchronous aclose raises.""" + + def __init__(self, response: AsyncGenerator[Any, None], error: Exception) -> None: + super().__init__(response) + self.error = error + + async def aclose(self) -> None: + self.closed = True + raise self.error + + +class _SyncRaisingAcloseIterator(_SyncClosableAsyncIterator): + """Non-generator async iterator whose synchronous aclose raises.""" + + def __init__(self, response: AsyncGenerator[Any, None], error: Exception) -> None: + super().__init__(response) + self.error = error + + def aclose(self) -> None: + self.closed = True + raise self.error + + +class RaisingAcloseIteratorCallback(CustomLogger): + """Iterator hook that returns a non-generator async iterator whose aclose raises.""" + + def __init__( + self, + iterator_type: type[_RaisingAcloseIterator] | type[_SyncRaisingAcloseIterator] = _RaisingAcloseIterator, + error: Exception | None = None, + ) -> None: + super().__init__() + self.iterator_type = iterator_type + self.error = error if error is not None else RuntimeError("cleanup failed") + self.returned: tuple[_RaisingAcloseIterator | _SyncRaisingAcloseIterator, ...] = () + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _RaisingAcloseIterator | _SyncRaisingAcloseIterator: + iterator = self.iterator_type(response, self.error) + self.returned = (*self.returned, iterator) + return iterator + + +@pytest.mark.asyncio +async def test_a_hook_iterator_whose_aclose_raises_still_finishes_the_stream(caplog: pytest.LogCaptureFixture) -> None: + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = RaisingAcloseIteratorCallback() + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + warnings_emitted = [ + record.getMessage() + for record in caplog.records + if record.levelname == "WARNING" and "RaisingAcloseIteratorCallback" in record.getMessage() + ] + assert len(warnings_emitted) == 1 + assert "RuntimeError" in warnings_emitted[0] + assert "cleanup failed" not in warnings_emitted[0] + + +@pytest.mark.asyncio +async def test_a_hook_iterator_whose_synchronous_aclose_raises_still_finishes_the_stream( + caplog: pytest.LogCaptureFixture, +) -> None: + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = RaisingAcloseIteratorCallback( + iterator_type=_SyncRaisingAcloseIterator, error=ValueError("sync cleanup failed") + ) + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + warnings_emitted = [ + record.getMessage() + for record in caplog.records + if record.levelname == "WARNING" and "RaisingAcloseIteratorCallback" in record.getMessage() + ] + assert len(warnings_emitted) == 1 + assert "ValueError" in warnings_emitted[0] + + +@pytest.mark.asyncio +async def test_a_hook_iterator_with_a_clean_aclose_streams_everything_without_warning( + caplog: pytest.LogCaptureFixture, +) -> None: + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = ClosableIteratorCallback() + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + assert not [ + record.getMessage() + for record in caplog.records + if record.levelname == "WARNING" and "ClosableIteratorCallback" in record.getMessage() + ] diff --git a/tests/unit/proxy/management_endpoints/policy_endpoints/test_endpoints.py b/tests/unit/proxy/management_endpoints/policy_endpoints/test_endpoints.py index 86f9aeafc08..62863163b8e 100644 --- a/tests/unit/proxy/management_endpoints/policy_endpoints/test_endpoints.py +++ b/tests/unit/proxy/management_endpoints/policy_endpoints/test_endpoints.py @@ -6,11 +6,21 @@ without needing a running proxy. """ import pytest +from fastapi import Request +from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.management_endpoints.policy_endpoints.endpoints import ( GuardrailTestResultEntry, _compute_overall_action, _test_guardrail_definitions, + list_policies, +) +from litellm.proxy.policy_engine.policy_registry import get_policy_registry +from litellm.types.proxy.policy_engine import ( + PolicyGuardrailsResponse, + PolicyListResponse, + PolicyScopeResponse, + PolicySummaryItem, ) @@ -293,3 +303,68 @@ class TestEnrichPolicyTemplateStreamKeepalive: assert b": ping\n\n" not in chunks assert chunks[0] == b'data: {"type": "competitor", "name": "Rival Air"}\n\n' assert chunks[-1].startswith(b'data: {"type": "done"') + + +def _policy_list_request() -> Request: + return Request({"type": "http", "method": "GET", "path": "/policy/list", "headers": []}) + + +def _summary_item(inherit, resolved_guardrails, inheritance_chain) -> PolicySummaryItem: + return PolicySummaryItem( + inherit=inherit, + scope=PolicyScopeResponse(), + guardrails=PolicyGuardrailsResponse(), + resolved_guardrails=resolved_guardrails, + inheritance_chain=inheritance_chain, + ) + + +@pytest.fixture +def policy_registry(): + registry = get_policy_registry() + registry.clear() + yield registry + registry.clear() + + +@pytest.mark.asyncio +async def test_list_policies_is_empty_before_any_policy_is_loaded(policy_registry): + response = await list_policies(request=_policy_list_request(), user_api_key_dict=UserAPIKeyAuth()) + + assert response == PolicyListResponse(policies={}, total_count=0) + + +@pytest.mark.parametrize( + ("policies_config", "expected_policies"), + [ + ({}, {}), + ( + {"solo": {"description": "standalone", "guardrails": {"add": ["pii"]}}}, + {"solo": _summary_item(None, ["pii"], ["solo"])}, + ), + ( + { + "base": {"guardrails": {"add": ["pii"]}}, + "child": {"inherit": "base", "guardrails": {"add": ["audit"], "remove": ["pii"]}}, + "conditional": {"guardrails": {"add": ["toxicity"]}, "condition": {"model": "gpt-4.*"}}, + "empty": {"guardrails": {"remove": ["pii"]}}, + }, + { + "base": _summary_item(None, ["pii"], ["base"]), + "child": _summary_item("base", ["audit"], ["base", "child"]), + "conditional": _summary_item(None, ["toxicity"], ["conditional"]), + "empty": _summary_item(None, [], ["empty"]), + }, + ), + ], +) +@pytest.mark.asyncio +async def test_list_policies_reports_each_loaded_policy_with_its_resolved_guardrails( + policy_registry, policies_config, expected_policies +): + policy_registry.load_policies(policies_config) + + response = await list_policies(request=_policy_list_request(), user_api_key_dict=UserAPIKeyAuth()) + + assert response == PolicyListResponse(policies=expected_policies, total_count=0) + assert list(response.policies) == list(expected_policies) diff --git a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py index 01e41e8b03f..0aff5852c96 100644 --- a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py @@ -150,6 +150,89 @@ async def _classifier_user_payload(body: Mapping[str, object], monkeypatch: pyte return router.recorded_calls[0]["messages"][1]["content"] +@pytest.mark.parametrize( + "tiers,config_default,explicit_default,expected", + ( + ({"SIMPLE": "cheap-model", "MEDIUM": "mid-model"}, None, None, "mid-model"), + ({"SIMPLE": ["cheap-model"], "MEDIUM": ["mid-model", "strong-model"]}, None, None, "mid-model"), + ({"SIMPLE": ["cheap-model"], "MEDIUM": []}, None, None, "cheap-model"), + ({"SIMPLE": "cheap-model"}, None, None, "cheap-model"), + ({"MEDIUM": "mid-model"}, "cheap-model", None, "cheap-model"), + ({"MEDIUM": "mid-model"}, "cheap-model", "strong-model", "strong-model"), + ({"MEDIUM": "mid-model"}, None, "strong-model", "strong-model"), + ({}, "cheap-model", None, "cheap-model"), + ({}, None, "strong-model", "strong-model"), + ), +) +@pytest.mark.asyncio +async def test_preview_and_serving_share_default_model_resolution( + monkeypatch: pytest.MonkeyPatch, + tiers: Mapping[str, object], + config_default: str | None, + explicit_default: str | None, + expected: str, +): + from litellm.router_utils.auto_router_model_naming import validate_complexity_router_config_write + from litellm.types.management_endpoints.auto_router_endpoints import ComplexityRouterConfigValidationRequest + + config: Final = { + "tiers": tiers, + "default_model": config_default, + "classifier_type": "llm", + "classifier_fallback": "default_model", + "classifier_llm_config": {"model": "unconfigured-classifier"}, + } + assert validate_complexity_router_config_write(config) is None + verdict: Final = await auto_router_endpoints.validate_complexity_router_config( + ComplexityRouterConfigValidationRequest(complexity_router_config=config), ADMIN + ) + assert verdict.valid and verdict.error is None + serving: Final = _router() + serving.init_complexity_router_deployment( + Deployment( + model_name="default-parity", + litellm_params={ + "model": "auto_router/complexity_router", + "complexity_router_config": config, + "complexity_router_default_model": explicit_default, + }, + model_info={"id": "default-parity"}, + ) + ) + strategy: Final = serving.complexity_routers["default-parity"][0].strategy + assert strategy.config.default_model == expected + decision: Final = await strategy.async_pre_routing_hook( + model="default-parity", messages=[{"role": "user", "content": "hello"}], request_kwargs={} + ) + assert decision is not None and decision.model == expected + monkeypatch.setattr(proxy_server, "llm_router", serving) + preview: Final = await preview_auto_router_routing( + http_request=ROUTING_HTTP_REQUEST, + data=AutoRouterRoutingTestRequest.model_validate( + {"prompt": "hello", "complexity_router_config": config, "default_model": explicit_default} + ), + user_api_key_dict=ADMIN, + ) + assert preview.routed_model == expected + assert preview.routing_decision["cause"] == "default_model_fallback" + + +@pytest.mark.asyncio +async def test_preview_missing_unresolvable_default_is_a_config_error(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(proxy_server, "llm_router", _router()) + with pytest.raises(HTTPException) as error: + await preview_auto_router_routing( + http_request=ROUTING_HTTP_REQUEST, + data=_request( + "hello", tiers={}, classifier_type="llm", classifier_fallback="default_model", + classifier_llm_config={"model": "unconfigured-classifier"}, + ), + user_api_key_dict=ADMIN, + ) + assert error.value.status_code == 400 + assert "requires a default model" in error.value.detail["error"] + + @pytest.mark.asyncio async def test_simple_prompt_routes_to_the_simple_tier(monkeypatch: pytest.MonkeyPatch): response = await _route("what is 2+2", monkeypatch) @@ -195,12 +278,17 @@ async def test_escalation_keyword_bumps_the_classified_tier(monkeypatch: pytest. assert response.routing_decision["escalation_keyword"] == "ultrathink" +@pytest.mark.parametrize("default_model,expected,configured", ((None, "mid-model", True), ("never-configured", "never-configured", False))) @pytest.mark.asyncio -async def test_tier_model_missing_from_the_proxy_is_reported(monkeypatch: pytest.MonkeyPatch): - response = await _route("what is 2+2", monkeypatch, tiers={**TIERS, "SIMPLE": ["never-configured"]}) +async def test_tier_model_missing_from_the_proxy_is_reported( + monkeypatch: pytest.MonkeyPatch, default_model: str | None, expected: str, configured: bool +): + response = await _route( + "what is 2+2", monkeypatch, tiers={**TIERS, "SIMPLE": ["never-configured"]}, default_model=default_model + ) - assert response.routed_model == "never-configured" - assert response.routed_model_configured is False + assert response.routed_model == expected + assert response.routed_model_configured is configured @pytest.mark.asyncio diff --git a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py index 2cbf1d578b2..081a2d8ce73 100644 --- a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1876,7 +1876,7 @@ class TestTemporaryMCPSessionEndpoints: url="https://temp.example.com", transport=MCPTransport.http, ) - existing_server = MagicMock() + existing_server = MagicMock(dcr_issuer=None, dcr_server_url=None, token_endpoint_auth_method=None) existing_server.authentication_token = "token-abc" existing_server.client_id = "client-123" existing_server.client_secret = "secret-xyz" @@ -1912,7 +1912,7 @@ class TestTemporaryMCPSessionEndpoints: @staticmethod def _inherit_with(payload_credentials, **server_overrides): - existing_server = MagicMock() + existing_server = MagicMock(dcr_issuer=None, dcr_server_url=None, token_endpoint_auth_method=None) existing_server.authentication_token = None existing_server.client_id = "client-123" existing_server.client_secret = "secret-xyz" @@ -2616,6 +2616,9 @@ class TestTemporaryMCPSessionEndpoints: user_id="admin-user", ) inherited_server = MagicMock( + dcr_issuer=None, + dcr_server_url=None, + token_endpoint_auth_method=None, authentication_token="token-abc", client_id="client-id", client_secret="client-secret", @@ -11145,3 +11148,177 @@ async def test_protocol_update_on_missing_server_preserves_not_found() -> None: user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), ) assert error.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_changed_upstream_session_gets_an_isolated_id_without_saved_client(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-upstream", name="saved", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + client_id="old-client", client_secret="old-secret", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + payload = NewMCPServerRequest( + server_id=saved.server_id, server_name="saved", transport="http", auth_type="oauth2", + oauth2_flow="authorization_code", url="https://new.example/mcp", issuer="https://new.example", + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert not (staged.credentials or {}).get("client_id") + assert not (staged.credentials or {}).get("client_secret") + assert await management._resolve_session_server_id(staged) != saved.server_id + assert saved.client_id == "old-client" + + +def test_staged_url_edit_clears_resubmitted_issuer_and_endpoints(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-oauth", name="saved", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + authorization_url="https://old.example/authorize", token_url="https://old.example/token", + registration_url="https://old.example/register", client_id="old-client", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + payload = NewMCPServerRequest( + server_id=saved.server_id, server_name="saved", transport="http", auth_type="oauth2", + oauth2_flow="authorization_code", url="https://new.example/mcp", issuer=saved.issuer, + authorization_url=saved.authorization_url, token_url=saved.token_url, registration_url=saved.registration_url, + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert staged.issuer is None + assert staged.authorization_url is None + assert staged.token_url is None + assert staged.registration_url is None + assert staged.oauth2_flow == "authorization_code" + assert saved.issuer == "https://old.example" + + +@pytest.mark.asyncio +async def test_distinct_issuer_identifier_edit_isolates_session_client(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-oauth", name="saved", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer="https://idp.example", + client_id="saved-client", client_secret="saved-secret", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + payload = NewMCPServerRequest( + server_id=saved.server_id, server_name="saved", transport="http", auth_type="oauth2", + url=saved.url, issuer="https://IDP.example:443/", + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert not (staged.credentials or {}).get("client_id") + assert not (staged.credentials or {}).get("client_secret") + assert await management._resolve_session_server_id(staged) != saved.server_id + + +@pytest.mark.asyncio +async def test_url_edit_stages_existing_client_with_previous_issuer_binding(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-oauth", name="saved", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer="https://idp.example", + client_id="static-client", client_secret="static-secret", authentication_token="old-token", + token_endpoint_auth_method="client_secret_basic", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + payload = NewMCPServerRequest( + server_id=saved.server_id, server_name="saved", transport="http", auth_type="oauth2", + url=saved.url + "?v=2", issuer=saved.issuer, + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert staged.credentials["client_id"] == "static-client" + assert staged.credentials["client_secret"] == "static-secret" + assert staged.credentials["dcr_issuer"] == saved.issuer + assert staged.credentials["dcr_server_url"] == saved.url + assert staged.credentials["token_endpoint_auth_method"] == "client_secret_basic" + assert "auth_value" not in staged.credentials + assert staged.issuer is None + assert await management._resolve_session_server_id(staged) != saved.server_id + + +@pytest.mark.asyncio +@pytest.mark.parametrize("upstream_changed", [False, True]) +async def test_session_identity_checks_saved_row_when_registry_is_empty(monkeypatch, upstream_changed): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + + saved = generate_mock_mcp_server_db_record(server_id="saved-oauth-row") + saved.auth_type = "oauth2" + saved.url = "https://old.example/mcp" + saved.issuer = "https://old.example" + saved.approval_status = "approved" + monkeypatch.setattr(management.global_mcp_server_manager, "get_mcp_server_by_id", lambda _: None) + monkeypatch.setattr(management, "_get_prisma_client_or_none", lambda: MagicMock()) + monkeypatch.setattr(management, "get_mcp_server", AsyncMock(return_value=saved)) + payload = NewMCPServerRequest( + server_id=saved.server_id, auth_type="oauth2", transport="http", + url="https://new.example/mcp" if upstream_changed else saved.url, + ) + resolved = await management._resolve_session_server_id(payload) + assert (resolved != saved.server_id) is upstream_changed + + +def test_new_oauth_session_does_not_look_up_saved_credentials(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + + lookup = MagicMock() + monkeypatch.setattr(management.global_mcp_server_manager, "get_mcp_server_by_id", lookup) + payload = NewMCPServerRequest(url="https://new.example/mcp", auth_type="oauth2", transport="http") + assert management._inherit_credentials_from_existing_server(payload) is payload + assert not payload.credentials + lookup.assert_not_called() + + +def test_edit_does_not_rebind_resubmitted_saved_client_to_new_issuer(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-static", name="saved", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + client_id="saved-client", client_secret="saved-secret", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + payload = NewMCPServerRequest( + server_id=saved.server_id, transport="http", auth_type="oauth2", + url="https://new.example/mcp", issuer="https://new.example", + credentials={"client_id": saved.client_id, "client_secret": saved.client_secret}, + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert not (staged.credentials or {}).get("client_id") + assert not (staged.credentials or {}).get("client_secret") + assert saved.client_id == "saved-client" + + +@pytest.mark.parametrize("replacement", [ + {"client_secret": "replacement-secret"}, + {"client_secret": "old-secret", "token_endpoint_auth_method": "client_secret_basic"}, + {"client_secret": None}, + {"dcr_issuer": "https://new.example", "dcr_server_url": "https://new.example/mcp"}, +]) +def test_staged_issuer_edit_preserves_replacement_with_same_client_id(monkeypatch, replacement): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-static", name="saved", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + client_id="shared-client", client_secret="old-secret", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + submitted = {"client_id": "shared-client", **replacement} + payload = NewMCPServerRequest( + server_id=saved.server_id, transport="http", auth_type="oauth2", + url="https://new.example/mcp", issuer="https://new.example", credentials=submitted, + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert staged.credentials == submitted + assert saved.client_secret == "old-secret" diff --git a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py index 8159890ef16..452fdec648f 100644 --- a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py @@ -6167,6 +6167,53 @@ class TestStrategyRouterWriteValidation: async with _auto_router_capability_slot(fake, effective_params=effective_params, model_id=candidate_id) as table: assert hasattr(table, "create") + @pytest.mark.asyncio + async def test_slot_counts_tuned_routers_whose_stored_params_are_json_strings( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from fastapi import HTTPException + + from litellm.proxy import proxy_server + from litellm.proxy.management_endpoints.model_management_endpoints import _auto_router_capability_slot + from litellm.router_utils.auto_router_tuning_baseline import snapshot_tuning_baselines + + fake = self._FakeDb([]) + fake.tx_obj.litellm_proxymodeltable.find_many = AsyncMock( + return_value=[ + { + "model_id": "a", + "model_name": "router-a", + "litellm_params": json.dumps( + {"model": "auto_router/complexity_router", "complexity_router_config": self._TUNED_A_EDITED} + ), + "model_info": json.dumps({"id": "a"}), + } + ] + ) + monkeypatch.setattr(proxy_server._license_check, "auto_router_capability_limit", lambda: 1) + monkeypatch.setattr(proxy_server, "llm_router", None) + monkeypatch.setattr( + proxy_server, + "heuristic_v1_tuning_baselines", + snapshot_tuning_baselines( + [self._db_router_row("a", self._TUNED_A), self._db_router_row("b", self._TUNED_B)] + ), + ) + + with pytest.raises(HTTPException) as refused: + async with _auto_router_capability_slot( + fake, + effective_params={ + "model": "auto_router/complexity_router", + "complexity_router_config": self._TUNED_B_EDITED, + }, + model_id="b", + ): + pass + + assert refused.value.status_code == 403 + assert "changed heuristic scoring rules" in str(refused.value.detail) + @pytest.mark.asyncio async def test_add_new_model_refuses_a_second_tuned_heuristic_v1_router_without_a_model_id(self) -> None: """A create request carries no model_info at all, yet the quota still judges it: Deployment mints the diff --git a/tests/unit/proxy/management_endpoints/test_organization_endpoints.py b/tests/unit/proxy/management_endpoints/test_organization_endpoints.py index 440f93d1387..7d3bd049a03 100644 --- a/tests/unit/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_organization_endpoints.py @@ -1219,6 +1219,34 @@ async def test_legacy_update_without_budget_fields_skips_budget_write(monkeypatc assert prisma.db.litellm_organizationtable.update.await_args.kwargs["data"]["organization_alias"] == "renamed" +@pytest.mark.asyncio +async def test_legacy_update_writes_sent_metadata_to_the_organization_row(monkeypatch): + prisma = await _run_legacy_update_organization( + monkeypatch, + body={"organization_id": "org-1", "metadata": {"team": "search", "limits": {"rpm": 5}}}, + existing_budget_id="budget-1", + ) + + organization_write = prisma.db.litellm_organizationtable.update.await_args + assert organization_write.kwargs["where"] == {"organization_id": "org-1"} + assert json.loads(organization_write.kwargs["data"]["metadata"]) == {"team": "search", "limits": {"rpm": 5}} + prisma.db.litellm_budgettable.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("body", [[], ["litellm_budget_table"], "organization_id", 5, 1.5, True, None]) +async def test_legacy_update_rejects_json_body_that_is_not_an_object_before_reading_the_database(monkeypatch, body): + from pydantic import ValidationError + + from litellm.proxy import proxy_server + + with pytest.raises(ValidationError): + await _run_legacy_update_organization(monkeypatch, body=body, existing_budget_id="budget-1") + + proxy_server.prisma_client.db.litellm_organizationtable.find_unique.assert_not_awaited() + proxy_server.prisma_client.db.litellm_organizationtable.update.assert_not_awaited() + + def test_build_budget_write_data_recomputes_reset_at_on_duration(): """A sent budget_duration recomputes budget_reset_at so the reset window follows the new duration.""" from litellm.proxy.management_endpoints.organization_endpoints import build_budget_write_data diff --git a/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py index 14ee9db6ffd..d5dbd3df6a6 100644 --- a/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_tag_management_endpoints.py @@ -2,6 +2,7 @@ import inspect import json from collections.abc import Mapping, Sequence from contextlib import contextmanager +from datetime import datetime from types import MappingProxyType, SimpleNamespace from typing import Final, cast from unittest.mock import AsyncMock, Mock, patch @@ -1521,3 +1522,67 @@ async def test_add_tag_to_deployment_model_not_found(): assert exc_info.value.status_code == 500 assert "not found in database" in str(exc_info.value.detail) + + +class _StoredTagTable: + def __init__(self, model_info: object) -> None: + self.model_info = model_info + + async def find_many(self, where: object = None, include: object = None) -> list[SimpleNamespace]: + return [ + SimpleNamespace( + tag_name="routed-tag", + description="Routes to one model", + models=["model-1"], + model_info=self.model_info, + budget_id=None, + created_at=datetime(2025, 1, 1), + updated_at=datetime(2025, 1, 2), + created_by="user-123", + litellm_budget_table=None, + ) + ] + + +class _NoDynamicTagSpend: + async def group_by(self, by: object, where: object, min: object, max: object) -> list[object]: + return [] + + +@pytest.mark.parametrize( + ("stored_model_info", "returned_model_info"), + [ + ('{"model-1": "gpt-4o"}', {"model-1": "gpt-4o"}), + ({"model-1": "gpt-4o"}, {"model-1": "gpt-4o"}), + (None, {}), + ], +) +def test_tag_info_and_tag_list_return_the_stored_model_info_decoded( + monkeypatch, stored_model_info, returned_model_info +): + from litellm.proxy import proxy_server + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + monkeypatch.setattr( + proxy_server, + "prisma_client", + SimpleNamespace( + db=SimpleNamespace( + litellm_tagtable=_StoredTagTable(stored_model_info), + litellm_dailytagspend=_NoDynamicTagSpend(), + ) + ), + ) + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="test-user-123", user_role=LitellmUserRoles.PROXY_ADMIN + ) + try: + info_response = client.post("/tag/info", json={"names": ["routed-tag"]}) + list_response = client.get("/tag/list") + finally: + app.dependency_overrides.clear() + + assert info_response.status_code == 200 + assert info_response.json()["routed-tag"]["model_info"] == returned_model_info + assert list_response.status_code == 200 + assert [tag["model_info"] for tag in list_response.json()] == [returned_model_info] diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py index 094320225cd..1c9bd470e52 100644 --- a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py @@ -2533,3 +2533,22 @@ class TestRecordPartialUsageForFailure: assert "combined_usage_object" not in logging_obj.model_call_details assert "response_cost" not in logging_obj.model_call_details + + +@pytest.mark.parametrize( + ("all_chunks", "interrupted"), + [ + (['data: {"type": "content_block_delta"}'], True), + (['data: {"type": "message_delta"}', 'data: {"type": "message_stop"}'], False), + ([b'data: {"type": "message_delta"}\ndata: {"type": "message_stop"}\n'], False), + (['data: {"type": "message_delta"}', "data: [1, 2]", 'data: "text"', "data: 7", "data: null"], False), + (['data: {"type": "content_block_stop"}', "data: [1, 2]", "data: not json"], True), + (['data: {"type": "message_start"}', 'data: {"type": ["message_delta"]}', "data: {}"], True), + (["data: [1, 2]", "data: null", "event: message_delta"], True), + ([], True), + ], +) +def test_stream_was_interrupted_skips_data_lines_that_are_not_json_objects( + all_chunks: list[str | bytes], interrupted: bool +): + assert AnthropicPassthroughLoggingHandler._stream_was_interrupted(all_chunks) is interrupted diff --git a/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_vertex_passthrough_logging_handler.py b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_vertex_passthrough_logging_handler.py new file mode 100644 index 00000000000..f9df5a120f5 --- /dev/null +++ b/tests/unit/proxy/pass_through_endpoints/llm_provider_handlers/test_vertex_passthrough_logging_handler.py @@ -0,0 +1,95 @@ +from datetime import datetime + +import httpx +import pytest +from pydantic import ValidationError + +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.proxy._types import PassThroughEndpointLoggingTypedDict +from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import ( + VertexPassthroughLoggingHandler, +) +from litellm.types.utils import EmbeddingResponse, ModelResponse + +_PREDICT_ROUTE = "/v1/projects/p/locations/us-central1/publishers/google/models/text-embedding-004:predict" +_INTERACTIONS_ROUTE = "https://aiplatform.googleapis.com/v1beta1/projects/p/locations/global/interactions" + + +def _handle(url_route: str, payload: object) -> PassThroughEndpointLoggingTypedDict: + logging_obj = Logging( + model="unknown", + messages=[{"role": "user", "content": "hi"}], + stream=False, + call_type="pass_through_endpoint", + start_time=datetime(2026, 1, 1), + litellm_call_id="call-1", + function_id="fn-1", + ) + logging_obj.optional_params = {} + response = httpx.Response(200, json=payload) + return VertexPassthroughLoggingHandler.vertex_passthrough_handler( + httpx_response=response, + logging_obj=logging_obj, + url_route=url_route, + result=response.text, + start_time=datetime(2026, 1, 1), + end_time=datetime(2026, 1, 1), + cache_hit=False, + request_body={"model": "gemini-omni-flash-preview"}, + ) + + +def test_predict_response_with_text_embeddings_is_logged_as_an_embedding_response(): + result = _handle( + _PREDICT_ROUTE, + { + "predictions": [ + {"embeddings": {"values": [0.1, 0.2], "statistics": {"token_count": 3}}}, + {"embeddings": {"values": [0.3, 0.4], "statistics": {"token_count": 4}}}, + ], + "metadata": {"billableCharacterCount": 9}, + }, + ) + + response = result["result"] + assert isinstance(response, EmbeddingResponse) + assert response.data == [ + {"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}, + {"object": "embedding", "index": 1, "embedding": [0.3, 0.4]}, + ] + assert response.usage.prompt_tokens == 7 + assert result["kwargs"]["model"] == "text-embedding-004" + assert result["kwargs"]["custom_llm_provider"] == "vertex_ai" + + +@pytest.mark.parametrize("payload", [["not", "an", "object"], "text", 7, 1.5, True]) +def test_predict_response_that_is_not_a_json_object_is_rejected_without_echoing_it(payload: object): + with pytest.raises(ValidationError) as exc_info: + _handle(_PREDICT_ROUTE, payload) + + assert "input_value" not in str(exc_info.value) + + +def test_interactions_usage_object_is_read_into_prompt_and_completion_tokens(): + result = _handle( + _INTERACTIONS_ROUTE, + { + "id": "interactions/abc", + "model": "gemini-omni-flash-preview", + "usage": { + "total_tokens": 41, + "total_input_tokens": 12, + "input_tokens_by_modality": [{"modality": "text", "tokens": 12}], + "total_output_tokens": 9, + "output_tokens_by_modality": [{"modality": "text", "tokens": 9}], + "total_thought_tokens": 20, + }, + }, + ) + + response = result["result"] + assert isinstance(response, ModelResponse) + assert response.usage.prompt_tokens == 12 + assert response.usage.completion_tokens == 29 + assert response.usage.completion_tokens_details.text_tokens == 9 + assert result["kwargs"]["custom_llm_provider"] == "vertex_ai" diff --git a/tests/unit/proxy/policy_engine/test_policy_registry.py b/tests/unit/proxy/policy_engine/test_policy_registry.py new file mode 100644 index 00000000000..db7775b6238 --- /dev/null +++ b/tests/unit/proxy/policy_engine/test_policy_registry.py @@ -0,0 +1,154 @@ +import json +from collections.abc import Mapping +from datetime import datetime, timezone +from types import SimpleNamespace + +import pytest + +from litellm.proxy.policy_engine.policy_registry import PolicyRegistry +from litellm.types.proxy.policy_engine import PolicyCreateRequest, PolicyUpdateRequest + +_NOW = datetime(2026, 1, 1, tzinfo=timezone.utc) + +_REQUESTED_PIPELINE = {"mode": "pre_call", "steps": [{"guardrail": "pii-guard", "on_fail": "next"}]} + +_STORED_PIPELINE = { + "mode": "pre_call", + "steps": [ + { + "guardrail": "pii-guard", + "on_fail": "next", + "on_pass": "allow", + "on_error": None, + "pass_data": False, + "modify_response_message": None, + } + ], +} + +_MALFORMED_PIPELINES = [ + pytest.param({"mode": "pre_call", "steps": []}, "steps", id="no-steps"), + pytest.param({"mode": "pre_call"}, "steps", id="steps-missing"), + pytest.param({"steps": [{"guardrail": "pii-guard"}]}, "mode", id="mode-missing"), + pytest.param({"mode": "during_call", "steps": [{"guardrail": "pii-guard"}]}, "mode", id="unknown-mode"), + pytest.param( + {"mode": "pre_call", "steps": [{"guardrail": "pii-guard"}], "order": "strict"}, "order", id="unknown-key" + ), + pytest.param({"mode": "pre_call", "steps": "pii-guard"}, "steps", id="steps-not-a-list"), +] + + +class _PolicyTable: + def __init__(self) -> None: + self.writes: list[Mapping[str, object]] = [] + + def _row(self, data: Mapping[str, object]) -> SimpleNamespace: + pipeline = data.get("pipeline") + return SimpleNamespace( + policy_id="policy-1", + policy_name=data.get("policy_name", "pii-policy"), + version_number=1, + version_status=data.get("version_status", "draft"), + parent_version_id=None, + is_latest=True, + published_at=None, + production_at=None, + inherit=None, + description=None, + guardrails_add=data.get("guardrails_add", []), + guardrails_remove=data.get("guardrails_remove", []), + condition=None, + pipeline=json.loads(pipeline) if isinstance(pipeline, str) else None, + created_at=_NOW, + updated_at=_NOW, + created_by=None, + updated_by=None, + ) + + async def create(self, data: Mapping[str, object]) -> SimpleNamespace: + self.writes.append(data) + return self._row(data) + + async def find_unique(self, where: Mapping[str, object]) -> SimpleNamespace: + return self._row({"policy_name": "pii-policy", "version_status": "draft"}) + + async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> SimpleNamespace: + self.writes.append(data) + return self._row(data) + + +def _prisma_client(table: _PolicyTable) -> SimpleNamespace: + return SimpleNamespace(db=SimpleNamespace(litellm_policytable=table)) + + +@pytest.mark.asyncio +async def test_add_policy_to_db_stores_the_pipeline_with_step_defaults_filled_in(): + table = _PolicyTable() + registry = PolicyRegistry() + + response = await registry.add_policy_to_db( + PolicyCreateRequest(policy_name="pii-policy", guardrails_add=["pii-guard"], pipeline=_REQUESTED_PIPELINE), + _prisma_client(table), + ) + + assert [json.loads(str(write["pipeline"])) for write in table.writes] == [_STORED_PIPELINE] + assert response.pipeline == _STORED_PIPELINE + stored_policy = registry.get_policy("pii-policy") + assert stored_policy is not None + assert stored_policy.pipeline is not None + assert stored_policy.pipeline.model_dump() == _STORED_PIPELINE + + +@pytest.mark.asyncio +async def test_update_policy_in_db_stores_the_pipeline_with_step_defaults_filled_in(): + table = _PolicyTable() + + response = await PolicyRegistry().update_policy_in_db( + "policy-1", + PolicyUpdateRequest(pipeline=_REQUESTED_PIPELINE), + _prisma_client(table), + ) + + assert [json.loads(str(write["pipeline"])) for write in table.writes] == [_STORED_PIPELINE] + assert response.pipeline == _STORED_PIPELINE + + +@pytest.mark.parametrize(("pipeline", "rejected_field"), _MALFORMED_PIPELINES) +@pytest.mark.asyncio +async def test_add_policy_to_db_rejects_a_malformed_pipeline_before_writing( + pipeline: dict[str, object], rejected_field: str +): + table = _PolicyTable() + registry = PolicyRegistry() + + with pytest.raises(Exception, match="Error adding policy to DB: 1 validation error") as raised: + await registry.add_policy_to_db( + PolicyCreateRequest(policy_name="pii-policy", pipeline=pipeline), + _prisma_client(table), + ) + + assert str(raised.value).startswith( + f"Error adding policy to DB: 1 validation error for GuardrailPipeline\n{rejected_field}\n" + ) + assert table.writes == [] + assert registry.get_policy("pii-policy") is None + + +@pytest.mark.parametrize(("pipeline", "rejected_field"), _MALFORMED_PIPELINES) +@pytest.mark.asyncio +async def test_update_policy_in_db_rejects_a_malformed_pipeline_before_writing( + pipeline: dict[str, object], rejected_field: str +): + table = _PolicyTable() + + with pytest.raises(Exception, match="Error updating policy in DB: 1 validation error") as raised: + await PolicyRegistry().update_policy_in_db( + "policy-1", + PolicyUpdateRequest(pipeline=pipeline), + _prisma_client(table), + ) + + assert str(raised.value).startswith( + f"Error updating policy in DB: 1 validation error for GuardrailPipeline\n{rejected_field}\n" + ) + assert table.writes == [] diff --git a/tests/unit/proxy/proxy_server/test_routes_model_cost_map.py b/tests/unit/proxy/proxy_server/test_routes_model_cost_map.py index 40bd66ea91c..31f69d686d4 100644 --- a/tests/unit/proxy/proxy_server/test_routes_model_cost_map.py +++ b/tests/unit/proxy/proxy_server/test_routes_model_cost_map.py @@ -103,7 +103,7 @@ def test_reload_model_cost_map_surfaces_the_blob_id_of_the_bytes_served_on_every import httpx import litellm - from litellm.litellm_core_utils.get_model_cost_map import git_blob_id + from litellm.litellm_core_utils.get_model_cost_map import _finalize_model_cost_map, git_blob_id from litellm.proxy import proxy_server as ps from litellm.proxy._types import LitellmUserRoles @@ -142,7 +142,7 @@ def test_reload_model_cost_map_surfaces_the_blob_id_of_the_bytes_served_on_every assert {key: status_response.json()[key] for key in expected} == expected assert public_response.status_code == 200 assert "gpt-4o" in public_response.json() - assert reload_body["models_count"] == len(litellm.model_cost) + assert reload_body["models_count"] == len(_finalize_model_cost_map(json.loads(body))) def test_reload_model_cost_map_fetch_failure_502_keeps_map( diff --git a/tests/unit/proxy/spend_tracking/test_budget_reservation.py b/tests/unit/proxy/spend_tracking/test_budget_reservation.py index 9df8e6f4d67..024fe5229be 100644 --- a/tests/unit/proxy/spend_tracking/test_budget_reservation.py +++ b/tests/unit/proxy/spend_tracking/test_budget_reservation.py @@ -11,6 +11,8 @@ from litellm.caching import DualCache from litellm.models.budget import LiteLLM_BudgetTable from litellm.proxy import proxy_server from litellm.proxy._types import ( + LiteLLM_OrganizationTable, + LiteLLM_ProjectTableCachedObj, LiteLLM_TeamMembership, LiteLLM_TeamTable, LiteLLM_UserTable, @@ -18,11 +20,13 @@ from litellm.proxy._types import ( ) from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, + project_cache_key, team_membership_reservation_cache_key, ) from litellm.proxy.spend_tracking.budget_reservation import ( _get_team_member_budget_counter, estimate_request_max_cost, + get_budget_window_start, release_unbound_budget_reservation, reserve_budget_for_request, ) @@ -342,3 +346,83 @@ async def test_release_unbound_budget_reservation_leaves_a_bound_one_to_its_call assert spend_counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(reservation["reserved_cost"]) assert reservation["finalized"] is False + + +@pytest.mark.asyncio +async def test_reservation_holds_cost_against_cached_org_and_project_budgets(spend_counter_cache: DualCache): + cache: Final = UserApiKeyCache() + await cache.async_set_cache( + key="org_id:org-budgeted:with_budget", + value=LiteLLM_OrganizationTable( + organization_id="org-budgeted", + budget_id="org-budget", + spend=1.5, + models=[], + created_by="admin", + updated_by="admin", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=10.0), + ), + model_type=LiteLLM_OrganizationTable, + ) + await cache.async_set_cache( + key=project_cache_key("project-budgeted"), + value=LiteLLM_ProjectTableCachedObj( + project_id="project-budgeted", spend=2.5, litellm_budget_table=LiteLLM_BudgetTable(max_budget=20.0) + ), + model_type=LiteLLM_ProjectTableCachedObj, + ) + + reservation: Final = await reserve_budget_for_request( + request_body={"model": "gpt-4o", "input": "hello"}, + route="/v1/responses", + llm_router=None, + valid_token=UserAPIKeyAuth(token="hashed-org-project", org_id="org-budgeted", project_id="project-budgeted"), + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=cache, + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + ) + + assert reservation is not None + reserved_cost: Final = reservation["reserved_cost"] + assert reserved_cost > 0 + assert reservation["entries"] == [ + { + "counter_key": "spend:org:org-budgeted", + "entity_type": "Organization", + "entity_id": "org-budgeted", + "reserved_cost": reserved_cost, + "applied_adjustment": 0.0, + }, + { + "counter_key": "spend:project:project-budgeted", + "entity_type": "Project", + "entity_id": "project-budgeted", + "reserved_cost": reserved_cost, + "applied_adjustment": 0.0, + }, + ] + assert spend_counter_cache.in_memory_cache.get_cache(key="spend:org:org-budgeted") == pytest.approx(reserved_cost) + assert spend_counter_cache.in_memory_cache.get_cache(key="spend:project:project-budgeted") == pytest.approx( + reserved_cost + ) + + +@pytest.mark.parametrize( + ("window", "expected"), + [ + ( + '{"budget_duration": "1h", "reset_at": "2030-01-01T01:00:00Z"}', + datetime(2030, 1, 1, 0, 0, tzinfo=timezone.utc), + ), + ('{"reset_at": "2030-01-01T01:00:00Z"}', None), + ("{}", None), + ('["budget_duration", "1h"]', None), + ('"1h"', None), + ("null", None), + ("not json", None), + ], +) +def test_budget_window_start_reads_json_encoded_windows(window: str, expected: datetime | None) -> None: + assert get_budget_window_start(window) == expected diff --git a/tests/unit/proxy/test__lazy_features.py b/tests/unit/proxy/test__lazy_features.py index d2bb8244c2f..c3c5068b890 100644 --- a/tests/unit/proxy/test__lazy_features.py +++ b/tests/unit/proxy/test__lazy_features.py @@ -50,6 +50,33 @@ def _has_lazy_middleware(app: FastAPI) -> bool: return any(middleware.cls is LazyFeatureMiddleware for middleware in app.user_middleware) +@pytest.mark.parametrize("disable_lazy_routes", ("false", "true")) +def test_retired_roi_calculator_cannot_be_loaded_or_called( + monkeypatch: pytest.MonkeyPatch, disable_lazy_routes: str +) -> None: + monkeypatch.setenv(FLAG, disable_lazy_routes) + app: Final = FastAPI() + attach_lazy_features(app) + + with TestClient(app) as client: + for method, path in ( + ("POST", "/lazy/warm/roi_calculator"), + ("GET", "/roi-calculator/settings"), + ("GET", "/roi-calculator/report?mode=demo"), + ("POST", "/roi-calculator/sync"), + ("GET", "/roi-calculator/observed/report"), + ("GET", "/roi-calculator/observed/apps"), + ("POST", "/roi-calculator/observed/sync"), + ("POST", "/roi-calculator/observed/oauth/github/start"), + ("GET", "/roi-calculator/observed/oauth/github/callback"), + ("GET", "/roi-calculator/observed/oauth/github/installed"), + ("POST", "/roi-calculator/observed/oauth/gitlab/start"), + ("GET", "/roi-calculator/observed/oauth/gitlab/callback"), + ): + response: Final = client.request(method, path) + assert response.status_code == 404, f"{method} {path}: {response.text}" + + @pytest.mark.parametrize("value", ("1", "true", "TRUE", "yes", "on")) def test_flag_registers_every_feature_at_startup(monkeypatch: pytest.MonkeyPatch, value: str) -> None: monkeypatch.setenv(FLAG, value) diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index 8485c286a30..8dc6c0477b0 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -26,6 +26,7 @@ from litellm.constants import ( CLIENT_REQUESTED_MODEL_SCOPE_KEY, MAX_LITELLM_CALL_ID_LENGTH, RETURN_RAW_MODEL_NAME_METADATA_KEY, + STREAM_SSE_KEEPALIVE_PING_BYTES, ) from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.opentelemetry import UserAPIKeyAuth @@ -4256,6 +4257,72 @@ class TestStreamCloseOnDisconnect: assert upstream.aclosed + async def test_async_streaming_data_generator_closes_the_guardrail_chain_on_client_disconnect( + self, + ): + cleanup_ran = [] + + async def guarded_chain(**_kwargs): + try: + yield {"type": "chunk"} + yield {"type": "chunk"} + finally: + cleanup_ran.append(True) + + proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) + proxy_logging_obj.async_post_call_streaming_iterator_hook = guarded_chain + gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( + response=MagicMock(), + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + request_data={"model": "mock-model"}, + proxy_logging_obj=proxy_logging_obj, + serialize_chunk=lambda c: "data: x\n\n", + serialize_error=lambda e: "data: error\n\n", + ) + + await gen.__anext__() + await gen.aclose() + + assert cleanup_ran == [True] + + async def test_async_streaming_data_generator_refunds_the_budget_when_closing_the_guardrail_chain_raises( + self, + ): + async def guarded_chain(**_kwargs): + try: + yield STREAM_SSE_KEEPALIVE_PING_BYTES + yield {"type": "chunk"} + finally: + raise RuntimeError("cleanup failed") + + proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) + proxy_logging_obj.async_post_call_streaming_iterator_hook = guarded_chain + reservation = object() + user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + user_api_key_dict.budget_reservation = reservation + gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( + response=MagicMock(), + user_api_key_dict=user_api_key_dict, + request_data={"model": "mock-model"}, + proxy_logging_obj=proxy_logging_obj, + serialize_chunk=lambda c: "data: x\n\n", + serialize_error=lambda e: "data: error\n\n", + ) + + released: list[object] = [] + + async def record_release(budget_reservation: object) -> None: + released.append(budget_reservation) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation_on_cancel", + new=record_release, + ): + await gen.__anext__() + await gen.aclose() + + assert released == [reservation] + async def test_async_streaming_data_generator_redacts_internal_details_on_error( self, ): diff --git a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py index 42a9dc441dd..e2fbd5dcd98 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -1771,6 +1771,108 @@ async def test_proxy_startup_refuses_an_unsafe_master_key_even_when_the_database assert ("could not be checked" in announced[0]) == key_can_have_encrypted_the_database +@pytest.mark.asyncio +async def test_proxy_startup_refuses_fips_mode_when_this_python_does_not_enforce_fips(monkeypatch, tmp_path): + from fastapi import FastAPI + + from litellm.proxy.common_utils.fips import FipsModeError + from litellm.proxy.proxy_server import proxy_startup_event + + _, announced = _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: False) + monkeypatch.setenv("LITELLM_FIPS_MODE", "true") + + with pytest.raises(FipsModeError): + async with proxy_startup_event(FastAPI()): + pass + + assert len(announced) == 1 + assert "does not enforce FIPS" in announced[0] + + +@pytest.mark.asyncio +async def test_proxy_startup_refuses_fips_mode_when_the_config_disables_tls_verification(monkeypatch, tmp_path): + import yaml + from fastapi import FastAPI + + from litellm.proxy.common_utils.fips import FipsModeError + from litellm.proxy.proxy_server import proxy_startup_event + + config_path, announced = _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + config_path.write_text( + yaml.dump({"general_settings": {"master_key": "sk-a-safe-master-key"}, "litellm_settings": {"ssl_verify": False}}) + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: True) + monkeypatch.setattr(litellm, "ssl_verify", True) + monkeypatch.setenv("LITELLM_FIPS_MODE", "true") + + with pytest.raises(FipsModeError): + async with proxy_startup_event(FastAPI()): + pass + + assert "TLS certificate verification is disabled by litellm_settings.ssl_verify" in announced[0] + + +class _PrismaClientWhoseUserTableCannotHash: + class _Table: + async def find_many(self, where): + raise ValueError("[digital envelope routines] unsupported") + + class _Db: + litellm_usertable = None + + def __init__(self, database_url, proxy_logging_obj): + self.db = self._Db() + self.db.litellm_usertable = self._Table() + self.writer_db = self.db + + async def connect(self): + pass + + async def disconnect(self): + pass + + def start_view_setup_task(self): + pass + + async def check_view_exists(self): + pass + + async def health_check(self): + pass + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fips_mode", ["true", "false"]) +async def test_proxy_startup_surfaces_a_password_migration_crypto_failure(monkeypatch, tmp_path, caplog, fips_mode): + from fastapi import FastAPI + + from litellm.proxy.proxy_server import proxy_startup_event + + _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setenv("DATABASE_URL", "postgresql://nobody:nothing@127.0.0.1:1/unreachable") + monkeypatch.setattr("litellm.proxy.proxy_server.PrismaClient", _PrismaClientWhoseUserTableCannotHash) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: True) + monkeypatch.setenv("LITELLM_FIPS_MODE", fips_mode) + + with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"): + if fips_mode == "true": + with pytest.raises(ValueError, match="digital envelope routines"): + async with proxy_startup_event(FastAPI()): + pass + else: + async with proxy_startup_event(FastAPI()): + await asyncio.sleep(0) + + failures = [r.getMessage() for r in caplog.records if "Password migration failed" in r.getMessage()] + assert len(failures) == 1 + assert "plaintext passwords stay unhashed" in failures[0] + assert "digital envelope routines" in failures[0] + + class _DatabaseWithOneStoredCredential: def __init__(self, ciphertext): self._ciphertext = ciphertext @@ -7323,6 +7425,69 @@ async def test_async_data_generator_cleanup_on_early_exit(): mock_response.aclose.assert_awaited_once() +def _guarded_chain_logging(chain): + from litellm.proxy.utils import ProxyLogging + + proxy_logging = MagicMock(spec=ProxyLogging) + proxy_logging.async_post_call_streaming_iterator_hook = chain + proxy_logging.async_post_call_streaming_hook = AsyncMock(side_effect=lambda **kwargs: kwargs.get("response")) + proxy_logging.post_call_failure_hook = AsyncMock() + return proxy_logging + + +@pytest.mark.asyncio +async def test_async_data_generator_closes_the_guardrail_chain_before_returning_on_client_disconnect(): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + + cleanup_ran = [] + + async def guarded_chain(**_kwargs): + try: + yield {"choices": [{"delta": {"content": "Hello"}}]} + yield {"choices": [{"delta": {"content": " world"}}]} + finally: + cleanup_ran.append(True) + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", _guarded_chain_logging(guarded_chain)): + gen = async_data_generator(MagicMock(), MagicMock(spec=UserAPIKeyAuth), {"model": "gpt-4o-mini"}) + first_chunk = await gen.__anext__() + await gen.aclose() + + assert first_chunk.startswith("data: ") + assert cleanup_ran == [True] + + +@pytest.mark.asyncio +async def test_async_data_generator_closes_the_guardrail_chain_while_a_keepalive_read_is_pending(): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + + cleanup_ran = [] + never_arrives = asyncio.Event() + + async def guarded_chain(**_kwargs): + try: + yield {"choices": [{"delta": {"content": "Hello"}}]} + await never_arrives.wait() + yield {"choices": [{"delta": {"content": " world"}}]} + finally: + cleanup_ran.append(True) + + with ( + patch.object(litellm, "sse_keepalive_ping_interval_seconds", 1.0), + patch("litellm.proxy.proxy_server.proxy_logging_obj", _guarded_chain_logging(guarded_chain)), + ): + gen = async_data_generator(MagicMock(), MagicMock(spec=UserAPIKeyAuth), {"model": "gpt-4o-mini"}) + first_chunk = await gen.__anext__() + heartbeat = await gen.__anext__() + await gen.aclose() + + assert first_chunk.startswith("data: ") + assert heartbeat == ": ping\n\n" + assert cleanup_ran == [True] + + @pytest.mark.asyncio async def test_async_data_generator_uses_direct_stream_fast_path_without_callbacks(): """ diff --git a/tests/unit/proxy/test_tracing_endpoints.py b/tests/unit/proxy/test_tracing_endpoints.py index 516a0415545..11d63ad545c 100644 --- a/tests/unit/proxy/test_tracing_endpoints.py +++ b/tests/unit/proxy/test_tracing_endpoints.py @@ -5,12 +5,15 @@ Tests for the agent tracing endpoints (litellm/proxy/tracing_endpoints.py). from collections.abc import AsyncGenerator, Mapping from contextlib import asynccontextmanager from types import ModuleType -from typing import Final, Literal -from unittest.mock import AsyncMock, MagicMock +from typing import Final, Literal, TypedDict +from unittest.mock import AsyncMock, MagicMock, call import pytest from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient +from httpx import Response +from pydantic import TypeAdapter +from typing_extensions import ReadOnly from litellm.constants import TRACE_READ_RETRY_AFTER_SECONDS from litellm.proxy import tracing_endpoints @@ -97,6 +100,29 @@ SPAN_DETAIL_RESPONSE: Final = { "output_ui": {"kind": "text", "text": ""}, "attributes": {}, } +SPAN_ERROR_RESPONSE: Final = { + "span_id": "s1", + "message": "span error", + "total_chars": 10, + "next_cursor": None, +} +NOW_MS: Final = 1_800_000_000_000 + + +class RequestValidationError(TypedDict): + type: ReadOnly[str] + loc: ReadOnly[list[str | int]] + + +def _validation_errors(response: Response) -> tuple[RequestValidationError, ...]: + return tuple(TypeAdapter(list[RequestValidationError]).validate_python(response.json()["detail"])) + + +def _assert_validation_error(response: Response, error_type: str, location: tuple[str | int, ...]) -> None: + assert response.status_code == 422, response.text + assert any( + error["type"] == error_type and tuple(error["loc"]) == location for error in _validation_errors(response) + ) @pytest.mark.parametrize( @@ -262,6 +288,43 @@ def test_list_traces_defaults_to_last_24h(client, receiver): assert kwargs["cursor"] is None +@pytest.mark.parametrize( + ("params", "expected_start_ms", "expected_end_ms"), + ( + ({}, NOW_MS - tracing_endpoints.MS_PER_DAY, NOW_MS), + ({"start_ms": 123}, 123, NOW_MS), + ({"end_ms": -7}, NOW_MS - tracing_endpoints.MS_PER_DAY, -7), + ), +) +def test_list_traces_resolves_default_bounds_from_injected_clock( + client: TestClient, + receiver: MagicMock, + params: Mapping[str, int], + expected_start_ms: int, + expected_end_ms: int, +) -> None: + client.app.dependency_overrides[tracing_endpoints.current_time_ms] = lambda: NOW_MS + response: Final = client.get("/v1/traces", params=params) + assert response.status_code == 200, response.text + receiver.list_traces.assert_awaited_once_with( + scope={"all_teams": 0, "user_id": "user", "team_ids": ()}, + start_ms=expected_start_ms, + end_ms=expected_end_ms, + cursor=None, + ) + + +def test_list_traces_forwards_large_and_negative_bounds_unchanged(client: TestClient, receiver: MagicMock) -> None: + response: Final = client.get("/v1/traces", params={"start_ms": 2**63, "end_ms": -1, "cursor": "next"}) + assert response.status_code == 200, response.text + receiver.list_traces.assert_awaited_once_with( + scope={"all_teams": 0, "user_id": "user", "team_ids": ()}, + start_ms=2**63, + end_ms=-1, + cursor="next", + ) + + def test_get_trace_404_and_200(client, receiver): assert client.get("/v1/traces/missing").status_code == 404 receiver.get_trace.return_value = TRACE_RESPONSE @@ -289,6 +352,125 @@ def test_trace_detail_passes_scoped_reference(client, receiver, suffix, cursor, ) +@pytest.mark.parametrize("page_size", (1, 500)) +def test_trace_detail_accepts_page_size_bounds(client: TestClient, receiver: MagicMock, page_size: int) -> None: + receiver.get_trace.return_value = TRACE_RESPONSE + response: Final = client.get("/v1/traces/t1", params={"trace_ref": "run-one", "page_size": page_size}) + assert response.status_code == 200, response.text + receiver.get_trace.assert_awaited_once_with( + "t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one", None, page_size + ) + + +@pytest.mark.parametrize( + ("value", "error_type"), + (("0", "greater_than_equal"), ("501", "less_than_equal"), ("abc", "int_parsing")), +) +def test_trace_detail_reports_page_size_validation( + client: TestClient, receiver: MagicMock, value: str, error_type: str +) -> None: + response: Final = client.get("/v1/traces/t1", params={"page_size": value}) + _assert_validation_error(response, error_type, ("query", "page_size")) + receiver.get_trace.assert_not_awaited() + + +def test_trace_read_routes_accept_and_forward_512_character_cursors( + client: TestClient, receiver: MagicMock +) -> None: + cursor: Final = "x" * 512 + receiver.get_trace.return_value = TRACE_RESPONSE + receiver.get_span_error = AsyncMock(return_value=SPAN_ERROR_RESPONSE) + client.app.dependency_overrides[tracing_endpoints.current_time_ms] = lambda: NOW_MS + + list_response: Final = client.get("/v1/traces", params={"cursor": cursor}) + detail_response: Final = client.get("/v1/traces/t1", params={"trace_ref": "run-one", "cursor": cursor}) + error_response: Final = client.get( + "/v1/traces/t1/spans/s1/error", params={"trace_ref": "run-one", "cursor": cursor} + ) + + assert list_response.status_code == 200, list_response.text + assert detail_response.status_code == 200, detail_response.text + assert error_response.status_code == 200, error_response.text + receiver.list_traces.assert_awaited_once_with( + scope={"all_teams": 0, "user_id": "user", "team_ids": ()}, + start_ms=NOW_MS - tracing_endpoints.MS_PER_DAY, + end_ms=NOW_MS, + cursor=cursor, + ) + receiver.get_trace.assert_awaited_once_with( + "t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one", cursor, None + ) + receiver.get_span_error.assert_awaited_once_with( + "t1", "s1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one", cursor + ) + + +@pytest.mark.parametrize( + "path", + ("/v1/traces", "/v1/traces/t1", "/v1/traces/t1/spans/s1/error"), +) +def test_trace_read_routes_reject_513_character_cursors( + client: TestClient, receiver: MagicMock, path: str +) -> None: + response: Final = client.get(path, params={"cursor": "x" * 513}) + _assert_validation_error(response, "string_too_long", ("query", "cursor")) + + +def test_trace_read_routes_ignore_unknown_query_parameters(client: TestClient, receiver: MagicMock) -> None: + receiver.get_trace.return_value = TRACE_RESPONSE + receiver.get_span.return_value = SPAN_DETAIL_RESPONSE + receiver.get_span_error = AsyncMock(return_value=SPAN_ERROR_RESPONSE) + + list_params: Final = {"start_ms": 1, "end_ms": 2, "cursor": "list-cursor"} + list_response: Final = client.get("/v1/traces", params=list_params) + list_unknown_response: Final = client.get("/v1/traces", params={**list_params, "foo": "bar"}) + detail_params: Final = {"trace_ref": "run-one", "cursor": "detail-cursor", "page_size": 10} + detail_response: Final = client.get("/v1/traces/t1", params=detail_params) + detail_unknown_response: Final = client.get("/v1/traces/t1", params={**detail_params, "foo": "bar"}) + span_response: Final = client.get("/v1/traces/t1/spans/s1", params={"trace_ref": "run-one"}) + span_unknown_response: Final = client.get( + "/v1/traces/t1/spans/s1", params={"trace_ref": "run-one", "foo": "bar"} + ) + error_params: Final = {"trace_ref": "run-one", "cursor": "error-cursor"} + error_response: Final = client.get("/v1/traces/t1/spans/s1/error", params=error_params) + error_unknown_response: Final = client.get( + "/v1/traces/t1/spans/s1/error", params={**error_params, "foo": "bar"} + ) + + assert list_response.status_code == 200, list_response.text + assert list_unknown_response.status_code == 200, list_unknown_response.text + assert detail_response.status_code == 200, detail_response.text + assert detail_unknown_response.status_code == 200, detail_unknown_response.text + assert span_response.status_code == 200, span_response.text + assert span_unknown_response.status_code == 200, span_unknown_response.text + assert error_response.status_code == 200, error_response.text + assert error_unknown_response.status_code == 200, error_unknown_response.text + receiver.list_traces.assert_has_awaits( + ( + call(scope={"all_teams": 0, "user_id": "user", "team_ids": ()}, start_ms=1, end_ms=2, cursor="list-cursor"), + call(scope={"all_teams": 0, "user_id": "user", "team_ids": ()}, start_ms=1, end_ms=2, cursor="list-cursor"), + ) + ) + receiver.get_trace.assert_has_awaits( + ( + call("t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one", "detail-cursor", 10), + call("t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one", "detail-cursor", 10), + ) + ) + receiver.get_span.assert_has_awaits( + ( + call("t1", "s1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one"), + call("t1", "s1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one"), + ) + ) + receiver.get_span_error.assert_has_awaits( + ( + call("t1", "s1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one", "error-cursor"), + call("t1", "s1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one", "error-cursor"), + ) + ) + + @pytest.mark.parametrize( "path,method", ( @@ -606,6 +788,8 @@ def test_sql_and_help_use_authenticated_scope( result: Final = client.post("/v1/traces/query", json={"sql": "SELECT * FROM otel_traces"}) assert result.status_code == 200, result.text assert result.json() == SQL_ENVELOPE + assert type(result.json()["data"][0]["value"]) is str + assert type(result.json()["statistics"]["elapsed"]) is float receiver.storage.query_sql.assert_awaited_once_with("SELECT * FROM otel_traces", expected_scope, "test-secret") help_result: Final = client.get("/v1/traces/query/help") assert help_result.status_code == 200, help_result.text @@ -616,6 +800,29 @@ def test_sql_and_help_use_authenticated_scope( assert receiver.storage.query_sql.await_count == 1 +@pytest.mark.parametrize( + ("body", "error_type", "location"), + ( + (b"{}", "missing", ("body", "sql")), + (b'{"sql": null}', "string_type", ("body", "sql")), + (b'{"sql": 1}', "string_type", ("body", "sql")), + (b'{"sql": "SELECT 1", "extra": true}', "extra_forbidden", ("body", "extra")), + (b"{", "json_invalid", ("body", 1)), + (b"[]", "model_attributes_type", ("body",)), + ), +) +def test_sql_query_rejects_invalid_request_bodies( + client: TestClient, receiver: MagicMock, body: bytes, error_type: str, location: tuple[str | int, ...] +) -> None: + client.app.dependency_overrides[tracing_endpoints.provide_trace_query_secret] = lambda: "test-secret" + receiver.storage.query_sql = AsyncMock() + response: Final = client.post( + "/v1/traces/query", content=body, headers={"content-type": "application/json"} + ) + _assert_validation_error(response, error_type, location) + receiver.storage.query_sql.assert_not_awaited() + + @pytest.mark.parametrize("auth", (UserAPIKeyAuth(), UserAPIKeyAuth(team_id="a", project_id="p"))) def test_sql_rejects_missing_identity_without_querying( client: TestClient, receiver: MagicMock, auth: UserAPIKeyAuth diff --git a/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py b/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py index 8bd9dc0df8a..dbb354526d8 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py +++ b/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py @@ -11,9 +11,9 @@ from __future__ import annotations import asyncio import json +from collections.abc import AsyncGenerator, Iterator from copy import deepcopy import logging -from collections.abc import Iterator from typing import Any, Callable, Dict, List from unittest.mock import AsyncMock, MagicMock, patch @@ -26,6 +26,7 @@ from litellm.integrations.custom_guardrail import ( CustomGuardrail, ModifyResponseException, ) +from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.prometheus import PrometheusLogger from litellm.llms.base_llm.guardrail_translation.utils import stream_item_field from litellm.proxy._types import UserAPIKeyAuth @@ -2933,3 +2934,62 @@ async def test_streaming_iterator_hook_pipeline_discards_dropped_tool_call_on_re assert delivered == _responses_function_call_events() assert any("'gr-post'" in message and "discarded" in message for message in _warnings(caplog)) + + +class _RaisingAcloseIterator: + """Non-generator async iterator whose aclose raises after the stream ended.""" + + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + self._response = response + + def __aiter__(self) -> "_RaisingAcloseIterator": + return self + + async def __anext__(self) -> Any: + return await self._response.__anext__() + + async def aclose(self) -> None: + raise RuntimeError("cleanup failed") + + +class RaisingAcloseCallback(CustomLogger): + """Iterator hook returning a non-generator async iterator whose aclose raises.""" + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _RaisingAcloseIterator: + return _RaisingAcloseIterator(response) + + +@pytest.mark.asyncio +async def test_streaming_iterator_hook_pipeline_releases_buffered_content_when_a_callback_aclose_raises( + proxy_logging: ProxyLogging, + make_user_api_key_auth: Callable[..., UserAPIKeyAuth], + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + monkeypatch.setattr( + litellm, "callbacks", [_rewriting_stream_guardrail(lambda inputs: {}), RaisingAcloseCallback()] + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + data = _post_call_pipeline_data(stream=True) + chunks = _tool_call_stream_chunks() + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + delivered = [ + item + async for item in proxy_logging.async_post_call_streaming_iterator_hook( + user_api_key_dict=make_user_api_key_auth(request_route="/v1/chat/completions"), + response=_async_chunk_iter(chunks), + request_data=data, + ) + ] + + assert [chunk.model_dump() for chunk in delivered] == [chunk.model_dump() for chunk in chunks] + assert any( + "RaisingAcloseCallback" in message and "RuntimeError" in message and "cleanup failed" not in message + for message in _warnings(caplog) + ) diff --git a/tests/unit/rag/ingestion/test_base_ingestion.py b/tests/unit/rag/ingestion/test_base_ingestion.py new file mode 100644 index 00000000000..d20b46d356b --- /dev/null +++ b/tests/unit/rag/ingestion/test_base_ingestion.py @@ -0,0 +1,101 @@ +from collections import UserString +from typing import Final + +import httpx +import pytest +import respx + +import litellm +from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion +from litellm.types.utils import CredentialItem + +_FILE_URL: Final = "https://files.example/docs/report.pdf" +_STORED_CREDENTIAL: Final = CredentialItem( + credential_name="gemini-prod", + credential_info={}, + credential_values={"api_key": "stored-key"}, +) + + +def _vector_store_options(litellm_credential_name: object) -> dict[str, object]: + return { + "custom_llm_provider": "gemini", + "litellm_credential_name": litellm_credential_name, + "api_key": "caller-key", + "api_base": "https://caller.example", + } + + +def test_a_stored_credential_named_by_the_vector_store_replaces_the_caller_supplied_values( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr(litellm, "credential_list", [_STORED_CREDENTIAL]) + vector_store: Final = _vector_store_options("gemini-prod") + + ingestion: Final = GeminiRAGIngestion(ingest_options={"vector_store": vector_store}) + + assert ingestion.vector_store_config is vector_store + assert vector_store == { + "custom_llm_provider": "gemini", + "litellm_credential_name": "gemini-prod", + "api_key": "stored-key", + } + + +@pytest.mark.parametrize( + "litellm_credential_name", + ["gemini-staging", "", None, 7, True, ["gemini-prod"], {"credential_name": "gemini-prod"}], +) +def test_a_credential_name_that_matches_no_stored_credential_leaves_the_vector_store_config_alone( + monkeypatch: pytest.MonkeyPatch, litellm_credential_name: object +): + monkeypatch.setattr(litellm, "credential_list", [_STORED_CREDENTIAL]) + + ingestion: Final = GeminiRAGIngestion( + ingest_options={"vector_store": _vector_store_options(litellm_credential_name)} + ) + + assert ingestion.vector_store_config == _vector_store_options(litellm_credential_name) + + +def test_a_credential_name_that_is_not_a_string_is_not_resolved_even_when_it_equals_a_stored_name( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr(litellm, "credential_list", [_STORED_CREDENTIAL]) + litellm_credential_name: Final = UserString("gemini-prod") + + ingestion: Final = GeminiRAGIngestion( + ingest_options={"vector_store": _vector_store_options(litellm_credential_name)} + ) + + assert litellm_credential_name == "gemini-prod" + assert ingestion.vector_store_config == _vector_store_options(litellm_credential_name) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("headers", "expected_content_type"), + [ + ([("content-type", "application/pdf")], "application/pdf"), + ([("Content-Type", "text/plain; charset=utf-8")], "text/plain; charset=utf-8"), + ([("content-type", "")], ""), + ([("content-type", "text/plain"), ("content-type", "text/html")], "text/plain, text/html"), + ([], "application/octet-stream"), + ([("content-length", "8")], "application/octet-stream"), + ], +) +async def test_upload_from_a_url_takes_the_content_type_from_the_response_header( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, + headers: list[tuple[str, str]], + expected_content_type: str, +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(litellm, "user_url_validation", False) + litellm.in_memory_llm_clients_cache.flush_cache() + respx_mock.get(_FILE_URL).mock(return_value=httpx.Response(200, content=b"%PDF-1.7", headers=headers)) + ingestion: Final = GeminiRAGIngestion(ingest_options={"vector_store": {"custom_llm_provider": "gemini"}}) + + uploaded: Final = await ingestion.upload(file_url=_FILE_URL) + + assert uploaded == ("report.pdf", b"%PDF-1.7", expected_content_type, None) diff --git a/tests/unit/rag/ingestion/test_gemini_ingestion.py b/tests/unit/rag/ingestion/test_gemini_ingestion.py new file mode 100644 index 00000000000..8122a9d334b --- /dev/null +++ b/tests/unit/rag/ingestion/test_gemini_ingestion.py @@ -0,0 +1,111 @@ +import json +from collections.abc import Mapping +from types import MappingProxyType +from typing import Final + +import httpx +import pytest +import respx +from pydantic import ValidationError + +import litellm +from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion + +_API_BASE: Final = "https://gemini.example" +_STORE: Final = "fileSearchStores/docs-1" +_START_UPLOAD_URL: Final = f"{_API_BASE}/upload/v1beta/{_STORE}:uploadToFileSearchStore" +_UPLOAD_SESSION_URL: Final = "https://gemini.example/upload/session-1" +_DOCUMENT: Final = "fileSearchStores/docs-1/documents/notes-1" + + +def _ingestion(chunking_strategy: object) -> GeminiRAGIngestion: + return GeminiRAGIngestion( + ingest_options={ + "chunking_strategy": chunking_strategy, + "vector_store": { + "custom_llm_provider": "gemini", + "vector_store_id": _STORE, + "api_key": "test-key", + "api_base": _API_BASE, + }, + } + ) + + +async def _store_notes(ingestion: GeminiRAGIngestion) -> tuple[str | None, str | None]: + return await ingestion.store( + file_content=b"first note", + filename="notes.txt", + content_type="text/plain", + chunks=[], + embeddings=None, + ) + + +@pytest.fixture +def start_upload(monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter) -> respx.Route: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + respx_mock.put(_UPLOAD_SESSION_URL).mock(return_value=httpx.Response(200, json={"name": _DOCUMENT})) + return respx_mock.post(_START_UPLOAD_URL).mock( + return_value=httpx.Response(200, headers={"x-goog-upload-url": _UPLOAD_SESSION_URL}) + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("white_space_config", "expected"), + [ + ( + {"max_tokens_per_chunk": 200, "max_overlap_tokens": 20}, + {"maxTokensPerChunk": 200, "maxOverlapTokens": 20}, + ), + ({"max_tokens_per_chunk": 200}, {"maxTokensPerChunk": 200, "maxOverlapTokens": 400}), + ({"unrelated": True}, {"maxTokensPerChunk": 800, "maxOverlapTokens": 400}), + ( + {"max_tokens_per_chunk": "200", "max_overlap_tokens": None}, + {"maxTokensPerChunk": "200", "maxOverlapTokens": None}, + ), + ({7: "not a field", "max_overlap_tokens": 1.5}, {"maxTokensPerChunk": 800, "maxOverlapTokens": 1.5}), + (MappingProxyType({"max_overlap_tokens": 0}), {"maxTokensPerChunk": 800, "maxOverlapTokens": 0}), + ], +) +async def test_white_space_config_is_sent_as_the_chunking_config_of_the_upload( + start_upload: respx.Route, white_space_config: Mapping[object, object], expected: Mapping[str, object] +): + stored: Final = await _store_notes(_ingestion({"white_space_config": white_space_config})) + + assert stored == (_STORE, _DOCUMENT) + assert json.loads(start_upload.calls.last.request.content) == { + "displayName": "notes.txt", + "chunkingConfig": {"whiteSpaceConfig": expected}, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "chunking_strategy", + [None, {"type": "auto"}, {"white_space_config": None}, {"white_space_config": {}}, {"white_space_config": 0}], +) +async def test_upload_without_a_white_space_config_sends_no_chunking_config( + start_upload: respx.Route, chunking_strategy: object +): + stored: Final = await _store_notes(_ingestion(chunking_strategy)) + + assert stored == (_STORE, _DOCUMENT) + assert json.loads(start_upload.calls.last.request.content) == {"displayName": "notes.txt"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "white_space_config", + ["private-setting", [800, 400], [("max_tokens_per_chunk", 200)], 800, True, 1.5], +) +async def test_a_white_space_config_that_is_not_a_mapping_is_rejected_before_any_upload( + start_upload: respx.Route, white_space_config: object +): + with pytest.raises(ValidationError) as raised: + await _store_notes(_ingestion({"white_space_config": white_space_config})) + + assert "private-setting" not in str(raised.value) + assert not start_upload.called diff --git a/tests/unit/rag/ingestion/test_openai_ingestion.py b/tests/unit/rag/ingestion/test_openai_ingestion.py new file mode 100644 index 00000000000..74b68c9a517 --- /dev/null +++ b/tests/unit/rag/ingestion/test_openai_ingestion.py @@ -0,0 +1,101 @@ +import json +from typing import Final + +import httpx +import pytest +import respx + +import litellm +from litellm.rag.ingestion.openai_ingestion import OpenAIRAGIngestion + +_API_BASE: Final = "https://openai.example/v1" +_STATIC_CHUNKING: Final = {"type": "static", "static": {"max_chunk_size_tokens": 800, "chunk_overlap_tokens": 400}} +_ATTACHED_FILE: Final = { + "id": "file-1", + "object": "vector_store.file", + "created_at": 1767323045, + "vector_store_id": "vs_1", + "status": "completed", + "usage_bytes": 10, +} +_UPLOADED_FILE: Final = { + "id": "file-1", + "object": "file", + "bytes": 10, + "created_at": 1767323045, + "filename": "notes.txt", + "purpose": "assistants", + "status": "processed", +} + + +def _ingestion(chunking_strategy: object) -> OpenAIRAGIngestion: + return OpenAIRAGIngestion( + ingest_options={ + "chunking_strategy": chunking_strategy, + "vector_store": { + "custom_llm_provider": "openai", + "vector_store_id": "vs_1", + "api_key": "sk-test", + "api_base": _API_BASE, + }, + } + ) + + +@pytest.fixture +def attach_file(monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter) -> respx.Route: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + return respx_mock.post(f"{_API_BASE}/vector_stores/vs_1/files").mock( + return_value=httpx.Response(200, json=_ATTACHED_FILE) + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("chunking_strategy", "sent_chunking_strategy"), + [(None, {"type": "auto"}), ({"type": "auto"}, {"type": "auto"}), (_STATIC_CHUNKING, _STATIC_CHUNKING)], +) +async def test_an_uploaded_file_is_attached_to_the_vector_store_with_the_chunking_strategy( + attach_file: respx.Route, respx_mock: respx.MockRouter, chunking_strategy: object, sent_chunking_strategy: object +): + respx_mock.post(f"{_API_BASE}/files").mock(return_value=httpx.Response(200, json=_UPLOADED_FILE)) + + stored: Final = await _ingestion(chunking_strategy).store( + file_content=b"first note", + filename="notes.txt", + content_type="text/plain", + chunks=[], + embeddings=None, + ) + + assert stored == ("vs_1", "file-1") + assert json.loads(attach_file.calls.last.request.content) == { + "file_id": "file-1", + "chunking_strategy": sent_chunking_strategy, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("chunking_strategy", "sent_chunking_strategy"), + [(None, {"type": "auto"}), (_STATIC_CHUNKING, _STATIC_CHUNKING)], +) +async def test_an_existing_file_is_attached_to_the_vector_store_with_the_chunking_strategy( + attach_file: respx.Route, chunking_strategy: object, sent_chunking_strategy: object +): + stored: Final = await _ingestion(chunking_strategy).store( + file_content=None, + filename=None, + content_type=None, + chunks=[], + embeddings=None, + existing_file_id="file-9", + ) + + assert stored == ("vs_1", "file-9") + assert json.loads(attach_file.calls.last.request.content) == { + "file_id": "file-9", + "chunking_strategy": sent_chunking_strategy, + } diff --git a/tests/unit/rag/test_main.py b/tests/unit/rag/test_main.py index 54748efb480..81318a1e113 100644 --- a/tests/unit/rag/test_main.py +++ b/tests/unit/rag/test_main.py @@ -12,12 +12,14 @@ aquery carries the completion response with real usage and cost. import asyncio import json +from types import MappingProxyType from typing import Final from unittest.mock import patch import httpx import pytest import respx +from pydantic import ValidationError import litellm from litellm._internal_context import is_internal_call @@ -538,6 +540,91 @@ async def test_aquery_forwards_vector_store_params_to_search_but_not_completion( assert not ({"milvus_text_field", "outputFields"} & set(completion_kwargs)) +_UNDECODABLE_FILE: Final = {"filename": "notes.txt", "content": "x"} + + +@pytest.mark.parametrize( + ("ingest_options", "expected_provider"), + [ + ({"vector_store": {"custom_llm_provider": "bedrock"}}, "bedrock"), + ({"vector_store": MappingProxyType({"custom_llm_provider": "bedrock"})}, "bedrock"), + ({"vector_store": {"vector_store_id": "vs_1", 7: "ignored"}}, None), + ({}, None), + ], +) +def test_ingest_failure_is_attributed_to_the_vector_store_provider( + ingest_options: dict[str, object], expected_provider: str | None +) -> None: + with pytest.raises(litellm.APIConnectionError, match="Invalid base64-encoded string") as raised: + litellm.ingest(ingest_options=ingest_options, file=_UNDECODABLE_FILE) + + assert raised.value.llm_provider == expected_provider + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("ingest_options", "expected_provider"), + [ + ({"vector_store": {"custom_llm_provider": "bedrock"}}, "bedrock"), + ({"vector_store": MappingProxyType({"custom_llm_provider": "bedrock"})}, "bedrock"), + ({"vector_store": {"vector_store_id": "vs_1", 7: "ignored"}}, None), + ({}, None), + ], +) +async def test_aingest_failure_is_attributed_to_the_vector_store_provider( + ingest_options: dict[str, object], expected_provider: str | None +) -> None: + with pytest.raises(litellm.APIConnectionError, match="Invalid base64-encoded string") as raised: + await litellm.aingest(ingest_options=ingest_options, file=_UNDECODABLE_FILE) + + assert raised.value.llm_provider == expected_provider + + +@pytest.mark.parametrize("vector_store", [None, "openai", ["openai"], [{"api_key": "sk-test"}]]) +def test_ingest_failure_with_a_vector_store_that_is_not_a_mapping_raises_a_validation_error( + vector_store: object, +) -> None: + with pytest.raises(ValidationError) as raised: + litellm.ingest(ingest_options={"vector_store": vector_store}, file=_UNDECODABLE_FILE) + + assert "sk-test" not in str(raised.value) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("vector_store", [None, "openai", ["openai"], [{"api_key": "sk-test"}]]) +async def test_aingest_failure_with_a_vector_store_that_is_not_a_mapping_raises_a_validation_error( + vector_store: object, +) -> None: + with pytest.raises(ValidationError) as raised: + await litellm.aingest(ingest_options={"vector_store": vector_store}, file=_UNDECODABLE_FILE) + + assert "sk-test" not in str(raised.value) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider", [None, 5]) +async def test_aquery_with_a_provider_that_is_not_a_string_bills_only_the_completion(provider: object) -> None: + messages: Final = [{"role": "user", "content": "hello"}] + default_provider_response: Final = await litellm.aquery( + model="gpt-4o-mini", + messages=messages, + retrieval_config={"vector_store_id": "vs_test_123"}, + mock_response="hi there", + ) + + response: Final = await litellm.aquery( + model="gpt-4o-mini", + messages=messages, + retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": provider}, + mock_response="hi there", + ) + + await _drain_logging_worker() + + assert response._hidden_params["response_cost"] > 0 + assert response._hidden_params["response_cost"] == default_provider_response._hidden_params["response_cost"] + + def test_rag_call_types_are_registered(): """ query/aquery/ingest/aingest are @client-decorated entry points, so their diff --git a/tests/unit/repositories/test_repositories.py b/tests/unit/repositories/test_repositories.py index bae6db9ee88..3aed934f40c 100644 --- a/tests/unit/repositories/test_repositories.py +++ b/tests/unit/repositories/test_repositories.py @@ -11,11 +11,13 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from prisma import models as prisma_models from prisma.builder import QueryBuilder +from pydantic import ValidationError from litellm.models.base import DomainModel from litellm.models.budget import LiteLLM_BudgetTable from litellm.models.credentials import CredentialItem from litellm.models.team import LiteLLM_TeamTable +from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper from litellm.repositories.base_repository import BaseRepository from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE @@ -308,6 +310,29 @@ class TestBudgetRepository: assert budget.budget_id == "budget-1" +_ROW_TIMESTAMP: Final = datetime(2026, 1, 1, 12, 0, 0) + + +def _stored_proxy_model_row(*, litellm_params: str, model_info: str) -> prisma_models.LiteLLM_ProxyModelTable: + return prisma_models.LiteLLM_ProxyModelTable( + model_id="other-model", + model_name="gpt-4o", + litellm_params=litellm_params, + model_info=model_info, + blocked=False, + created_at=_ROW_TIMESTAMP, + created_by="admin", + updated_at=_ROW_TIMESTAMP, + updated_by="admin", + ) + + +def _proxy_model_client(rows: list[prisma_models.LiteLLM_ProxyModelTable]) -> SimpleNamespace: + return SimpleNamespace( + db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=AsyncMock(return_value=rows))) + ) + + class TestModelRepository: @pytest.fixture def repo(self): @@ -331,6 +356,56 @@ class TestModelRepository: ).build_query() assert 'where: { model_id: { not: "current-model" } }' in " ".join(query.split()) + @pytest.mark.parametrize("double_encoded", [False, True]) + @pytest.mark.asyncio + async def test_find_all_except_returns_stored_rows_with_decrypted_params( + self, monkeypatch: pytest.MonkeyPatch, double_encoded: bool + ) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt") + stored_params: Final = json.dumps( + {"model": "openai/gpt-4o", "api_key": encrypt_value_helper("sk-secret"), "rpm": 5, "tags": ["prod"]} + ) + stored_info: Final = json.dumps({"id": "other-model", "team_id": "team-1"}) + row: Final = _stored_proxy_model_row( + litellm_params=json.dumps(stored_params) if double_encoded else stored_params, + model_info=json.dumps(stored_info) if double_encoded else stored_info, + ) + + models: Final = await ModelRepository(_proxy_model_client([row])).find_all_except("current-model") + + assert [model.model_dump() for model in models] == [ + { + "model_id": "other-model", + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-secret", "rpm": 5, "tags": ["prod"]}, + "model_info": {"id": "other-model", "team_id": "team-1"}, + "blocked": False, + "created_at": _ROW_TIMESTAMP, + "created_by": "admin", + "updated_at": _ROW_TIMESTAMP, + "updated_by": "admin", + } + ] + + @pytest.mark.asyncio + async def test_find_all_except_keeps_empty_params_and_missing_model_info(self) -> None: + row: Final = _stored_proxy_model_row(litellm_params="{}", model_info="null") + + models: Final = await ModelRepository(_proxy_model_client([row])).find_all_except("current-model") + + assert [(model.litellm_params, model.model_info) for model in models] == [({}, None)] + + @pytest.mark.parametrize("stored_params", ['["sk-secret"]', "1", "true", json.dumps('["sk-secret"]'), '"1"']) + @pytest.mark.asyncio + async def test_find_all_except_rejects_rows_whose_params_are_not_a_json_object(self, stored_params: str) -> None: + row: Final = _stored_proxy_model_row(litellm_params=stored_params, model_info="null") + + with pytest.raises(ValidationError) as rejected: + await ModelRepository(_proxy_model_client([row])).find_all_except("current-model") + + assert [error["type"] for error in rejected.value.errors()] == ["dict_type"] + assert "sk-secret" not in str(rejected.value) + def test_table_is_wrapped_for_config_sync(self, repo): from litellm.proxy.common_utils.config_sync_pubsub import ( _PublishOnWriteActions, diff --git a/tests/unit/router_strategy/adaptive_router/test_adaptive_router.py b/tests/unit/router_strategy/adaptive_router/test_adaptive_router.py index d717c4e8c89..9144bbc1788 100644 --- a/tests/unit/router_strategy/adaptive_router/test_adaptive_router.py +++ b/tests/unit/router_strategy/adaptive_router/test_adaptive_router.py @@ -7,6 +7,7 @@ from litellm.router_strategy.adaptive_router import adaptive_router as ar_module import pytest from litellm.router_strategy.adaptive_router.adaptive_router import AdaptiveRouter +from litellm.router_strategy.adaptive_router.bandit import initial_cell from litellm.router_strategy.adaptive_router.signals import Turn from litellm.types.router import ( AdaptiveRouterConfig, @@ -416,3 +417,11 @@ def test_session_state_expiry_is_refreshed_on_access(): second_exp = r._session_states_expiry[("sess-A", "fast")] assert second_exp > first_exp + + +@pytest.mark.parametrize("model", ["fast", "smart"]) +def test_cell_returns_the_cold_start_prior_for_an_available_model(model): + r = _make_router() + cell = r.cell(RequestType.CODE_GENERATION, model) + assert cell == initial_cell(r.model_to_prefs[model], RequestType.CODE_GENERATION) + assert cell.total_samples == 0 diff --git a/tests/unit/router_strategy/test_complexity_router.py b/tests/unit/router_strategy/test_complexity_router.py index 333524ffffc..dfeb9f8d961 100644 --- a/tests/unit/router_strategy/test_complexity_router.py +++ b/tests/unit/router_strategy/test_complexity_router.py @@ -1671,6 +1671,40 @@ class TestRouterComplexityDeploymentMethods: router.init_complexity_router_deployment(deployment) assert "auto_router/complexity_router/test-router" in router.complexity_routers + @pytest.mark.parametrize("tier_models", ("custom-model", ["custom-model", "other-model"])) + @pytest.mark.parametrize("explicit,expected", ((None, "configured-model"), ("top-model", "top-model"))) + def test_custom_default_resolution_preserves_explicit_precedence(self, tier_models, explicit, expected): + from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig + + config: Final = ComplexityRouterConfig.model_validate({ + "tiers": {"CUSTOM": tier_models, "OTHER": "other-model"}, + "tier_definitions": [ + {"name": "CUSTOM", "description": "Custom work"}, + {"name": "OTHER", "description": "Other work"}, + ], + "fallback_tier": " CUSTOM ", + "classifier_type": "llm", + "classifier_llm_config": {"model": "classifier"}, + "default_model": "configured-model", + }) + assert config.resolve_default_model(explicit) == expected + inferred: Final = config.model_copy(update={"default_model": None}) + assert inferred.resolve_default_model() == "custom-model" + assert config.default_model == "configured-model" + + @pytest.mark.parametrize("config", (None, {})) + def test_absent_deployment_config_still_requires_explicit_default(self, config): + from litellm.types.router import Deployment + + router: Final = Router(model_list=[]) + deployment: Final = Deployment( + model_name="no-default", + litellm_params={"model": "auto_router/complexity_router", "complexity_router_config": config}, + model_info={"id": "no-default"}, + ) + with pytest.raises(ValueError, match="complexity_router_default_model is required"): + router.init_complexity_router_deployment(deployment) + @staticmethod def _forecast_row(model_name: str, model_id: str, classifier_type: str) -> dict[str, object]: settings: Final = ( diff --git a/tests/unit/secret_managers/test_aws_secret_manager.py b/tests/unit/secret_managers/test_aws_secret_manager.py new file mode 100644 index 00000000000..d9922b43ec8 --- /dev/null +++ b/tests/unit/secret_managers/test_aws_secret_manager.py @@ -0,0 +1,36 @@ +import base64 + +import pytest +from botocore.stub import Stubber + +from litellm.secret_managers.aws_secret_manager import AWSKeyManagementService_V2 + +_CIPHERTEXT = b"ciphertext-bytes" +_KEY_ID = "arn:aws:kms:us-west-2:111122223333:key/1234abcd-12ab-34cd-56ef-1234567890ab" + + +@pytest.mark.parametrize( + ("plaintext", "expected"), + [ + pytest.param(b"True", True, id="true literal"), + pytest.param(b"False", False, id="false literal"), + pytest.param(b" True\n", True, id="padded true literal"), + pytest.param(b"1", "1", id="integer literal"), + pytest.param(b"[True]", "[True]", id="list literal"), + pytest.param(b"'True'", "'True'", id="quoted string literal"), + pytest.param(b"sk-not-a-literal", "sk-not-a-literal", id="plain secret"), + pytest.param(b"", "", id="empty secret"), + ], +) +def test_decrypt_value_turns_only_boolean_literals_into_bools(monkeypatch, plaintext, expected): + monkeypatch.setenv("AWS_REGION_NAME", "us-west-2") + monkeypatch.setenv("LITELLM_LICENSE", "license-for-test") + monkeypatch.setenv("ENCRYPTED_SETTING", "aws_kms/" + base64.b64encode(_CIPHERTEXT).decode()) + kms = AWSKeyManagementService_V2() + + with Stubber(kms.kms_client) as stubber: + stubber.add_response("decrypt", {"KeyId": _KEY_ID, "Plaintext": plaintext}, {"CiphertextBlob": _CIPHERTEXT}) + decrypted = kms.decrypt_value(secret_name="ENCRYPTED_SETTING") + + assert decrypted == expected + assert type(decrypted) is type(expected) diff --git a/tests/unit/secret_managers/test_secret_managers_main.py b/tests/unit/secret_managers/test_secret_managers_main.py index acc91b691d4..32040251795 100644 --- a/tests/unit/secret_managers/test_secret_managers_main.py +++ b/tests/unit/secret_managers/test_secret_managers_main.py @@ -431,3 +431,38 @@ def test_secret_manager_would_be_consulted_is_false_without_a_client(monkeypatch monkeypatch.setattr(litellm, "secret_manager_client", None) assert secret_manager_would_be_consulted("os.environ/ANY_NAME") is False + + +class _FixedValueSecretManager(CustomSecretManager): + def __init__(self, value): + self.value = value + + def sync_read_secret(self, secret_name, optional_params=None, timeout=None): + return self.value + + async def async_read_secret(self, secret_name, optional_params=None, timeout=None): + return self.value + + +@pytest.mark.parametrize( + ("stored", "expected"), + [ + ("True", True), + ("False", False), + ("1", "1"), + ("[True]", "[True]"), + ("'True'", "'True'"), + ("sk-not-a-literal", "sk-not-a-literal"), + ("", ""), + ], +) +def test_get_secret_turns_only_boolean_literals_from_the_secret_manager_into_bools(monkeypatch, stored, expected): + monkeypatch.setattr(litellm, "secret_manager_client", _FixedValueSecretManager(stored)) + monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.CUSTOM) + monkeypatch.setattr(litellm, "_key_management_settings", KeyManagementSettings(access_mode="read_only")) + monkeypatch.delenv("STORED_FLAG", raising=False) + + secret = get_secret("os.environ/STORED_FLAG") + + assert secret == expected + assert type(secret) is type(expected) diff --git a/tests/unit/test_assert_ci_coverage.py b/tests/unit/test_assert_ci_coverage.py index 5f3903a5feb..4498dbe1498 100644 --- a/tests/unit/test_assert_ci_coverage.py +++ b/tests/unit/test_assert_ci_coverage.py @@ -84,8 +84,52 @@ def test_integration_groups_require_exclusive_scheduled_circleci_owner(tmp_path: workflow.write_text(yaml.safe_dump({"jobs": {}})) _, missing_invocation = coverage._integration_ownership(tmp_path) assert [(finding.subject, finding.detail) for finding in missing_invocation] == [ - (github_path, "GitHub-owned integration contract has no invoking workflow") + (github_path, "GitHub-owned integration contract has no invoking job") ] + circle_config: Final = yaml.safe_load(circle.read_text()) + circle_config["jobs"]["postgres_suite"] = {"parameters": {"test_path": {"type": "string"}}} + circle_config["workflows"]["integration"]["jobs"].append({"postgres_suite": {"test_path": github_path}}) + circle.write_text(yaml.safe_dump(circle_config)) + _, circle_invocation = coverage._integration_ownership(tmp_path) + assert circle_invocation == () + + +def test_circleci_postgres_suites_count_their_test_path_as_invoked() -> None: + config: Final = { + "workflows": { + "integration": { + "jobs": [ + {"postgres_suite": {"name": "proxy-behavior", "test_path": "tests/proxy_behavior", "seed": True}}, + { + "postgres_suite": { + "name": "roi-database", + "test_path": "tests/integration/database/test_roi_observed.py", + "seed": False, + } + }, + ] + } + } + } + assert coverage._invoked_test_tokens(coverage._scalars(config, "config.yml")) == frozenset( + {"tests/proxy_behavior", "tests/integration/database/test_roi_observed.py"} + ) + + +def test_every_circleci_postgres_suite_is_credited_on_the_repo_as_it_stands() -> None: + circle: Final = yaml.safe_load(coverage.CIRCLECI_CONFIG.read_text()) + paths: Final = tuple( + job["postgres_suite"]["test_path"] + for job in circle["workflows"]["integration"]["jobs"] + if isinstance(job, dict) and "postgres_suite" in job + ) + assert paths == ( + "tests/proxy_behavior", + "tests/proxy_security_tests", + "tests/proxy_migration_tests", + "tests/integration/database/test_roi_observed.py", + ) + assert set(paths) <= coverage._invoked_test_tokens(coverage._all_scalars()) def test_an_ancestor_directory_covers_a_file_but_does_not_name_it(): @@ -183,37 +227,39 @@ def test_every_sharded_root_named_in_the_script_exists_on_disk(): def test_the_repo_as_it_stands_has_every_shard_child_assigned(): findings = coverage._unassigned_shard_children( - coverage._shard_tokens(coverage._all_scalars(), coverage._unit_selection_arms()) + coverage._invoked_test_tokens(coverage._all_scalars()) ) assert [f.subject for f in findings] == [] -def test_shard_tokens_credits_only_wired_unit_flags(tmp_path): - root = tmp_path / "tests" / "tree" - (root / "wired").mkdir(parents=True) - (root / "wired" / "test_a.py").write_text("def test_a(): assert True\n") - (root / "unwired").mkdir(parents=True) - (root / "unwired" / "test_b.py").write_text("def test_b(): assert True\n") - script = tmp_path / ".circleci" / "scripts" / "unit_selection.sh" - script.parent.mkdir(parents=True) - script.write_text( - "legacy_paths() {\n" - " case \"$1\" in\n" - " wired-flag) echo tests/tree/wired ;;\n" - " unwired-flag)\n" - " echo tests/tree/unwired ;;\n" - " esac\n" - "}\n" +def test_shards_are_credited_only_from_explicit_test_paths() -> None: + scalars: Final = ( + coverage.Scalar(key="test-path", value="tests/unit/wired\ntests/unit/also_wired"), + coverage.Scalar(key="shard", value="tests/unit/not_a_test_path"), ) + assert coverage._invoked_test_tokens(scalars) == frozenset({"tests/unit/wired", "tests/unit/also_wired"}) - scalars: Final = (coverage.Scalar(key="unit-flag", value="wired-flag"),) - findings = coverage._unassigned_shard_children( - coverage._shard_tokens(scalars, coverage._unit_selection_arms(tmp_path)), - roots=("tests/tree",), - repo_root=tmp_path, + +def test_an_ignored_path_is_not_credited_as_invoked() -> None: + scalars: Final = ( + coverage.Scalar(key="test-path", value="tests/unit/a\n--ignore=tests/unit/b/test_x.py"), ) + assert coverage._invoked_test_tokens(scalars) == frozenset({"tests/unit/a"}) - assert tuple(f.subject for f in findings) == ("tests/tree/unwired",) + +def test_a_file_its_only_shard_ignores_is_not_covered_by_that_shards_glob() -> None: + selections: Final = coverage._invoked_selections( + ( + coverage.Scalar( + key="test-path", + value="tests/unit/proxy/test_*.py --ignore=tests/unit/proxy/test_update_spend.py", + ), + ) + ) + assert len(selections) == 1 + selection: Final = selections[0] + assert selection.covers("tests/unit/proxy/test_other.py") is True + assert selection.covers("tests/unit/proxy/test_update_spend.py") is False def test_check_shards_passes_on_the_repo_as_it_stands(capsys): diff --git a/tests/unit/test_circleci_path_filter.py b/tests/unit/test_circleci_path_filter.py index 3776aea28e4..281790ee13e 100644 --- a/tests/unit/test_circleci_path_filter.py +++ b/tests/unit/test_circleci_path_filter.py @@ -137,6 +137,21 @@ def test_classify_decisions(category: str, changed: list[str], expected: str) -> assert classify(category, changed) == expected +@pytest.mark.parametrize( + ("changed", "expected"), + ( + ("litellm/caching/redis_cache.py", "run"), + ("tests/unit/caching/test_redis_cluster_cache.py", "run"), + (".circleci/config.yml", "run"), + ("uv.lock", "run"), + ("litellm/router.py", "skip"), + ("docs/my-website/redis.md", "skip"), + ), +) +def test_redis_compat_path_filter(changed: str, expected: str) -> None: + assert classify("redis-compat", [changed]) == expected + + def test_markdown_under_ui_counts_as_client_not_docs() -> None: assert classify("client", ["ui/litellm-dashboard/README.md"]) == "run" assert classify("backend", ["ui/litellm-dashboard/README.md"]) == "skip" diff --git a/tests/unit/test_dashscope_image_generation.py b/tests/unit/test_dashscope_image_generation.py index 6f91fe9a0e0..960807a033b 100644 --- a/tests/unit/test_dashscope_image_generation.py +++ b/tests/unit/test_dashscope_image_generation.py @@ -9,6 +9,7 @@ from unittest.mock import MagicMock, patch import httpx import pytest +from pydantic import ValidationError import litellm from litellm.llms.dashscope.image_generation.transformation import ( @@ -455,3 +456,98 @@ def test_litellm_image_generation_dashscope_end_to_end(model: str): assert "input" in body assert "messages" in body["input"] assert body["parameters"]["size"] == "1024*1024" + + +def _transform_response(payload: object) -> ImageResponse: + return DashScopeImageGenerationConfig().transform_image_generation_response( + model="qwen-image-2.0", + raw_response=httpx.Response(200, json=payload), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +@pytest.mark.parametrize( + "payload", + [ + {}, + {"output": {}}, + {"output": {"choices": []}}, + {"output": {"choices": ""}}, + {"output": {"choices": [{}]}}, + {"output": {"choices": [{"message": {}}]}}, + {"output": {"choices": [{"message": {"content": {}}}]}}, + {"output": {"choices": [{"message": {"content": [{}]}}]}}, + {"output": {"choices": [{"message": {"content": [{"image": ""}]}}]}}, + {"output": {"choices": [{"message": {"content": [{"text": "hi"}]}}]}}, + {"code": "Partial", "output": {"choices": []}}, + ], +) +def test_transform_response_without_image_content_has_no_images( + payload: dict[str, object], +): + assert _transform_response(payload).data == [] + + +def test_transform_response_skips_content_items_without_an_image(): + response = _transform_response( + { + "output": { + "choices": [ + {"message": {"content": [{"text": "caption"}, {"image": "https://a.example/1.png"}]}}, + {"message": {"content": [{"image": None}, {"image": "https://a.example/2.png"}]}}, + ] + } + } + ) + + assert [image.url for image in response.data] == [ + "https://a.example/1.png", + "https://a.example/2.png", + ] + + +@pytest.mark.parametrize( + ("payload", "message"), + [ + ({"code": "InvalidParameter", "message": "Size not supported"}, "Size not supported"), + ({"code": "Throttled"}, "{'code': 'Throttled'}"), + ({"code": 429, "message": {"detail": "slow down"}}, "{'detail': 'slow down'}"), + ], +) +def test_transform_response_reports_api_error_bodies( + payload: dict[str, object], message: str +): + with pytest.raises(BaseLLMException) as exc_info: + _transform_response(payload) + + assert exc_info.value.message == message + assert exc_info.value.status_code == 200 + + +@pytest.mark.parametrize( + "payload", + [ + ["code"], + "code", + 7, + {"output": None}, + {"output": ["not", "an", "object"]}, + {"output": {"choices": 7}}, + {"output": {"choices": ["not an object"]}}, + {"output": {"choices": [{"message": "not an object"}]}}, + {"output": {"choices": [{"message": {"content": None}}]}}, + {"output": {"choices": [{"message": {"content": ["not an object"]}}]}}, + ], +) +def test_transform_response_rejects_malformed_payloads_without_echoing_them( + payload: object, +): + with pytest.raises(ValidationError) as exc_info: + _transform_response(payload) + + assert "input_value" not in str(exc_info.value) diff --git a/tests/unit/test_unit_shard_missing_paths.py b/tests/unit/test_unit_shard_missing_paths.py index 0360a227142..4fa9c5bd3c1 100644 --- a/tests/unit/test_unit_shard_missing_paths.py +++ b/tests/unit/test_unit_shard_missing_paths.py @@ -38,7 +38,6 @@ def _run_shard(tmp_path: Path, test_path: str, workers: str) -> subprocess.Compl "PATH": f"{shim_dir}{os.pathsep}{os.environ['PATH']}", "GITHUB_OUTPUT": str(tmp_path / "github_output"), "TEST_PATH": test_path, - "UNIT_FLAG": "", "WORKERS": workers, }, capture_output=True, diff --git a/tests/unit/vector_stores/test_vector_store_registry.py b/tests/unit/vector_stores/test_vector_store_registry.py index 762176d6a81..acfaccf8e2d 100644 --- a/tests/unit/vector_stores/test_vector_store_registry.py +++ b/tests/unit/vector_stores/test_vector_store_registry.py @@ -1,10 +1,14 @@ import json +from collections.abc import Mapping, Sequence +from types import SimpleNamespace +from typing import Final from unittest.mock import patch import httpx import pytest import respx from fastapi.testclient import TestClient +from pydantic import BaseModel, ValidationError from datetime import datetime, timezone @@ -13,7 +17,7 @@ from unittest.mock import AsyncMock, MagicMock import litellm from litellm.types.vector_stores import LiteLLM_ManagedVectorStore from litellm.vector_stores.main import search -from litellm.vector_stores.vector_store_registry import VectorStoreRegistry +from litellm.vector_stores.vector_store_registry import VectorStoreIndexRegistry, VectorStoreRegistry @pytest.fixture(autouse=True) @@ -249,3 +253,93 @@ async def test_config_owned_store_survives_db_liveness_check_while_missing_db_st prisma_client.db.litellm_managedvectorstorestable.find_unique.assert_awaited_once_with( where={"vector_store_id": "vs_from_db"} ) + + +_INDEX_PARAMS: Final = {"vector_store_index": "real-index-name", "vector_store_name": "azure-ai-search-store"} +_INDEX_CREATED_AT: Final = datetime(2026, 1, 2, 3, 4, 5, tzinfo=timezone.utc) +_INDEX_ROW: Final = { + "id": "idx-1", + "index_name": "team-docs", + "litellm_params": _INDEX_PARAMS, + "index_info": {"dimensions": 1536}, + "created_at": _INDEX_CREATED_AT, + "created_by": "user-1", + "updated_at": _INDEX_CREATED_AT, + "updated_by": "user-2", +} + + +class GeneratedIndexRow(BaseModel): + id: str + index_name: str + litellm_params: dict[str, str] + index_info: dict[str, int] | None = None + created_at: datetime + created_by: str | None = None + updated_at: datetime + updated_by: str | None = None + + +def _database_listing_index_rows(rows: Sequence[object]) -> SimpleNamespace: + async def find_many(order: Mapping[str, str]) -> Sequence[object]: + return rows if order == {"created_at": "desc"} else [] + + return SimpleNamespace( + db=SimpleNamespace(litellm_managedvectorstoreindextable=SimpleNamespace(find_many=find_many)) + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "row", + [ + _INDEX_ROW, + GeneratedIndexRow(**_INDEX_ROW), + list(_INDEX_ROW.items()), + {**_INDEX_ROW, "column_added_by_a_later_migration": "ignored"}, + ], +) +async def test_vector_store_index_rows_from_the_db_are_returned_as_indexes(row: object) -> None: + indexes: Final = await VectorStoreIndexRegistry._get_vector_store_indexes_from_db( + _database_listing_index_rows([row]) + ) + + assert [index.model_dump() for index in indexes] == [_INDEX_ROW] + + +@pytest.mark.asyncio +async def test_vector_store_index_row_without_optional_columns_gets_empty_defaults() -> None: + indexes: Final = await VectorStoreIndexRegistry._get_vector_store_indexes_from_db( + _database_listing_index_rows([{"id": "idx-1", "index_name": "team-docs", "litellm_params": _INDEX_PARAMS}]) + ) + + assert [index.model_dump() for index in indexes] == [ + { + "id": "idx-1", + "index_name": "team-docs", + "litellm_params": _INDEX_PARAMS, + "index_info": None, + "created_at": None, + "created_by": None, + "updated_at": None, + "updated_by": None, + } + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "row", + [ + {"index_name": "team-docs", "litellm_params": _INDEX_PARAMS}, + {**_INDEX_ROW, "id": 7}, + {**_INDEX_ROW, "litellm_params": None}, + {**_INDEX_ROW, "litellm_params": {"vector_store_index": "real-index-name"}}, + {**_INDEX_ROW, "index_info": ["not", "a", "mapping"]}, + {**_INDEX_ROW, 7: "column names must be strings"}, + {**_INDEX_ROW, b"index_name": "column names are not decoded"}, + ], +) +async def test_malformed_vector_store_index_row_from_the_db_raises_a_validation_error(row: object) -> None: + with pytest.raises(ValidationError): + await VectorStoreIndexRegistry._get_vector_store_indexes_from_db(_database_listing_index_rows([row])) diff --git a/ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg b/ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg new file mode 100644 index 00000000000..0dd78b078c9 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/llm_shield_proxy.svg @@ -0,0 +1,6 @@ + + + + + + diff --git a/ui/litellm-dashboard/src/app/(dashboard)/components/SidebarProvider.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/components/SidebarProvider.test.tsx new file mode 100644 index 00000000000..a3d75df9606 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/components/SidebarProvider.test.tsx @@ -0,0 +1,45 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render, screen, within } from "@testing-library/react"; +import type { ComponentProps } from "react"; +import { expect, it, vi } from "vitest"; +import Sidebar from "@/components/leftnav"; +import SidebarProvider from "./SidebarProvider"; + +const { getSettings } = vi.hoisted(() => ({ + getSettings: vi.fn().mockResolvedValue({ values: { enable_projects_ui: true } }), +})); + +vi.mock("@/components/networking", () => ({ getUiSettings: getSettings, getUISettings: getSettings })); +vi.mock("@/components/leftnav", () => ({ + default: ({ enableProjectsUI }: ComponentProps) => ( + + ), +})); + +it("reuses loaded navigation settings immediately when the mobile sidebar mounts again", async () => { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + const shell = (mobileOpen: boolean) => ( + +
+ +
+ {mobileOpen && ( +
+ +
+ )} +
+ ); + const { rerender } = render(shell(false)); + expect(await screen.findByRole("link", { name: "Projects" })).toBeVisible(); + + for (const open of [true, false, true]) { + rerender(shell(open)); + if (open) { + expect( + within(screen.getByRole("region", { name: "Mobile navigation" })).getByRole("link", { name: "Projects" }), + ).toBeVisible(); + } + } + expect(getSettings).toHaveBeenCalledTimes(1); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/components/SidebarProvider.tsx b/ui/litellm-dashboard/src/app/(dashboard)/components/SidebarProvider.tsx index 6aaf08dae79..c7794233ded 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/components/SidebarProvider.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/components/SidebarProvider.tsx @@ -1,9 +1,7 @@ "use client"; import Sidebar from "@/components/leftnav"; -import { getUISettings } from "@/components/networking"; -import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -import { useEffect, useState } from "react"; +import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; interface SidebarProviderProps { sidebarCollapsed: boolean; @@ -11,66 +9,19 @@ interface SidebarProviderProps { } const SidebarProvider = ({ sidebarCollapsed, onToggleCollapsed }: SidebarProviderProps) => { - const { accessToken } = useAuthorized(); - const [enabledPagesInternalUsers, setEnabledPagesInternalUsers] = useState(null); - const [enableProjectsUI, setEnableProjectsUI] = useState(false); - const [disableAgentsForInternalUsers, setDisableAgentsForInternalUsers] = useState(false); - const [allowAgentsForTeamAdmins, setAllowAgentsForTeamAdmins] = useState(false); - const [disableVectorStoresForInternalUsers, setDisableVectorStoresForInternalUsers] = useState(false); - const [allowVectorStoresForTeamAdmins, setAllowVectorStoresForTeamAdmins] = useState(false); - - useEffect(() => { - const fetchUISettings = async () => { - if (!accessToken) { - return; - } - - try { - const settings = await getUISettings(accessToken); - - // API returns 'values' not 'settings' - if (settings?.values?.enabled_ui_pages_internal_users !== undefined) { - setEnabledPagesInternalUsers(settings.values.enabled_ui_pages_internal_users); - } else { - } - - if (settings?.values?.enable_projects_ui !== undefined) { - setEnableProjectsUI(Boolean(settings.values.enable_projects_ui)); - } - - if (settings?.values?.disable_agents_for_internal_users !== undefined) { - setDisableAgentsForInternalUsers(Boolean(settings.values.disable_agents_for_internal_users)); - } - - if (settings?.values?.allow_agents_for_team_admins !== undefined) { - setAllowAgentsForTeamAdmins(Boolean(settings.values.allow_agents_for_team_admins)); - } - - if (settings?.values?.disable_vector_stores_for_internal_users !== undefined) { - setDisableVectorStoresForInternalUsers(Boolean(settings.values.disable_vector_stores_for_internal_users)); - } - - if (settings?.values?.allow_vector_stores_for_team_admins !== undefined) { - setAllowVectorStoresForTeamAdmins(Boolean(settings.values.allow_vector_stores_for_team_admins)); - } - } catch (error) { - console.error("[SidebarProvider] Failed to fetch UI settings:", error); - } - }; - - fetchUISettings(); - }, [accessToken]); + const { data: settings } = useUISettings(); + const values = settings?.values; return ( ); }; 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 29df7c8bf3d..a02cae097a7 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 @@ -73,7 +73,9 @@ interface GuardrailPreset { provider: string; categoryName?: string; guardrailNameSuggestion: string; - mode: string; + // A guardrail that both rewrites the request and repairs the response needs two + // modes seeded, not one; the form already normalises either shape. + mode: string | string[]; defaultOn: boolean; } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts index 10ca58294b9..684049c54a9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts @@ -2,7 +2,9 @@ export interface GuardrailPreset { provider: string; categoryName?: string; guardrailNameSuggestion: string; - mode: string; + // A guardrail that both rewrites the request and repairs the response needs two + // modes seeded, not one; the form already normalises either shape. + mode: string | string[]; defaultOn: boolean; } @@ -325,6 +327,14 @@ export const GUARDRAIL_PRESETS: Record = { // MCP-only: default_on is the only activation path on the MCP hook defaultOn: true, }, + llm_shield_proxy: { + provider: "LLM Shield Proxy", + guardrailNameSuggestion: "LLM Shield Proxy", + // Both halves are required. With only pre_call the request is redacted and the + // placeholders are handed straight back to the caller. + mode: ["pre_call", "post_call"], + defaultOn: false, + }, conduct: { provider: "Conduct", guardrailNameSuggestion: "Conduct Guard", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts index a27dd344c95..e756c8c3e58 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts @@ -29,6 +29,7 @@ const EXPECTED_PARTNER_LOGO_FILES: Record = { straiker: "straiker.svg", alice: "alice.svg", agent_365: "microsoft_azure.svg", + llm_shield_proxy: "llm_shield_proxy.svg", conduct: "conduct.png", }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts index d88a333d6f1..d1ead5f589a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts @@ -484,6 +484,16 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ tags: ["Agentic", "MCP", "Tool Misuse", "Observability"], providerKey: "Agent365", }, + { + id: "llm_shield_proxy", + name: "LLM Shield Proxy", + description: + "Self-hosted PII redaction that puts the original values back into the model's response, so the provider never receives personal data while the end user still sees it.", + category: "partner", + logo: guardrailLogoMap["LLM Shield Proxy"], + tags: ["PII", "Data Privacy", "Compliance", "Streaming"], + providerKey: "LLM Shield Proxy", + }, { id: "conduct", name: "Conduct Guard", 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 8df7dfb1403..09e0f21c670 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 @@ -1,6 +1,7 @@ import aimSecurityLogo from "../../../../../public/assets/logos/aim_security.jpeg"; import aktoLogo from "../../../../../public/assets/logos/akto.svg"; import aliceLogo from "../../../../../public/assets/logos/alice.svg"; +import llmShieldProxyLogo from "../../../../../public/assets/logos/llm_shield_proxy.svg"; import conductLogo from "../../../../../public/assets/logos/conduct.png"; import aporiaLogo from "../../../../../public/assets/logos/aporia.png"; import bedrockLogo from "../../../../../public/assets/logos/bedrock.svg"; @@ -86,6 +87,7 @@ export const guardrail_provider_map: Record = { QostodianNexus: "qostodian_nexus", Repelloai: "repelloai", Alice: "alice", + "LLM Shield Proxy": "llm_shield_proxy", Conduct: "conduct", }; @@ -211,6 +213,7 @@ export const guardrailLogoMap = { Straiker: straikerLogo.src, Alice: aliceLogo.src, "Microsoft Agent 365": microsoftAzureLogo.src, + "LLM Shield Proxy": llmShieldProxyLogo.src, "Conduct Guard": conductLogo.src, } satisfies Record; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/layout.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/layout.test.tsx index 7cdf7aa0489..876a4e24d84 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/layout.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/layout.test.tsx @@ -1,5 +1,5 @@ import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; -import { render, screen, waitFor } from "@testing-library/react"; +import { fireEvent, render, screen, waitFor, within } from "@testing-library/react"; import { usePathname } from "next/navigation"; import { AuthProvider } from "@/contexts/AuthContext"; import Layout from "./layout"; @@ -19,12 +19,17 @@ vi.mock("@/components/liteadmin/LiteAdmin", () => ({ })); vi.mock("@/components/DashboardHeader", () => ({ - DashboardHeader: () =>
, + DashboardHeader: ({ navigationTrigger }: { navigationTrigger?: React.ReactNode }) => ( +
{navigationTrigger}
+ ), })); vi.mock("@/app/(dashboard)/components/SidebarProvider", () => ({ - default: ({ sidebarCollapsed }: { sidebarCollapsed: boolean }) => ( -
+ default: ({ sidebarCollapsed, onToggleCollapsed }: { sidebarCollapsed: boolean; onToggleCollapsed: () => void }) => ( +
+ + Settings +
), })); @@ -89,6 +94,47 @@ describe("(dashboard) Layout", () => { vi.mocked(usePathname).mockReturnValue("/ui/guardrails"); }); + it("starts mobile navigation closed, opens a modal drawer and closes it after choosing a page", async () => { + render( + + +

Gateway content

+
+
, + ); + pendingUiConfig.resolve(); + + const trigger = await screen.findByRole("button", { name: "Open navigation" }); + expect(screen.queryByRole("dialog", { name: "Navigation" })).not.toBeInTheDocument(); + + fireEvent.click(trigger); + const navigation = await screen.findByRole("dialog", { name: "Navigation" }); + expect(within(navigation).getByRole("button", { name: "Close navigation" })).toBeInTheDocument(); + fireEvent.click(within(navigation).getByRole("link", { name: "Settings" })); + await waitFor(() => expect(screen.queryByRole("dialog", { name: "Navigation" })).not.toBeInTheDocument()); + }); + + it("closes the mobile drawer when navigation changes outside the drawer", async () => { + const dashboard = () => ( + + +

Gateway content

+
+
+ ); + const { rerender } = render(dashboard()); + pendingUiConfig.resolve(); + fireEvent.click(await screen.findByRole("button", { name: "Open navigation" })); + expect(await screen.findByRole("dialog", { name: "Navigation" })).toBeInTheDocument(); + + vi.mocked(usePathname).mockReturnValue("/ui/api-keys"); + rerender(dashboard()); + await waitFor(() => expect(screen.queryByRole("dialog", { name: "Navigation" })).not.toBeInTheDocument()); + vi.mocked(usePathname).mockReturnValue("/ui/guardrails"); + rerender(dashboard()); + expect(screen.queryByRole("dialog", { name: "Navigation" })).not.toBeInTheDocument(); + }); + it("collapses the sidebar on Logs for a full-screen view and expands it again after leaving", async () => { const dashboard = () => ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx index d705089cee8..4762dd2fed1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx @@ -19,6 +19,10 @@ import { routeSegmentForPathname, uiHref } from "@/utils/uiHref"; import { PluginModeProvider, usePluginMode } from "@/contexts/PluginModeContext"; import { createApiClient } from "@/lib/http/client"; import { getProxyBaseUrl } from "@/components/networking"; +import { Sheet, SheetContent, SheetTitle, SheetTrigger } from "@/components/ui/sheet"; +import { Button } from "@/components/ui/button"; +import { Menu } from "lucide-react"; +import { useMediaQuery } from "usehooks-ts"; const pluginApiClient = createApiClient({ getBaseUrl: () => getProxyBaseUrl() ?? "" }); @@ -104,7 +108,16 @@ const FULL_BLEED_SEGMENTS = new Set(["logs"]); function DashboardShell({ children }: { children: React.ReactNode }) { const { accessToken } = useAuth(); const { mode } = usePluginMode(); - const routeSegment = routeSegmentForPathname(usePathname()); + const pathname = usePathname(); + const routeSegment = routeSegmentForPathname(pathname); + const searchParams = useSearchParams(); + const navigationKey = `${pathname}?${searchParams.toString()}`; + const isDesktop = useMediaQuery("(min-width: 768px)", { initializeWithValue: false }); + const [mobileNavigationKey, setMobileNavigationKey] = useState(null); + if (mobileNavigationKey !== null && (isDesktop || mobileNavigationKey !== navigationKey)) { + setMobileNavigationKey(null); + } + const mobileNavigationOpen = !isDesktop && mobileNavigationKey === navigationKey; const isFullBleed = FULL_BLEED_SEGMENTS.has(routeSegment); // A manual toggle holds only for the route it was made on; full-bleed routes default to collapsed. const [sidebarOverride, setSidebarOverride] = useState<{ segment: string; collapsed: boolean } | null>(null); @@ -138,21 +151,46 @@ function DashboardShell({ children }: { children: React.ReactNode }) { // sidebar owns its own scroll and the content column scrolls independently, // so the page can't be dragged past the end of the nav. return ( -
- - -
- - - - - - - -
{children}
+ setMobileNavigationKey(open ? navigationKey : null)}> +
+
+
- -
+ { + if (event.target instanceof Element && event.target.closest("a[href]")) setMobileNavigationKey(null); + }} + > + Navigation + setMobileNavigationKey(null)} /> + + +
+ + } + > + + + } + /> + + + + + + +
{children}
+
+
+
+ ); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/legacyPageRoutes.ts b/ui/litellm-dashboard/src/app/(dashboard)/legacyPageRoutes.ts index eecb897634b..09cc4ef0a75 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/legacyPageRoutes.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/legacyPageRoutes.ts @@ -37,7 +37,6 @@ const LEGACY_PAGE_ROUTES: ReadonlyMap = new Map( usage: "old-usage", "cost-optimization": "cost-optimization", "model-insights": "model-insights", - "roi-calculator": "roi-calculator", agents: "agents", "router-settings": "router-settings", users: "users", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx index 3cf2bcdf238..f6ad2e38658 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx @@ -28,7 +28,10 @@ const oauthHook = vi.hoisted(() => ({ tokenResponse: null as Record | null, reset: vi.fn(), onTokenReceived: null as - | ((token: Record | null, registeredClient?: { clientId?: string; clientSecret?: string }) => void) + | (( + token: Record | null, + registeredClient?: { client_id: string; client_secret?: string }, + ) => void) | null, getCredentials: null as (() => Record | undefined) | null, getTemporaryPayload: null as (() => Record | null) | null, @@ -37,7 +40,7 @@ vi.mock("@/hooks/useMcpOAuthFlow", () => ({ useMcpOAuthFlow: (opts: { onTokenReceived: ( token: Record | null, - registeredClient?: { clientId?: string; clientSecret?: string }, + registeredClient?: { client_id: string; client_secret?: string }, ) => void; getCredentials?: () => Record | undefined; getTemporaryPayload?: () => Record | null; @@ -628,7 +631,7 @@ describe("CreateMCPServer", () => { await act(async () => { oauthHook.onTokenReceived!( { access_token: "oauth2-minted-tok", refresh_token: "oauth2-minted-refresh", token_type: "Bearer" }, - { clientId: "dcr-minted-client", clientSecret: "dcr-minted-secret" }, + { client_id: "dcr-minted-client", client_secret: "dcr-minted-secret" }, ); }); @@ -675,7 +678,7 @@ describe("CreateMCPServer", () => { await act(async () => { oauthHook.onTokenReceived!( { access_token: "oauth2-tok", token_type: "Bearer" }, - { clientId: "dcr-client", clientSecret: "dcr-secret" }, + { client_id: "dcr-client", client_secret: "dcr-secret" }, ); }); @@ -702,7 +705,7 @@ describe("CreateMCPServer", () => { await act(async () => { oauthHook.onTokenReceived!( { access_token: "oauth2-tok", token_type: "Bearer" }, - { clientId: "leak-client", clientSecret: "leak-secret" }, + { client_id: "leak-client", client_secret: "leak-secret" }, ); }); // Ref is held while the modal is open. @@ -731,7 +734,7 @@ describe("CreateMCPServer", () => { await act(async () => { oauthHook.onTokenReceived!( { access_token: "oauth2-tok", token_type: "Bearer" }, - { clientId: "dcr-client", clientSecret: "dcr-secret" }, + { client_id: "dcr-client", client_secret: "dcr-secret" }, ); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx index 2d6e83aff9b..d53d11ddcfb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx @@ -56,7 +56,7 @@ import EnvVarsSection from "./EnvVarsSection"; import { isAdminRole } from "@/utils/roles"; import { validateMCPServerUrl, validateMCPServerName } from "./utils"; import { toast } from "@/lib/toast"; -import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; +import { useMcpOAuthFlow, type McpDcrCredentials } from "@/hooks/useMcpOAuthFlow"; import { useTestMCPConnection } from "@/hooks/useTestMCPConnection"; import { MountedFormField, @@ -144,7 +144,7 @@ const CreateMCPServer: React.FC = ({ // it can never be collected as a client-forwarded server's declared app; injected into the payload // only on an oauth2 submit (where persisting the registered client is correct), and cleared on any // invalidation or modal close. An abandoned authorize leaves it null, which is the desired asymmetry. - const dcrClientRef = React.useRef<{ client_id: string; client_secret?: string } | null>(null); + const dcrClientRef = React.useRef(null); // Set when the upstream identity (url/endpoints) changed while a declared app is present, so the // section can warn that the saved app may not match the new upstream (the app is kept, not wiped). const [appMayNotMatchUpstream, setAppMayNotMatchUpstream] = useState(false); @@ -264,12 +264,7 @@ const CreateMCPServer: React.FC = ({ // The DCR-minted client is held in a ref, NOT written into form.credentials, so it can never be // collected as a client-forwarded server's declared app; it is injected into the payload only on // an oauth2 submit. An admin-typed client already lives in form.credentials and is left untouched. - dcrClientRef.current = registeredClient?.clientId - ? { - client_id: registeredClient.clientId, - ...(registeredClient.clientSecret && { client_secret: registeredClient.clientSecret }), - } - : null; + dcrClientRef.current = registeredClient ?? null; const current = (allFieldsValue(form).credentials as Record | undefined) ?? {}; const nextCredentials = { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createServerPayload.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createServerPayload.ts index f45857fcc8b..30c4c502c42 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createServerPayload.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createServerPayload.ts @@ -29,7 +29,7 @@ export const AUTH_TYPES_REQUIRING_CREDENTIALS = [ export interface DcrClient { readonly client_id: string; - readonly client_secret?: string; + readonly client_secret?: string | null; } export interface CreateServerUiState { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.test.ts index 66786fac1c3..8a9bac582d6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.test.ts @@ -370,3 +370,31 @@ void isClientForwardedTokenMode; void normalizeEnvVars; void preservedAdminCredentials; export type { MCPServer }; + +it.each([ + [AUTH_TYPE.OAUTH2, OAUTH_FLOW.INTERACTIVE, true], + [AUTH_TYPE.OAUTH2, OAUTH_FLOW.M2M, false], + [AUTH_TYPE.TRUE_PASSTHROUGH, OAUTH_FLOW.INTERACTIVE, false], +])("applies a pending DCR client only to an interactive gateway OAuth save (%s, %s)", (authType, flow, useDcr) => { + const values: EditServerFormValues = { + auth_type: authType, + oauth_flow_type: flow, + transport: "http", + credentials: { client_id: "configured-client", client_secret: "configured-secret" }, + }; + const ui: EditServerUiState = { + ...baseUi, + dcrClient: { + client_id: "new-client", + client_secret: null, + dcr_issuer: "https://new.example", + dcr_server_url: "https://new.example/mcp", + }, + }; + const result = buildEditServerPayload(values, ui); + expect(result.kind).toBe("ok"); + if (result.kind !== "ok") throw new Error("Expected valid MCP server payload"); + expect(result.payload.credentials?.client_id).toBe(useDcr ? "new-client" : "configured-client"); + expect(result.payload.credentials?.client_secret).toBe(useDcr ? null : "configured-secret"); + expect(result.payload.credentials?.dcr_issuer).toBe(useDcr ? "https://new.example" : undefined); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.ts index 3930f1edb9b..f166fa20d1f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.ts @@ -1,3 +1,4 @@ +import type { McpDcrCredentials } from "@/hooks/useMcpOAuthFlow"; import { ADMIN_CONFIG_CREDENTIAL_KEYS, AUTH_TYPE, @@ -73,6 +74,7 @@ export interface EditServerPayload { } export interface EditServerUiState { + readonly dcrClient?: McpDcrCredentials | null; readonly mcpServer: MCPServer; readonly logoUrl: string | undefined; readonly costConfig: MCPServerCostInfo; @@ -213,6 +215,8 @@ const buildCredentials = (credentialValues: unknown): Readonly> | undefined; readonly includeCredentials: boolean; @@ -224,11 +228,16 @@ interface CredentialsEntryInput { // updates), so removal must be an explicit-null write: encrypt skips nulls and the merge overrides // the stored keys, returning the server to dynamic client registration. const resolveCredentialsEntry = ({ + dcrClient, + oauthFlow, authType, credentials, includeCredentials, removeStoredApp, }: CredentialsEntryInput): { readonly credentials?: Readonly> } => { + if (dcrClient && authType === AUTH_TYPE.OAUTH2 && oauthFlow !== OAUTH_FLOW.M2M) { + return { credentials: { ...credentials, ...dcrClient } }; + } if (removeStoredApp && isClientForwardedTokenMode(authType)) { return { credentials: { client_id: null, client_secret: null } }; } @@ -330,6 +339,8 @@ export const buildEditServerPayload = (values: EditServerFormValues, ui: EditSer const includeCredentials = restValues.auth_type && AUTH_TYPES_REQUIRING_CREDENTIALS.includes(restValues.auth_type); const credentialsEntryInput: CredentialsEntryInput = { + dcrClient: ui.dcrClient, + oauthFlow: restValues.oauth_flow_type, authType: restValues.auth_type, credentials: submitCredentials, includeCredentials: Boolean(includeCredentials), diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx index f79dd178571..91db38db389 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx @@ -16,21 +16,40 @@ vi.mock("@/components/networking", () => ({ })); const mockOauth: { + status: string; tokenResponse: any; getTemporaryPayload: (() => Record | null) | null; - onTokenReceived: ((token: Record | null) => void) | null; + onTokenReceived: + | (( + token: Record | null, + registeredClient?: { + client_id: string; + client_secret?: string; + dcr_issuer?: string; + dcr_server_url?: string; + }, + ) => void) + | null; reset: ReturnType; -} = { tokenResponse: null, getTemporaryPayload: null, onTokenReceived: null, reset: vi.fn() }; +} = { status: "idle", tokenResponse: null, getTemporaryPayload: null, onTokenReceived: null, reset: vi.fn() }; vi.mock("@/hooks/useMcpOAuthFlow", () => ({ useMcpOAuthFlow: (opts: { getTemporaryPayload?: () => Record | null; - onTokenReceived?: (token: Record | null) => void; + onTokenReceived?: ( + token: Record | null, + registeredClient?: { + client_id: string; + client_secret?: string; + dcr_issuer?: string; + dcr_server_url?: string; + }, + ) => void; }) => { mockOauth.getTemporaryPayload = opts?.getTemporaryPayload ?? null; mockOauth.onTokenReceived = opts?.onTokenReceived ?? null; return { startOAuthFlow: vi.fn(), - status: "idle", + status: mockOauth.status, error: null, tokenResponse: mockOauth.tokenResponse, reset: mockOauth.reset, @@ -547,6 +566,7 @@ describe("MCPServerEdit (auth type switch)", () => { describe("MCPServerEdit OAuth token invalidation", () => { beforeEach(() => { vi.clearAllMocks(); + mockOauth.status = "idle"; }); const renderOAuthEdit = () => @@ -560,6 +580,77 @@ describe("MCPServerEdit OAuth token invalidation", () => { />, ); + it.each(["authorizing", "exchanging"])("blocks Save while OAuth is %s", async (status) => { + mockOauth.status = status; + vi.mocked(networking.updateMCPServer).mockResolvedValue({ ...interactiveOAuthServer }); + const view = renderOAuthEdit(); + const save = screen.getAllByRole("button", { name: "Save Changes" })[0]; + expect(save).toBeDisabled(); + await act(async () => { + fireEvent.submit(screen.getByRole("form", { name: "Edit MCP server" })); + }); + expect(networking.updateMCPServer).not.toHaveBeenCalled(); + act(() => { + mockOauth.onTokenReceived?.( + { access_token: "new-token" }, + { + client_id: "new-client", + dcr_issuer: "https://new.example", + dcr_server_url: "https://new.example/mcp", + }, + ); + }); + mockOauth.status = "success"; + view.rerender( + , + ); + expect(screen.getAllByRole("button", { name: "Save Changes" })[0]).toBeEnabled(); + await act(async () => { + fireEvent.click(screen.getAllByRole("button", { name: "Save Changes" })[0]); + }); + await waitFor(() => expect(networking.updateMCPServer).toHaveBeenCalledTimes(1)); + expect(vi.mocked(networking.updateMCPServer).mock.calls[0][1].credentials).toMatchObject({ + client_id: "new-client", + dcr_issuer: "https://new.example", + }); + }); + + it.each(["Save Changes", "Cancel"])("keeps a newly registered client isolated until %s", async (action) => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ ...interactiveOAuthServer }); + renderOAuthEdit(); + act(() => { + mockOauth.onTokenReceived?.( + { access_token: "new-token" }, + { + client_id: "new-client", + client_secret: "new-secret", + dcr_issuer: "https://new.example", + dcr_server_url: "https://new.example/mcp", + }, + ); + }); + expect(networking.updateMCPServer).not.toHaveBeenCalled(); + await act(async () => { + fireEvent.click(screen.getAllByRole("button", { name: action })[0]); + }); + if (action === "Cancel") { + expect(networking.updateMCPServer).not.toHaveBeenCalled(); + return; + } + await waitFor(() => expect(networking.updateMCPServer).toHaveBeenCalledTimes(1)); + expect(vi.mocked(networking.updateMCPServer).mock.calls[0][1].credentials).toMatchObject({ + client_id: "new-client", + client_secret: "new-secret", + dcr_issuer: "https://new.example", + }); + }); + it("invalidates a session-authorized token when the transport switches to stdio", async () => { // Switching to stdio clears url/auth_type via programmatic form.setFieldsValue, which antd does // not report through onValuesChange; the explicit recheck in handleTransportChange must catch it. @@ -1687,6 +1778,51 @@ describe("MCPServerEdit (OAuth token persistence on save)", () => { expect(screen.getByText(/registered for the previous upstream/)).toBeInTheDocument(); }); + it("discards a canceled OAuth snapshot before saved server data loads", async () => { + setSecureItem( + EDIT_OAUTH_UI_STATE_KEY, + JSON.stringify({ + serverId: interactiveOAuthServer.server_id, + formValues: { ...interactiveOAuthServer, url: "https://new.example/mcp" }, + }), + ); + const props = { accessToken: "access-token", onCancel: vi.fn(), onSuccess: vi.fn(), availableAccessGroups: [] }; + const view = render(); + await act(async () => { + fireEvent.click(screen.getAllByRole("button", { name: "Cancel" })[0]); + }); + expect(props.onCancel).toHaveBeenCalledOnce(); + expect(window.sessionStorage.getItem(EDIT_OAUTH_UI_STATE_KEY)).toBeNull(); + expect(mockOauth.reset).toHaveBeenCalled(); + view.unmount(); + render(); + await waitFor(() => expect(screen.getByLabelText("MCP Server URL")).toHaveValue(interactiveOAuthServer.url)); + expect(networking.updateMCPServer).not.toHaveBeenCalled(); + }); + + it("restores the edited upstream after OAuth when saved server data loads later", async () => { + setSecureItem( + EDIT_OAUTH_UI_STATE_KEY, + JSON.stringify({ + serverId: interactiveOAuthServer.server_id, + formValues: { ...interactiveOAuthServer, url: "https://new.example/mcp", issuer: "https://new.example" }, + }), + ); + const props = { + accessToken: "access-token", + userID: "user-1", + onCancel: vi.fn(), + onSuccess: vi.fn(), + availableAccessGroups: [], + }; + const { rerender } = render( + , + ); + rerender(); + await waitFor(() => expect(screen.getByLabelText("MCP Server URL")).toHaveValue("https://new.example/mcp")); + expect(screen.getByLabelText("Issuer (optional)")).toHaveValue("https://new.example"); + }); + it("preserves a stored client_id on OAuth-resume restore even when the saved snapshot is token-only", async () => { // Post-redirect restore: the sessionStorage snapshot carries only a minted token (no client keys), // while the loaded server has a stored client_id. The restore must merge the server's declared app diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx index 3894cb6cd0b..1b1e4c27bde 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx @@ -55,7 +55,7 @@ import { EditServerFormValues, buildEditServerPayload, editPayloadErrorMessage } import { DUPLICATE_IDENTIFIER_MESSAGE, findDuplicateMcpServer, mcpSubmitErrorReason } from "./duplicateServerCheck"; import { toast } from "@/lib/toast"; import { getEditToolPreview } from "./editToolPreview"; -import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; +import { useMcpOAuthFlow, type McpDcrCredentials } from "@/hooks/useMcpOAuthFlow"; import { MountedFormField, MountedFormProvider, @@ -250,6 +250,7 @@ const MCPServerEdit: React.FC = ({ // in this edit session; undefined when none is held. If a mint-relevant field later diverges from it, // the held token (hook response + sessionStorage) is discarded so the admin must re-authorize. const authorizedIdentityRef = React.useRef(undefined); + const dcrClientRef = React.useRef(null); const { startOAuthFlow, @@ -300,7 +301,7 @@ const MCPServerEdit: React.FC = ({ env: values.env, }; }, - onTokenReceived: (token) => { + onTokenReceived: (token, registeredClient) => { if (!token?.access_token) { return; } @@ -319,6 +320,7 @@ const MCPServerEdit: React.FC = ({ return; } + dcrClientRef.current = registeredClient?.dcr_server_url ? registeredClient : null; const current = (allFieldsValue(form).credentials as Record | undefined) ?? {}; const nextCredentials = { ...(preservedAdminCredentials(current) ?? {}), @@ -393,6 +395,9 @@ const MCPServerEdit: React.FC = ({ if (!parsed || parsed.serverId !== mcpServer.server_id) { return; } + // The saved server may still be loading on the first render after the redirect. + // Consume this snapshot only after the matching server can restore it. + window.sessionStorage.removeItem(EDIT_OAUTH_UI_STATE_KEY); if (parsed.formValues) { // Rebuild credentials from the declared app in EITHER the loaded server or the saved snapshot, // then strip minted token material. Merging the two (server under snapshot) before stripping is @@ -426,7 +431,6 @@ const MCPServerEdit: React.FC = ({ } } catch (err) { console.error("Failed to restore MCP edit state", err); - } finally { window.sessionStorage.removeItem(EDIT_OAUTH_UI_STATE_KEY); } }, [form, mcpServer]); @@ -497,6 +501,7 @@ const MCPServerEdit: React.FC = ({ removeToken(mcpServer.server_id, userID); } setTools([]); + dcrClientRef.current = null; resetOAuthFlow(); // The admin-typed app is upstream-scoped config, not minted material, so it survives every // invalidation; only the held token is discarded. Token-shaped keys are excluded by the filter. @@ -720,7 +725,16 @@ const MCPServerEdit: React.FC = ({ return () => subscription.unsubscribe(); }, [form]); + const isOAuthPending = ["authorizing", "exchanging"].includes(oauthStatus); + + const handleCancel = () => { + window.sessionStorage.removeItem(EDIT_OAUTH_UI_STATE_KEY); + resetOAuthFlow(); + onCancel(); + }; + const submitForm = async () => { + if (isOAuthPending) return; const isValid = await form.trigger(mountedPaths(registry) as string[]); if (!isValid) { return; @@ -742,7 +756,8 @@ const MCPServerEdit: React.FC = ({ return; } try { - const built = buildEditServerPayload(values, { + const uiState = { + dcrClient: dcrClientRef.current, mcpServer, logoUrl, costConfig, @@ -752,7 +767,8 @@ const MCPServerEdit: React.FC = ({ toolNameToDisplayName, toolNameToDescription, removeStoredApp, - }); + }; + const built = buildEditServerPayload(values, uiState); if (built.kind !== "ok") { toast.fromError(editPayloadErrorMessage(built)); return; @@ -820,6 +836,7 @@ const MCPServerEdit: React.FC = ({
{ event.preventDefault(); void submitForm(); @@ -1308,10 +1325,12 @@ const MCPServerEdit: React.FC = ({
- - +
@@ -1323,10 +1342,12 @@ const MCPServerEdit: React.FC = ({
- - +
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx index 37d265c4e35..7aad92f9814 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx @@ -20,7 +20,6 @@ vi.mock("recharts", () => ({
), CartesianGrid: () => null, - Treemap: () => null, XAxis: () => null, YAxis: () => null, })); @@ -45,48 +44,26 @@ const response = { daily_totals: [{ date: "2026-09-28", spend: 2.5, prompt_tokens: 1000, completion_tokens: 2000, requests: 12 }], }; -const taskResponse = { - start_date: "2025-09-29", - end_date: "2026-09-28", - tasks: [ - { - task_type: "code_generation", - label: "Code Generation", - category: "Code", - value: 2.5, - share: 100, - leader: "fast-chat", - provider: "openai", - }, - ], -}; - -const mockApi = (tasks: unknown = taskResponse) => - vi - .mocked(apiClient.get) - .mockImplementation((path: string) => - path === "/model-insights/tasks" ? (tasks as Promise) : Promise.resolve(response), - ); - describe("ModelInsightsView", () => { beforeEach(() => { vi.mocked(apiClient.get).mockReset(); - mockApi(Promise.resolve(taskResponse)); + vi.mocked(apiClient.get).mockResolvedValue(response); }); - it("shows the ranking with share and the task legend from the API response", async () => { + it("shows the ranking with share from the API response", async () => { render(); expect(await screen.findByText("fast-chat")).toBeInTheDocument(); expect(screen.getByText("by openai")).toBeInTheDocument(); - expect(await screen.findByText("Code")).toBeInTheDocument(); - expect(screen.getAllByText("100.0%")).toHaveLength(2); + expect(screen.getAllByText("100.0%")).toHaveLength(1); + expect(screen.queryByText("Top models by task")).not.toBeInTheDocument(); expect(screen.getByRole("tab", { name: "tokens" })).toHaveAttribute("aria-selected", "true"); expect(screen.getByRole("tab", { name: "log" })).toBeInTheDocument(); expect(apiClient.get).toHaveBeenCalledWith("/model-insights", { accessToken: "token", query: { metric: "tokens" }, }); + expect(apiClient.get).not.toHaveBeenCalledWith("/model-insights/tasks", expect.anything()); }); it("refetches with the selected metric so top models are ranked by it", async () => { @@ -103,24 +80,6 @@ describe("ModelInsightsView", () => { ); }); - it("does not refetch the task breakdown when the chart metric changes", async () => { - render(); - await screen.findByText("Code"); - const taskCalls = () => - vi.mocked(apiClient.get).mock.calls.filter(([path]) => path === "/model-insights/tasks").length; - const before = taskCalls(); - - await userEvent.click(screen.getByRole("tab", { name: "requests" })); - await waitFor(() => - expect(apiClient.get).toHaveBeenCalledWith("/model-insights", { - accessToken: "token", - query: { metric: "requests" }, - }), - ); - - expect(taskCalls()).toBe(before); - }); - it("shows the API error instead of loading forever", async () => { vi.mocked(apiClient.get).mockRejectedValue(new Error("Only proxy admins can view deployment-wide model insights")); render(); @@ -133,11 +92,7 @@ describe("ModelInsightsView", () => { render(); await screen.findByText("fast-chat"); let resolve: (value: typeof response) => void = () => {}; - vi.mocked(apiClient.get).mockImplementation((path: string) => - path === "/model-insights/tasks" - ? Promise.resolve(taskResponse) - : new Promise((done) => (resolve = done as typeof resolve)), - ); + vi.mocked(apiClient.get).mockImplementation(() => new Promise((done) => (resolve = done as typeof resolve))); await userEvent.click(screen.getByRole("tab", { name: "spend" })); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx index ab406597589..21a707b8c42 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx @@ -2,8 +2,8 @@ import { Page } from "@/components/shared/Page"; import React from "react"; -import { Bar, BarChart, CartesianGrid, Treemap, XAxis, YAxis } from "recharts"; -import { ArrowDownRight, ArrowUpRight, BarChart3, Layers, Minus } from "lucide-react"; +import { Bar, BarChart, CartesianGrid, XAxis, YAxis } from "recharts"; +import { ArrowDownRight, ArrowUpRight, BarChart3, Minus } from "lucide-react"; import { apiClient } from "@/components/networking"; import { extractErrorMessage } from "@/utils/errorUtils"; @@ -12,7 +12,6 @@ import { PageHeader, PageHeaderDescription, PageHeaderTitle } from "@/components import { Alert, AlertDescription, AlertTitle } from "@/components/ui/alert"; import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card"; import { ChartConfig, ChartContainer, ChartTooltip, ChartTooltipContent } from "@/components/ui/chart"; -import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; import { Skeleton } from "@/components/ui/skeleton"; import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { @@ -22,8 +21,6 @@ import { Granularity, Metric, ModelInsightsResponse, - ModelInsightTasksResponse, - TaskSummary, modelOrder, rankModels, RankedModel, @@ -41,13 +38,6 @@ const PALETTE = [ "#6366f1", "#f97316", ]; -const FALLBACK_COLOR = "#64748b"; -const CATEGORY_COLORS: Record = { - General: "#ee8650", - Agent: "#7666e4", - Code: "#5fb074", - Data: "#3b82f6", -}; const SCALES = ["linear", "log"] as const; const GRANULARITIES = ["day", "week"] as const; const GRANULARITY_LABELS: Record = { day: "Daily", week: "Weekly" }; @@ -90,37 +80,11 @@ const RankingRow = ({ model, rank }: { model: RankedModel; rank: number }) => ( ); -type TileProps = TaskSummary & { x: number; y: number; width: number; height: number; index: number }; - -const TaskTileContent = ({ x, y, width, height, category, label, leader }: TileProps) => { - if (width <= 0 || height <= 0) return null; - const color = CATEGORY_COLORS[category] ?? FALLBACK_COLOR; - const fits = width > 90 && height > 44; - return ( - - - {fits && ( - <> - - {label} - - - {leader} - - - )} - - ); -}; - export default function ModelInsightsView({ accessToken }: { accessToken: string | null }) { const [loaded, setLoaded] = React.useState<{ metric: Metric; response: ModelInsightsResponse } | null>(null); const [metric, setMetric] = React.useState("tokens"); const [scale, setScale] = React.useState("linear"); const [granularity, setGranularity] = React.useState("day"); - const [taskMetric, setTaskMetric] = React.useState("spend"); - const [taskData, setTaskData] = React.useState(null); - const [taskError, setTaskError] = React.useState(null); const [error, setError] = React.useState(null); React.useEffect(() => { @@ -141,24 +105,6 @@ export default function ModelInsightsView({ accessToken }: { accessToken: string }; }, [accessToken, metric]); - React.useEffect(() => { - if (!accessToken) return; - let cancelled = false; - apiClient - .get("/model-insights/tasks", { accessToken, query: { metric: taskMetric } }) - .then((response) => { - if (cancelled) return; - setTaskError(null); - setTaskData(response); - }) - .catch((err: unknown) => { - if (!cancelled) setTaskError(extractErrorMessage(err)); - }); - return () => { - cancelled = true; - }; - }, [accessToken, taskMetric]); - const data = loaded?.response ?? null; const shown = loaded?.metric ?? metric; const isStale = loaded !== null && loaded.metric !== metric; @@ -176,15 +122,6 @@ export default function ModelInsightsView({ accessToken }: { accessToken: string () => (data ? rankModels(data.top_models, data.daily, shown, range) : []), [data, shown, range], ); - const tiles = React.useMemo(() => taskData?.tasks ?? [], [taskData]); - const categoryShares = React.useMemo( - () => - [...new Set(tiles.map((tile) => tile.category))].map((category) => ({ - category, - share: tiles.filter((tile) => tile.category === category).reduce((sum, tile) => sum + tile.share, 0), - })), - [tiles], - ); if (error) { return ( @@ -317,57 +254,6 @@ export default function ModelInsightsView({ accessToken }: { accessToken: string - - -
- - Top models by task - - - Each task's share of {METRIC_LABELS[taskMetric]}, labelled with its leading model - -
- -
- - {taskError && ( - - Could not load tasks - {taskError} - - )} - - ({ ...tile, name: tile.task_type }))} - dataKey="value" - isAnimationActive={false} - content={} - /> - -
    - {categoryShares.map(({ category, share }) => ( -
  • - - {category} - {share.toFixed(1)}% -
  • - ))} -
-
-
- Cost per session diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts index 0e1a5fa7aac..9e0aecc9e8e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts @@ -21,16 +21,6 @@ export type ModelInsightsResponse = { daily_totals: DailyTotal[]; top_models: ModelMetric[]; }; -export type TaskSummary = { - task_type: string; - label: string; - category: string; - value: number; - share: number; - leader: string; - provider: string; -}; -export type ModelInsightTasksResponse = { start_date: string; end_date: string; tasks: TaskSummary[] }; export type RankedModel = { model_group: string; provider: string; share: number; delta: number }; export type Granularity = "day" | "week"; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/MatchedPeopleToggle.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/MatchedPeopleToggle.tsx deleted file mode 100644 index 5c7eaf22384..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/MatchedPeopleToggle.tsx +++ /dev/null @@ -1,12 +0,0 @@ -import { useId } from "react"; -import { Switch } from "@/components/ui/switch"; - -export function MatchedPeopleToggle({ checked, onChange }: { checked: boolean; onChange: (checked: boolean) => void }) { - const id = useId(); - return ( - - ); -} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.integration.test.tsx deleted file mode 100644 index bc242ebe14a..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.integration.test.tsx +++ /dev/null @@ -1,88 +0,0 @@ -import { fireEvent, render, screen, waitFor } from "@testing-library/react"; -import userEvent from "@testing-library/user-event"; -import { afterEach, expect, it, vi } from "vitest"; -import ObservedAccounts from "./ObservedAccounts"; - -afterEach(() => vi.unstubAllGlobals()); -it("links several usernames on both providers to one email in a single save", async () => { - const writes: unknown[] = []; - const saved = vi.fn(); - vi.stubGlobal( - "fetch", - vi.fn(async (_input: string, init: RequestInit) => { - if (init.method === "PUT") { - writes.push(JSON.parse(String(init.body))); - return Response.json({ report: null }); - } - const identities = { - gateway_emails: ["ari@example.test"], - identity_map: { old: "ari@example.test" }, - unmatched_logins: ["new"], - connections: [ - { - id: "github-id", - source_provider: "github", - api_url: "https://api.github.com", - identity_map: { old: "ari@example.test" }, - unmatched_logins: ["new"], - }, - { - id: "gitlab-id", - source_provider: "gitlab", - api_url: "https://gitlab.com/api/v4", - identity_map: {}, - unmatched_logins: ["new"], - }, - ], - }; - return Response.json(identities); - }), - ); - const user = userEvent.setup(); - render(); - await screen.findByLabelText(/GitHub usernames/); - fireEvent.change(screen.getByLabelText("Internal email"), { target: { value: "ari@example.test" } }); - expect(screen.getByLabelText(/GitHub usernames/)).toHaveValue("old"); - fireEvent.change(screen.getByLabelText(/GitHub usernames/), { target: { value: "@Old, new, NEW" } }); - fireEvent.change(screen.getByLabelText(/GitLab usernames/), { target: { value: "new" } }); - await user.click(screen.getByRole("button", { name: "Save accounts" })); - await waitFor(() => expect(saved).toHaveBeenCalledOnce()); - expect(writes).toEqual([ - { - email: "ari@example.test", - accounts: [ - { connection_id: "github-id", login: "old" }, - { connection_id: "github-id", login: "new" }, - { connection_id: "gitlab-id", login: "new" }, - ], - }, - ]); -}); -it("keeps a conflicting link editable", async () => { - const saved = vi.fn(); - vi.stubGlobal( - "fetch", - vi.fn(async (_input: string, init: RequestInit) => - init.method === "PUT" - ? Response.json({ detail: "An account is already linked to another email. Unlink it first" }, { status: 409 }) - : Response.json({ gateway_emails: ["ari@example.test"], identity_map: {}, unmatched_logins: [] }), - ), - ); - const user = userEvent.setup(); - render( - , - ); - fireEvent.change(screen.getByLabelText("Source usernames"), { target: { value: "old, new" } }); - const button = screen.getByRole("button", { name: "Save accounts" }); - await waitFor(() => expect(button).toBeEnabled()); - await user.click(button); - expect(await screen.findByRole("alert")).toHaveTextContent("An account is already linked to another email"); - expect(screen.getByLabelText("Source usernames")).toHaveValue("old, new"); - expect(saved).not.toHaveBeenCalled(); -}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.tsx deleted file mode 100644 index ab6c4cf38df..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.tsx +++ /dev/null @@ -1,209 +0,0 @@ -"use client"; - -import { useEffect, useState } from "react"; -import { z } from "zod"; -import { apiClient } from "@/components/networking"; -import { extractProxyErrorMessage } from "@/lib/http/client"; -import { Button } from "@/components/ui/button"; -import { Input } from "@/components/ui/input"; -import { Dialog, DialogContent, DialogHeader, DialogTitle, DialogDescription } from "@/components/ui/dialog"; -import { accountLogins, type ObservedPerson } from "./observedData"; - -const connectionIdentityFields = { - id: z.string(), - source_provider: z.enum(["github", "gitlab"]), - api_url: z.string(), - identity_map: z.record(z.string(), z.string()), - unmatched_logins: z.array(z.string()), -}; -const identitiesFields = { - gateway_emails: z.array(z.string()), - identity_map: z.record(z.string(), z.string()), - unmatched_logins: z.array(z.string()), - connections: z.array(z.object(connectionIdentityFields)).optional(), -}; -const identitiesSchema = z.object(identitiesFields); - -function matches(identities: z.infer, email: string, people: ObservedPerson[]) { - const person = people.find((entry) => entry.email === email); - return Object.fromEntries( - (identities.connections ?? []).map((entry) => [ - entry.id, - [ - ...new Set([ - ...Object.entries(entry.identity_map) - .filter(([, address]) => address === email) - .map(([login]) => login), - ...(person?.accounts ?? []) - .filter((account) => account.connection_id === entry.id) - .map((account) => account.login), - ]), - ].join(", "), - ]), - ); -} - -export default function ObservedAccounts({ - accessToken, - people, - initialEmail = "", - onClose, - onSaved, -}: { - accessToken: string; - people: ObservedPerson[]; - initialEmail?: string; - onClose: () => void; - onSaved: () => void; -}) { - const [identities, setIdentities] = useState | null>(null); - const [email, setEmail] = useState(initialEmail); - const [logins, setLogins] = useState(people.find((person) => person.email === initialEmail)?.logins.join(", ") ?? ""); - const [linked, setLinked] = useState>({}); - const [saving, setSaving] = useState(false); - const [error, setError] = useState(""); - useEffect(() => { - const controller = new AbortController(); - apiClient - .get("/roi-calculator/observed/identities", { accessToken, signal: controller.signal }) - .then((data) => { - if (!controller.signal.aborted) { - const parsed = identitiesSchema.parse(data); - setIdentities(parsed); - setLinked(matches(parsed, initialEmail, people)); - } - }) - .catch((reason: unknown) => { - if (!controller.signal.aborted) setError(extractProxyErrorMessage(reason)); - }); - return () => controller.abort(); - }, [accessToken, initialEmail, people]); - function selectEmail(value: string) { - setEmail(value); - if (identities) setLinked(matches(identities, value, people)); - const automatic = people.find((person) => person.email === value)?.logins ?? []; - const manual = Object.entries(identities?.identity_map ?? {}) - .filter(([, address]) => address === value) - .map(([login]) => login); - setLogins([...new Set([...automatic, ...manual])].join(", ")); - } - async function save() { - setSaving(true); - setError(""); - try { - await apiClient.put("/roi-calculator/observed/identities", { - accessToken, - body: { - email: email.trim().toLowerCase(), - ...(identities?.connections?.length - ? { - accounts: identities.connections.flatMap((entry) => - accountLogins(linked[entry.id] ?? "").map((login) => ({ connection_id: entry.id, login })), - ), - } - : { logins: accountLogins(logins) }), - }, - }); - onSaved(); - onClose(); - } catch (reason) { - setError(extractProxyErrorMessage(reason)); - } finally { - setSaving(false); - } - } - return ( - { - if (!open) onClose(); - }} - > - - - Link accounts - Match one internal user to all their source accounts - -
-
- - selectEmail(event.target.value)} - /> - - {identities?.gateway_emails.map((address) => -
- {identities?.connections?.length ? ( - identities.connections.map((entry) => ( -
- - setLinked({ ...linked, [entry.id]: event.target.value })} - placeholder="current-account, old-account" - /> - {entry.unmatched_logins.length > 0 && ( -
- {entry.unmatched_logins.length} unmatched accounts -
- {entry.unmatched_logins.map((login) => ( - - ))} -
-
- )} -
- )) - ) : ( -
- - setLogins(event.target.value)} - placeholder="current-account, old-account" - /> -
- )} -

- Separate accounts with commas. Their merged changes are combined, and gateway spend is counted once -

- {error && ( -

- {error} -

- )} - -
-
-
- ); -} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.integration.test.tsx deleted file mode 100644 index 730f0c8faf5..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.integration.test.tsx +++ /dev/null @@ -1,198 +0,0 @@ -import { fireEvent, render, screen, waitFor } from "@testing-library/react"; -import userEvent from "@testing-library/user-event"; -import { afterEach, describe, expect, it, vi } from "vitest"; -import ObservedConnections from "./ObservedConnections"; -import type { ObservedSettings } from "./observedData"; - -const settings: ObservedSettings = { - id: "initial-github", - source_provider: "github", - api_url: "https://api.github.com", - repos: [], - has_token: false, - connection_type: "token", - update_interval_minutes: 1440, - ready: false, -}; -const app = { configured: true, api_url: null, callback_url: null }; -afterEach(() => vi.unstubAllGlobals()); - -describe("observed ROI connections", () => { - it("identifies the saved connection when editing its host", async () => { - const writes: unknown[] = []; - const existing = { ...settings, id: "saved-github", has_token: true, repos: ["org/service"], ready: true }; - vi.stubGlobal( - "fetch", - vi.fn(async (input: string, init: RequestInit) => { - const path = new URL(input, "http://localhost").pathname; - if (path.endsWith("/apps")) return Response.json({ github: app, gitlab: app }); - if (path.endsWith("/repositories")) return Response.json({ repositories: [], has_more: false }); - if (path.endsWith("/settings")) { - writes.push(JSON.parse(String(init.body))); - return Response.json({ ...existing, id: "enterprise-github", api_url: "https://git.example.test/api/v3" }); - } - throw new Error(path); - }), - ); - const user = userEvent.setup(); - render( - , - ); - await user.click(screen.getByRole("button", { name: "Edit GitHub api.github.com" })); - await user.click(screen.getByRole("button", { name: "Change connection" })); - await user.click(screen.getByText("Self-hosted instance")); - fireEvent.change(screen.getByLabelText("API URL"), { target: { value: "https://git.example.test/api/v3" } }); - fireEvent.change(screen.getByLabelText("GitHub access token"), { target: { value: "enterprise-test-token" } }); - await user.click(screen.getByRole("button", { name: "Continue" })); - expect(await screen.findByRole("heading", { name: "Choose repositories" })).toBeInTheDocument(); - expect(writes).toEqual([ - { - connection_id: "saved-github", - source_provider: "github", - api_url: "https://git.example.test/api/v3", - token: "enterprise-test-token", - repos: [], - update_interval_minutes: 1440, - }, - ]); - }); - it("starts a GitHub installation when switching from a GitLab app connection", async () => { - const starts = vi.fn(); - vi.stubGlobal( - "fetch", - vi.fn(async (input: string, init: RequestInit) => { - const url = new URL(input, "http://localhost"); - if (url.pathname.endsWith("/apps")) - return Response.json({ github: { ...app, can_install: true }, gitlab: app }); - if (url.pathname.endsWith("/repositories")) return Response.json({ repositories: [], has_more: false }); - starts(url.searchParams.get("install"), init.credentials); - return Response.json({ detail: "Authorization test stopped before redirect" }, { status: 502 }); - }), - ); - const user = userEvent.setup(); - const connected: ObservedSettings = { - ...settings, - source_provider: "gitlab", - connection_type: "app", - has_token: true, - }; - render(); - await user.click(screen.getByRole("button", { name: "Change connection" })); - await user.click(screen.getByRole("button", { name: "GitHub", exact: true })); - const connect = screen.getByRole("button", { name: "Connect GitHub", exact: true }); - await waitFor(() => expect(connect).toBeEnabled()); - await user.click(connect); - expect(await screen.findByRole("alert")).toHaveTextContent("Authorization test stopped before redirect"); - expect(starts).toHaveBeenCalledWith("true", "include"); - }); - it.each(["GitHub", "GitLab"] as const)( - "connects %s with a token, saves repositories, and starts a sync", - async (label) => { - const provider = label === "GitHub" ? "github" : "gitlab"; - const apiUrl = provider === "github" ? "https://api.github.com" : "https://gitlab.com/api/v4"; - const saved = vi.fn(); - const closed = vi.fn(); - const writes: { path: string; body: unknown }[] = []; - vi.stubGlobal( - "fetch", - vi.fn(async (input: string, init: RequestInit) => { - const path = new URL(input, "http://localhost").pathname; - if (init.method === "PUT" || init.method === "POST") - writes.push({ path, body: init.body ? JSON.parse(String(init.body)) : undefined }); - if (path.endsWith("/apps")) return Response.json({ github: app, gitlab: app }); - if (path.endsWith("/repositories")) - return Response.json({ - repositories: [{ name: "org/service", visibility: "private", archived: false }], - has_more: false, - }); - if (path.endsWith("/settings")) { - const request = JSON.parse(String(init.body)) as { repos: string[] }; - const connected = { - ...settings, - id: `saved-${provider}`, - source_provider: provider, - api_url: apiUrl, - has_token: true, - repos: request.repos, - ready: request.repos.length > 0, - }; - return Response.json(connected); - } - if (path.endsWith("/sync")) return Response.json({ running: true }, { status: 202 }); - throw new Error(path); - }), - ); - const user = userEvent.setup(); - render( - , - ); - await user.click(screen.getByRole("button", { name: label, exact: true })); - await user.click(screen.getByRole("button", { name: "Access token", exact: true })); - fireEvent.change(screen.getByLabelText(`${label} access token`), { target: { value: "source-test-token" } }); - await user.click(screen.getByRole("button", { name: "Continue" })); - await user.click(await screen.findByRole("checkbox", { name: /org\/service/ })); - await user.click(screen.getByRole("button", { name: "Save and sync" })); - await waitFor(() => expect(saved).toHaveBeenCalledOnce()); - expect(closed).toHaveBeenCalledOnce(); - expect(writes).toEqual([ - { - path: "/roi-calculator/observed/settings", - body: { - source_provider: provider, - api_url: apiUrl, - token: "source-test-token", - repos: [], - update_interval_minutes: 1440, - }, - }, - { - path: "/roi-calculator/observed/settings", - body: { - connection_id: `saved-${provider}`, - source_provider: provider, - api_url: apiUrl, - repos: ["org/service"], - update_interval_minutes: 1440, - }, - }, - { path: "/roi-calculator/observed/sync", body: undefined }, - ]); - }, - ); - it.each(["GitHub", "GitLab"] as const)( - "starts %s app authorization with a browser cookie and displays provider failures", - async (label) => { - const provider = label.toLowerCase(); - vi.stubGlobal( - "fetch", - vi.fn(async (input: string, init: RequestInit) => { - if (String(input).endsWith("/apps")) return Response.json({ github: app, gitlab: app }); - expect(String(input)).toContain(`/oauth/${provider}/start`); - expect(init.credentials).toBe("include"); - expect(init.method).toBe("POST"); - return Response.json({ detail: "Provider unavailable. Try again" }, { status: 502 }); - }), - ); - const user = userEvent.setup(); - render( - , - ); - await user.click(screen.getByRole("button", { name: label, exact: true })); - const button = await screen.findByRole("button", { name: `Connect ${label}`, exact: true }); - await waitFor(() => expect(button).toBeEnabled()); - await user.click(button); - expect(await screen.findByRole("alert")).toHaveTextContent("Provider unavailable. Try again"); - expect(button).toBeEnabled(); - }, - ); -}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.tsx deleted file mode 100644 index 591b7066e48..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.tsx +++ /dev/null @@ -1,606 +0,0 @@ -"use client"; - -import { useEffect, useState } from "react"; -import { Github, Gitlab, ArrowLeft, KeyRound } from "lucide-react"; -import { z } from "zod"; -import { apiClient } from "@/components/networking"; -import { extractProxyErrorMessage } from "@/lib/http/client"; -import { Button } from "@/components/ui/button"; -import { Input } from "@/components/ui/input"; -import { Dialog, DialogContent, DialogHeader, DialogTitle, DialogDescription } from "@/components/ui/dialog"; -import { - observedSettingsSchema, - repositoryNames, - type ObservedSettings, - type ObservedConnection, -} from "./observedData"; - -const appFields = { - configured: z.boolean(), - can_install: z.boolean().optional().default(false), - api_url: z.string().nullable(), - callback_url: z.string().nullable(), -}; -const appSchema = z.object(appFields); -const appsSchema = z.object({ github: appSchema, gitlab: appSchema }); -const repositoriesSchema = z.object({ - repositories: z.array(z.object({ name: z.string(), visibility: z.string(), archived: z.boolean() })), - has_more: z.boolean(), -}); -const defaultUrl = { github: "https://api.github.com", gitlab: "https://gitlab.com/api/v4" }; - -function preferredConnectionMethod(selected: "app" | "token" | null, configured: boolean | undefined) { - return selected ?? (configured ? "app" : "token"); -} - -function TokenFields({ - label, - provider, - token, - onToken, - apiUrl, - onUrl, - hasToken, -}: { - label: string; - provider: ObservedSettings["source_provider"]; - token: string; - onToken: (value: string) => void; - apiUrl: string; - onUrl: (value: string) => void; - hasToken: boolean; -}) { - return ( - <> - - onToken(event.target.value)} - placeholder={hasToken ? "Leave blank to keep the saved token" : "Optional for public repositories"} - /> -

- {provider === "github" - ? "Fine-grained token: read access to pull requests, issues, and metadata" - : "Token with read_api scope"} -

-
- Self-hosted instance - - onUrl(event.target.value)} /> -
- - ); -} -function AppMessage({ configured, label }: { configured: boolean; label: string }) { - return ( -

- {configured - ? `Youโ€™ll authorize ${label}, then choose repositories` - : `Register the ${label} app in gateway settings, or connect with a token`} -

- ); -} - -function RepositoryChoices({ - available, - repos, - setRepos, - query, - setQuery, - page, - setPage, -}: { - available: z.infer | null; - repos: string; - setRepos: (value: string) => void; - query: string; - setQuery: (value: string) => void; - page: number; - setPage: (value: number) => void; -}) { - return ( - <> - { - setQuery(event.target.value); - setPage(1); - }} - placeholder="Find repositoriesโ€ฆ" - /> -
- {!available && ( -

- Loading repositoriesโ€ฆ -

- )} - {available?.repositories.length === 0 && ( -

No repositories found

- )} - {available?.repositories - .filter((repo) => !repo.archived) - .map((repo) => ( - - ))} -
-
- - -
- - ); -} - -function ConnectionMethod({ - method, - setMethod, - label, - provider, - token, - setToken, - apiUrl, - setApiUrl, - connected, - apps, - busy, - connect, -}: { - method: "app" | "token"; - setMethod: (value: "app" | "token") => void; - label: string; - provider: ObservedSettings["source_provider"]; - token: string; - setToken: (value: string) => void; - apiUrl: string; - setApiUrl: (value: string) => void; - connected: ObservedConnection; - apps: z.infer | null; - busy: boolean; - connect: () => void; -}) { - const connectLabel = method === "app" ? `Connect ${label}` : "Continue"; - const sameSource = connected.source_provider === provider && connected.api_url === apiUrl; - const hasSavedToken = sameSource && connected.has_token && connected.connection_type === "token"; - return ( - <> -
- - -
- {method === "token" && ( - - )} - {method === "app" && } - - - ); -} - -type ConnectionStep = "list" | "connect" | "repos"; -const stepTitles = { list: "Connections", connect: "Connect your code", repos: "Choose repositories" }; - -function initialStep(settings: ObservedSettings, afterAuthorization: boolean): ConnectionStep { - if (afterAuthorization) return "repos"; - if (settings.connections?.length) return "list"; - return settings.has_token || settings.ready ? "repos" : "connect"; -} - -function stepDescription(step: ConnectionStep, label: string) { - if (step === "list") return "All selected repositories appear in one report"; - if (step === "connect") return "Connect GitHub and GitLab with an app or access token"; - return `Select ${label} repositories to compare`; -} - -function connectionMethodLabel(entry: ObservedConnection) { - if (entry.connection_type === "app") return "App"; - return entry.has_token ? "Token" : "Public access"; -} - -function ConnectionList({ - connections, - onEdit, - onAdd, -}: { - connections: ObservedConnection[]; - onEdit: (entry: ObservedConnection) => void; - onAdd: () => void; -}) { - return ( - <> - {connections.map((entry) => ( -
-
-

- {entry.source_provider === "github" ? : } - {entry.source_provider === "github" ? "GitHub" : "GitLab"} -

-

{new URL(entry.api_url).host}

-

- {entry.repos.length} repositories ยท {connectionMethodLabel(entry)} -

-
- -
- ))} - - - ); -} - -function initialMethod(settings: ObservedSettings) { - return settings.has_token ? settings.connection_type : null; -} - -function hasConnections(settings: ObservedSettings) { - return Boolean(settings.connections?.length); -} - -function canManageApp(connected: ObservedConnection, apps: z.infer | null) { - return ( - connected.connection_type === "app" && connected.source_provider === "github" && Boolean(apps?.github.can_install) - ); -} - -export default function ObservedConnections({ - accessToken, - settings, - onClose, - onSaved, - initialError = "", - afterAuthorization = false, -}: { - accessToken: string; - settings: ObservedSettings; - onClose: () => void; - onSaved: () => void; - initialError?: string; - afterAuthorization?: boolean; -}) { - const [savedSettings, setSavedSettings] = useState(settings); - const [connected, setConnected] = useState(settings); - const [provider, setProvider] = useState(settings.source_provider); - const [apiUrl, setApiUrl] = useState(settings.api_url); - const [selectedMethod, setMethod] = useState<"app" | "token" | null>(initialMethod(settings)); - const [step, setStep] = useState(() => initialStep(settings, afterAuthorization)); - const [token, setToken] = useState(""); - const [repos, setRepos] = useState(settings.repos.join(", ")); - const [apps, setApps] = useState | null>(null); - const method = preferredConnectionMethod(selectedMethod, apps?.[provider].configured); - const [available, setAvailable] = useState | null>(null); - const [query, setQuery] = useState(""); - const [page, setPage] = useState(1); - const [busy, setBusy] = useState(false); - const [error, setError] = useState(initialError); - const label = provider === "github" ? "GitHub" : "GitLab"; - const saveLabel = repositoryNames(repos).length ? "Save and sync" : "Save repositories"; - const manageApp = canManageApp(connected, apps); - const showConnections = hasConnections(savedSettings); - useEffect(() => { - const controller = new AbortController(); - apiClient - .get("/roi-calculator/observed/apps", { accessToken, signal: controller.signal }) - .then((data) => { - if (!controller.signal.aborted) setApps(appsSchema.parse(data)); - }) - .catch((reason: unknown) => { - if (!controller.signal.aborted) setError(extractProxyErrorMessage(reason)); - }); - return () => controller.abort(); - }, [accessToken]); - useEffect(() => { - if (step !== "repos" || (!connected.has_token && connected.source_provider === "github")) return; - const controller = new AbortController(); - const timer = setTimeout(() => { - apiClient - .get("/roi-calculator/observed/repositories", { - accessToken, - signal: controller.signal, - query: { query, page, connection: connected.id }, - }) - .then((data) => { - if (!controller.signal.aborted) setAvailable(repositoriesSchema.parse(data)); - }) - .catch((reason: unknown) => { - if (!controller.signal.aborted) setError(extractProxyErrorMessage(reason)); - }); - }, 250); - return () => { - clearTimeout(timer); - controller.abort(); - }; - }, [accessToken, step, connected, query, page]); - function selectProvider(value: ObservedSettings["source_provider"]) { - setProvider(value); - const existing = savedSettings.connections?.find( - (entry) => entry.source_provider === value && entry.api_url === defaultUrl[value], - ); - setApiUrl(existing?.api_url ?? defaultUrl[value]); - setConnected( - existing ?? { - ...settings, - source_provider: value, - api_url: defaultUrl[value], - repos: [], - has_token: false, - ready: false, - connection_type: "token", - id: undefined, - }, - ); - setMethod(existing?.connection_type ?? null); - setToken(""); - setError(""); - } - function searchRepositories(value: string) { - setAvailable(null); - setError(""); - setQuery(value); - } - function changePage(value: number) { - setAvailable(null); - setError(""); - setPage(value); - } - async function connect(install = false) { - setBusy(true); - setError(""); - try { - if (method === "app" || install) { - const sameApp = connected.source_provider === provider && connected.connection_type === "app"; - const firstInstallation = provider === "github" && !sameApp && apps?.github.can_install; - const result = z.object({ url: z.string().url() }).parse( - await apiClient.post(`/roi-calculator/observed/oauth/${provider}/start`, { - accessToken, - credentials: "include", - query: { install: install || Boolean(firstInstallation) }, - }), - ); - window.location.assign(result.url); - return; - } - const same = provider === connected.source_provider && apiUrl === connected.api_url; - const keepToken = same && connected.has_token && connected.connection_type === "token"; - const result = observedSettingsSchema.parse( - await apiClient.put("/roi-calculator/observed/settings", { - accessToken, - body: { - connection_id: savedSettings.connections?.find((entry) => entry.id === connected.id)?.id, - source_provider: provider, - api_url: apiUrl, - token: token || (keepToken ? undefined : ""), - repos: same ? connected.repos : [], - update_interval_minutes: connected.update_interval_minutes, - }, - }), - ); - setSavedSettings(result); - setConnected(result); - setToken(""); - setRepos(result.repos.join(", ")); - setAvailable(null); - setStep("repos"); - } catch (reason) { - setError(extractProxyErrorMessage(reason)); - } finally { - setBusy(false); - } - } - async function save() { - setBusy(true); - setError(""); - try { - const result = observedSettingsSchema.parse( - await apiClient.put("/roi-calculator/observed/settings", { - accessToken, - body: { - connection_id: connected.id, - source_provider: connected.source_provider, - api_url: connected.api_url, - repos: repositoryNames(repos), - update_interval_minutes: connected.update_interval_minutes, - }, - }), - ); - if (result.ready) await apiClient.post("/roi-calculator/observed/sync", { accessToken }); - onSaved(); - onClose(); - } catch (reason) { - setError(extractProxyErrorMessage(reason)); - } finally { - setBusy(false); - } - } - function edit(entry: ObservedConnection) { - setConnected(entry); - setProvider(entry.source_provider); - setApiUrl(entry.api_url); - setMethod(entry.connection_type); - setRepos(entry.repos.join(", ")); - setToken(""); - setQuery(""); - setPage(1); - setAvailable(null); - setError(""); - setStep("repos"); - } - return ( - { - if (!open) onClose(); - }} - > - - - {stepTitles[step]} - {stepDescription(step, label)} - -
- {step === "list" && ( - { - selectProvider( - savedSettings.connections?.some((entry) => entry.source_provider === "github") ? "gitlab" : "github", - ); - setStep("connect"); - }} - /> - )} - {step === "connect" && ( - <> - {showConnections && ( - - )} -
- {(["github", "gitlab"] as const).map((value) => ( - - ))} -
- connect()} - /> - - )} - {step === "repos" && ( - <> - {showConnections && ( - - )} - - {manageApp && ( - - )} - - setRepos(event.target.value)} - placeholder={ - provider === "github" ? "owner/repo, owner/another-repo" : "group/project, group/subgroup/project" - } - /> - {(connected.has_token || connected.source_provider === "gitlab") && ( - <> - - - )} - - - )} - {error && ( -

- {error} -

- )} -
-
-
- ); -} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedDetails.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedDetails.tsx deleted file mode 100644 index 84a4762f046..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedDetails.tsx +++ /dev/null @@ -1,239 +0,0 @@ -"use client"; - -import { useState } from "react"; -import { ExternalLink, GitPullRequest } from "lucide-react"; -import { Badge } from "@/components/ui/badge"; -import { Button } from "@/components/ui/button"; -import { Input } from "@/components/ui/input"; -import { Sheet, SheetContent, SheetHeader, SheetTitle, SheetDescription } from "@/components/ui/sheet"; -import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; -import { - dateRange, - changeTerms, - duration, - money, - number, - recordedBranches, - type Comparison, - type ObservedPerson, - type ObservedPull, - type ObservedSnapshot, - type Period, -} from "./observedData"; - -export function PullList({ - pulls, - provider, - matchedOnly = false, -}: { - matchedOnly?: boolean; - pulls: ObservedPull[]; - provider: ObservedSnapshot["source_provider"]; -}) { - const terms = changeTerms(provider); - const emptyMessage = matchedOnly - ? "No merged changes from matched people in this period" - : `No ${terms.lower} in this period`; - const [query, setQuery] = useState(""); - const [limit, setLimit] = useState(20); - const filtered = pulls.filter((pull) => - `${pull.number} ${pull.title} ${pull.author} ${pull.repo} ${pull.source_repo} ${pull.source_branch}` - .toLowerCase() - .includes(query.toLowerCase()), - ); - return ( -
- { - setQuery(event.target.value); - setLimit(20); - }} - className="max-w-sm" - /> - - - - {terms.requests} - Author - Opened to merged - Tagged spend - - - - {filtered.slice(0, limit).map((pull) => ( - - - - - - #{pull.number} - {pull.title} - {pull.repo} - - - - - - {pull.author || "Deleted author"} - {pull.agent && ( - - Agent - - )} - - {duration(pull.merge_hours)} - - {money(pull.branch_cost.spend)} - - - ))} - -
-

- Elapsed time from opening to merge, not engineering effort or time saved -

- {filtered.length === 0 && ( -

- {query ? `No ${terms.lower} match this search` : emptyMessage} -

- )} -
- - {number(Math.min(limit, filtered.length))} of {number(filtered.length)} {terms.lower} - - {limit < filtered.length && ( - - )} -
-
- ); -} - -export function PersonDetails({ - person, - snapshot, - comparison, - onClose, - onEdit, -}: { - person: ObservedPerson; - snapshot: ObservedSnapshot; - comparison: Comparison; - onClose: () => void; - onEdit?: () => void; -}) { - const terms = changeTerms(snapshot.source_provider); - const [period, setPeriod] = useState("current"); - const current = person.periods.current; - const baseline = person.periods[comparison]; - const urls = new Set(person.periods[period].pr_urls); - const pulls = snapshot.pulls[period].filter((pull) => urls.has(pull.url)); - return ( - { - if (!open) onClose(); - }} - > - - - {person.name} - {[person.email, person.logins.join(", ")].filter(Boolean).join(" ยท ")} - - {onEdit && ( - - )} -
-
-

Merged {terms.plural}

-

{number(current.merged_prs)}

-

{number(baseline.merged_prs)} in comparison

-
-
-

Recorded spend

-

- {money(current.spend_observation === "no_records" ? null : current.gateway_recorded_spend)} -

-

Gateway only

-
-
-

Spend / matched {terms.singular}

-

{money(current.recorded_spend_per_attributed_pr)}

-

Period average

-
-
-
-

Merged {terms.plural}

-
- - -
-
-

{dateRange(snapshot.periods[period].window)} ยท UTC

- -
-
- ); -} - -export function BranchSpend({ snapshot, matchedOnly = true }: { snapshot: ObservedSnapshot; matchedOnly?: boolean }) { - const rows = recordedBranches(snapshot, matchedOnly); - return ( -
- - - - Repository - Branch - Requests - Tagged spend - - - - {rows.map((row) => ( - - {row.repo} - {row.branch} - {number(row.requests)} - {money(row.spend)} - - ))} - -
- {rows.length === 0 && ( -

- {matchedOnly - ? "No tagged branch spend for matched people in this period" - : "No tagged branch spend in this period"} -

- )} -
- ); -} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.integration.test.tsx deleted file mode 100644 index b16fe5aa7d3..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.integration.test.tsx +++ /dev/null @@ -1,384 +0,0 @@ -import { fireEvent, render as renderWithoutNuqs, screen, waitFor, within } from "@testing-library/react"; -import userEvent from "@testing-library/user-event"; -import { afterEach, describe, expect, it, vi } from "vitest"; -import { NuqsAdapter } from "nuqs/adapters/react"; -import ObservedROIView from "./ObservedROIView"; -import { createObservedDemo } from "./observedDemo"; -import type { ObservedSettings, ObservedSnapshot, ObservedStatus } from "./observedData"; - -const render = (ui: Parameters[0]) => renderWithoutNuqs(ui, { wrapper: NuqsAdapter }); - -const settings: ObservedSettings = { - source_provider: "gitlab", - api_url: "https://gitlab.com/api/v4", - repos: ["org/service"], - has_token: true, - connection_type: "app", - update_interval_minutes: 1440, - ready: true, -}; -const idle: ObservedStatus = { - running: false, - phase: "complete", - stage: "", - done: 0, - total: 0, - error: null, - finished_at: null, -}; -const period = { - window: { start: "2026-09-01", end: "2026-09-28" }, - merged_prs: 1, - median_merge_hours: 16 / 3600, - human_authored: 1, - agent_authored: 0, - missing_author: 0, - agents_without_requester: 0, - matched_internal_prs: 0, - new_bug_labeled_issues: 0, - new_regression_labeled_issues: 0, - explicitly_titled_revert_prs: 0, - matched_users_recorded_spend: 0, - spend_observation: "no_records" as const, - human_summary: { median_merge_hours: 16 / 3600 }, -}; -const report: ObservedSnapshot = { - source_provider: "gitlab", - repos: ["org/service"], - unmatched_logins: [], - unlinked_branches: [], - captured_at: "2026-09-29T00:00:00Z", - periods: { current: period, previous: period, last_year: period }, - people: [], - pulls: { current: [], previous: [], last_year: [] }, -}; -const app = { configured: true, api_url: null, callback_url: null }; - -afterEach(() => { - vi.unstubAllGlobals(); - window.history.replaceState(null, "", "/"); -}); - -describe("observed ROI dashboard", () => { - it("defaults every contributor list to matched people and keeps the switch across tabs", async () => { - const sample = createObservedDemo(7); - const outside = { - ...sample.pulls.current[0], - url: "https://gitlab.com/outside/api/-/merge_requests/999", - author: "outside", - agent: false, - title: "Outside change", - source_branch: "outside-only", - branch_cost: { - repo: "gitlab.com/outside/api", - branch: "outside-only", - spend: 123, - requests: 5, - status: "matched" as const, - }, - }; - const report = { - ...sample, - pulls: { ...sample.pulls, current: [...sample.pulls.current, outside] }, - unlinked_branches: [{ repo: "gitlab.com/outside/api", branch: "orphan-only", spend: 5, requests: 1 }], - }; - const requests = vi.fn(async (input: string, _init: RequestInit) => { - const path = new URL(input, "http://localhost").pathname; - if (path.endsWith("/settings")) return Response.json(settings); - if (path.endsWith("/report")) return Response.json({ report }); - if (path.endsWith("/sync")) return Response.json(idle); - throw new Error(path); - }); - vi.stubGlobal("fetch", requests); - const user = userEvent.setup(); - render(); - expect(await screen.findByRole("switch", { name: "Matched people only" })).toBeChecked(); - expect(screen.getByRole("tab", { name: "Engineers 3" })).toBeInTheDocument(); - expect(screen.queryByText("outside")).not.toBeInTheDocument(); - await user.click(screen.getByRole("switch", { name: "Matched people only" })); - expect(screen.getByRole("tab", { name: "Engineers 4" })).toBeInTheDocument(); - await user.click(screen.getByRole("button", { name: "View outside's merged changes" })); - expect(await screen.findByRole("dialog", { name: "outside" })).toHaveTextContent("Outside change"); - expect(screen.queryByRole("button", { name: "Edit linked accounts" })).not.toBeInTheDocument(); - await user.click(screen.getByRole("button", { name: "Close" })); - await user.click(screen.getByRole("switch", { name: "Matched people only" })); - await user.click(screen.getByRole("tab", { name: "Merged changes" })); - expect(screen.queryByText("Outside change")).not.toBeInTheDocument(); - await user.click(screen.getByRole("tab", { name: "Branch spend" })); - expect(screen.getAllByText(/feature\/sample-/).length).toBeGreaterThan(0); - expect(screen.queryByText("outside-only")).not.toBeInTheDocument(); - expect(screen.queryByText("orphan-only")).not.toBeInTheDocument(); - await user.click(screen.getByRole("switch", { name: "Matched people only" })); - expect(screen.getByText("outside-only")).toBeInTheDocument(); - expect(screen.getByText("orphan-only")).toBeInTheDocument(); - await user.click(screen.getByRole("tab", { name: "Merged changes" })); - expect(screen.getByText("Outside change")).toBeInTheDocument(); - expect(requests.mock.calls.every(([, init]) => init.method === "GET")).toBe(true); - }); - - it("shows a useful empty matched view and exposes all changes when the filter is off", async () => { - const sample = { ...createObservedDemo(7), people: [] }; - vi.stubGlobal( - "fetch", - vi.fn(async (input: string) => { - const path = new URL(input, "http://localhost").pathname; - if (path.endsWith("/settings")) return Response.json(settings); - if (path.endsWith("/report")) return Response.json({ report: sample }); - return Response.json(idle); - }), - ); - const user = userEvent.setup(); - render(); - expect(await screen.findByText("No merged changes from matched people in this period")).toBeInTheDocument(); - await user.click(screen.getByRole("switch", { name: "Matched people only" })); - expect(screen.queryByText("No merged changes from matched people in this period")).not.toBeInTheDocument(); - expect(screen.getAllByRole("link", { name: /Add repository search/ }).length).toBeGreaterThan(0); - }); - - it("previews every sample view before setup, changes sample periods without writes, and exits back to setup", async () => { - const requests = vi.fn(async (input: string, _init: RequestInit) => { - const path = new URL(input, "http://localhost").pathname; - const disconnected = { ...settings, ready: false, repos: [], has_token: false }; - if (path.endsWith("/settings")) return Response.json(disconnected); - if (path.endsWith("/report")) return Response.json({ report: null }); - if (path.endsWith("/sync")) return Response.json(idle); - throw new Error(path); - }); - vi.stubGlobal("fetch", requests); - const user = userEvent.setup(); - render(); - expect(await screen.findByRole("heading", { name: "Connect your repositories" })).toBeInTheDocument(); - expect(screen.queryByText("Ready to sync")).not.toBeInTheDocument(); - await user.click(screen.getByRole("button", { name: "Preview sample report" })); - expect(screen.getByRole("status")).toHaveTextContent("Youโ€™re viewing demo data"); - expect(screen.getByRole("tab", { name: "Engineers 3", selected: true })).toBeInTheDocument(); - expect(screen.queryByRole("button", { name: "Connections" })).not.toBeInTheDocument(); - expect(screen.queryByRole("button", { name: "Link accounts" })).not.toBeInTheDocument(); - await waitFor(() => expect(window.location.search).toBe("?demo=1")); - await user.click(screen.getByRole("button", { name: "View Alex Rivera's merged changes" })); - expect(await screen.findByRole("dialog", { name: "Alex Rivera" })).toHaveTextContent("alex-demo@example.com"); - expect(screen.getByRole("heading", { name: "Merged changes" })).toBeInTheDocument(); - expect(screen.queryByRole("button", { name: "Edit linked accounts" })).not.toBeInTheDocument(); - await user.click(screen.getByRole("button", { name: "Close" })); - await user.click(screen.getByRole("tab", { name: "Merged changes" })); - expect(screen.getByRole("img", { name: /Merged changes by week/ })).toBeInTheDocument(); - await user.click(screen.getByRole("tab", { name: "Quality" })); - expect(screen.getByText("New regression-labeled issues")).toBeInTheDocument(); - await user.click(screen.getByRole("tab", { name: "Branch spend" })); - expect(screen.getAllByText(/feature\/sample-/).length).toBeGreaterThan(0); - await user.click(screen.getByRole("combobox", { name: "Reporting period" })); - await user.click(await screen.findByRole("option", { name: "Last 7 days" })); - expect(screen.getByRole("combobox", { name: "Reporting period" })).toHaveTextContent("Last 7 days"); - await user.click(screen.getByRole("combobox", { name: "Comparison period" })); - await user.click(await screen.findByRole("option", { name: "vs. same period last year" })); - expect(screen.getByRole("combobox", { name: "Comparison period" })).toHaveTextContent("vs. same period last year"); - await user.click(screen.getByRole("button", { name: "Exit demo" })); - expect(screen.getByRole("heading", { name: "Connect your repositories" })).toBeInTheDocument(); - expect(screen.getByRole("button", { name: "Connect GitHub or GitLab" })).toBeEnabled(); - await waitFor(() => expect(window.location.search).toBe("")); - expect(requests.mock.calls.every(([, init]) => init.method === "GET")).toBe(true); - }); - - it("keeps live sync and its report intact when entering and exiting the demo", async () => { - const requests = vi.fn(async (input: string, _init: RequestInit) => { - const path = new URL(input, "http://localhost").pathname; - if (path.endsWith("/settings")) return Response.json(settings); - if (path.endsWith("/report")) return Response.json({ report }); - if (path.endsWith("/sync")) return Response.json({ ...idle, running: true, stage: "Reading changes" }); - throw new Error(path); - }); - vi.stubGlobal("fetch", requests); - window.history.replaceState(null, "", "/roi-calculator/?from=review#report"); - const user = userEvent.setup(); - render(); - expect(await screen.findByRole("button", { name: "Cancel sync" })).toBeEnabled(); - await user.click(screen.getByRole("button", { name: "Preview sample report" })); - expect(screen.queryByRole("button", { name: "Cancel sync" })).not.toBeInTheDocument(); - expect(screen.getByRole("button", { name: "2 repositories" })).toBeInTheDocument(); - await waitFor(() => expect(window.location.search).toBe("?from=review&demo=1")); - await user.click(screen.getByRole("button", { name: "Exit demo" })); - expect(screen.getByRole("button", { name: "Cancel sync" })).toBeEnabled(); - expect(screen.getByRole("button", { name: "1 repository" })).toBeInTheDocument(); - expect(screen.getByRole("tab", { name: "Merge requests", selected: true })).toBeInTheDocument(); - await waitFor(() => expect(window.location.search).toBe("?from=review")); - expect(window.location.hash).toBe("#report"); - expect(requests.mock.calls.every(([, init]) => init.method === "GET")).toBe(true); - }); - - it.each(["failed", "pending"])("opens a demo URL even when live requests are %s", async (state) => { - window.history.replaceState(null, "", "/roi-calculator/?demo=1"); - const pending = Promise.withResolvers(); - vi.stubGlobal( - "fetch", - vi.fn(() => (state === "failed" ? Promise.reject(new Error("Live data unavailable")) : pending.promise)), - ); - const user = userEvent.setup(); - render(); - expect(screen.getByRole("status")).toHaveTextContent("Youโ€™re viewing demo data"); - expect(screen.getByText("Alex Rivera")).toBeInTheDocument(); - expect(screen.queryByRole("alert")).not.toBeInTheDocument(); - await user.click(screen.getByRole("button", { name: "Exit demo" })); - expect(screen.queryByText("Alex Rivera")).not.toBeInTheDocument(); - expect(screen.queryByRole("heading", { name: "Connect your repositories" })).not.toBeInTheDocument(); - await waitFor(() => expect(window.location.search).toBe("")); - if (state === "failed") expect(await screen.findByRole("alert")).toHaveTextContent("Live data unavailable"); - }); - - it("retries failures, keeps the report during cancellation, and refreshes after completion", async () => { - let status: ObservedStatus = { ...idle, phase: "error", error: "Provider temporarily unavailable" }; - let completeOnPoll = false; - let currentReport = { ...report, repos: ["org/service", "org/docs"] }; - vi.stubGlobal( - "fetch", - vi.fn(async (input: string, init: RequestInit) => { - const path = new URL(input, "http://localhost").pathname; - if (path.endsWith("/settings")) return Response.json(settings); - if (path.endsWith("/report")) return Response.json({ report: currentReport }); - if (path.endsWith("/sync")) { - if (init.method === "POST") status = { ...idle, running: true, phase: "pulls" }; - if (init.method === "DELETE") status = { ...idle, phase: "cancelled" }; - if (init.method === "GET" && completeOnPoll) { - status = { ...idle, finished_at: "2026-09-29T00:01:00Z" }; - currentReport = { ...report, repos: ["org/updated"] }; - } - return Response.json(status); - } - throw new Error(path); - }), - ); - const user = userEvent.setup(); - render(); - expect(await screen.findByRole("tab", { name: "Merge requests", selected: true })).toBeInTheDocument(); - await user.click(screen.getByRole("tab", { name: "Quality" })); - expect(screen.getByRole("alert")).toHaveTextContent("Provider temporarily unavailable"); - await user.click(screen.getByRole("button", { name: "Retry" })); - await user.click(await screen.findByRole("button", { name: "Cancel sync" })); - expect(await screen.findByRole("button", { name: "Sync now" })).toBeEnabled(); - expect(screen.queryByText("org/service")).not.toBeInTheDocument(); - await user.click(screen.getByRole("button", { name: "2 repositories" })); - const repositories = await screen.findByRole("dialog", { name: "Repositories" }); - expect(within(repositories).getByText("org/service")).toBeInTheDocument(); - expect(within(repositories).getByText("org/docs")).toBeInTheDocument(); - await user.keyboard("{Escape}"); - await waitFor(() => expect(screen.queryByRole("dialog", { name: "Repositories" })).not.toBeInTheDocument()); - await user.click(screen.getByRole("button", { name: "Sync now" })); - expect(await screen.findByRole("button", { name: "Cancel sync" })).toBeEnabled(); - completeOnPoll = true; - await user.click(await screen.findByRole("button", { name: "1 repository" }, { timeout: 4000 })); - expect( - await within(screen.getByRole("dialog", { name: "Repositories" })).findByText("org/updated"), - ).toBeInTheDocument(); - await user.keyboard("{Escape}"); - expect(screen.getByRole("tab", { name: "Quality", selected: true })).toBeInTheDocument(); - expect(screen.getByRole("button", { name: "Sync now" })).toBeEnabled(); - }); - - it("syncs the selected range and labels its equal-length comparison", async () => { - let currentReport = report; - const requested: string[] = []; - vi.stubGlobal( - "fetch", - vi.fn(async (input: string, init: RequestInit) => { - const url = new URL(input, "http://localhost"); - if (url.pathname.endsWith("/settings")) return Response.json(settings); - if (url.pathname.endsWith("/report")) return Response.json({ report: currentReport }); - if (url.pathname.endsWith("/sync")) { - if (init.method === "POST") { - requested.push(url.searchParams.get("days") ?? ""); - currentReport = { - ...report, - periods: { - ...report.periods, - current: { ...period, window: { start: "2026-09-22", end: "2026-09-28" } }, - previous: { ...period, window: { start: "2026-09-15", end: "2026-09-21" } }, - }, - }; - } - return Response.json(idle); - } - throw new Error(url.pathname); - }), - ); - const user = userEvent.setup(); - render(); - await user.click(await screen.findByRole("combobox", { name: "Reporting period" })); - await user.click(await screen.findByRole("option", { name: "Last 7 days" })); - await waitFor(() => expect(requested).toEqual(["7"])); - expect(await screen.findByText(/Comparing with Sep 15.*Sep 21/)).toBeInTheDocument(); - expect(screen.getByRole("combobox", { name: "Reporting period" })).toHaveTextContent("Last 7 days"); - expect(screen.getByRole("combobox", { name: "Comparison period" })).toHaveTextContent("vs. previous period"); - }); - - it("shows a successful empty repository without a setup prompt or invented durations", async () => { - const emptyPeriod = { - ...period, - merged_prs: 0, - human_authored: 0, - median_merge_hours: null, - human_summary: { median_merge_hours: null }, - }; - const empty = { ...report, periods: { current: emptyPeriod, previous: emptyPeriod, last_year: emptyPeriod } }; - vi.stubGlobal( - "fetch", - vi.fn(async (input: string) => { - const path = new URL(input, "http://localhost").pathname; - if (path.endsWith("/settings")) return Response.json(settings); - if (path.endsWith("/report")) return Response.json({ report: empty }); - if (path.endsWith("/sync")) return Response.json(idle); - throw new Error(path); - }), - ); - render(); - expect(await screen.findByRole("heading", { name: "No merged changes yet" })).toBeInTheDocument(); - expect(screen.getByText("No merges")).toBeInTheDocument(); - expect(screen.getByRole("button", { name: "Sync now" })).toBeEnabled(); - expect(screen.queryByRole("alert")).not.toBeInTheDocument(); - expect(screen.queryByRole("heading", { name: "Connect your repositories" })).not.toBeInTheDocument(); - expect(screen.queryByText("0h")).not.toBeInTheDocument(); - }); - - it.each([ - { query: "connected=gitlab", alerts: [] }, - { query: "connection_failed=1", alerts: ["Connection failed or expired. Try again or use a token"] }, - { query: "connection_cancelled=1", alerts: ["Connection cancelled. Choose an app or token to try again"] }, - ])("resumes setup after $query and refreshes saved changes after closing", async ({ query, alerts }) => { - window.history.replaceState(null, "", `/roi-calculator/?${query}`); - let currentSettings = settings; - vi.stubGlobal( - "fetch", - vi.fn(async (input: string, init: RequestInit) => { - const path = new URL(input, "http://localhost").pathname; - if (path.endsWith("/apps")) return Response.json({ github: app, gitlab: app }); - if (path.endsWith("/repositories")) return Response.json({ repositories: [], has_more: false }); - if (path.endsWith("/settings")) { - if (init.method === "PUT") currentSettings = { ...settings, repos: ["org/changed"] }; - return Response.json(currentSettings); - } - if (path.endsWith("/report")) return Response.json({ report: { ...report, repos: currentSettings.repos } }); - if (path.endsWith("/sync")) - return init.method === "POST" - ? Response.json({ detail: "Provider unavailable" }, { status: 502 }) - : Response.json(idle); - throw new Error(path); - }), - ); - const user = userEvent.setup(); - render(); - const dialog = await screen.findByRole("dialog", { name: "Choose repositories" }); - expect( - within(dialog) - .queryAllByRole("alert") - .map((alert) => alert.textContent), - ).toEqual(alerts); - await waitFor(() => expect(window.location.search).toBe("")); - fireEvent.change(within(dialog).getByLabelText("Repositories"), { target: { value: "org/changed" } }); - await user.click(within(dialog).getByRole("button", { name: "Save and sync" })); - expect(await within(dialog).findByRole("alert")).toHaveTextContent("Provider unavailable"); - await user.click(within(dialog).getByRole("button", { name: "Close" })); - await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument()); - await user.click(await screen.findByRole("button", { name: "1 repository" })); - expect( - await within(screen.getByRole("dialog", { name: "Repositories" })).findByText("org/changed"), - ).toBeInTheDocument(); - }); -}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.tsx deleted file mode 100644 index 4a9d20fb870..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.tsx +++ /dev/null @@ -1,252 +0,0 @@ -"use client"; - -import { useEffect, useState } from "react"; -import { parseAsString, useQueryStates } from "nuqs"; -import { Link2, RefreshCw } from "lucide-react"; -import { apiClient } from "@/components/networking"; -import { extractProxyErrorMessage } from "@/lib/http/client"; -import { Page } from "@/components/shared/Page"; -import { PageHeader, PageHeaderDescription, PageHeaderTitle } from "@/components/shared/PageHeader"; -import { DemoNotice } from "@/components/shared/DemoNotice"; -import { Button } from "@/components/ui/button"; -import { Skeleton } from "@/components/ui/skeleton"; -import ObservedConnections from "./ObservedConnections"; -import ObservedReport from "./ObservedReport"; -import { useObservedReport, type ObservedViewData } from "./useObservedReport"; -import { syncMessage, type ObservedSnapshot } from "./observedData"; -import { createObservedDemo } from "./observedDemo"; -import { parseAsDemoFlag } from "./demoUrlState"; - -const OBSERVED_ROI_QUERY_PARSERS = { - demo: parseAsDemoFlag, - connected: parseAsString, - connection_cancelled: parseAsString, - connection_failed: parseAsString, -}; - -function SyncActions({ - data, - error, - busy, - readOnly, - compact = false, - onSync, - onRetry, -}: { - data: ObservedViewData | null; - error: string; - busy: boolean; - readOnly: boolean; - compact?: boolean; - onSync: (cancel: boolean) => void; - onRetry: () => void; -}) { - const message = error || data?.status.error; - const statusMessage = data ? syncMessage(data.status, data.report) : ""; - const canSync = data?.settings.ready && !readOnly; - if (!message && !statusMessage && !canSync) return null; - return ( -
- {message && ( -
- {message} - -
- )} - {data && ( -
- {statusMessage && ( - - {statusMessage} - - )} - {!readOnly && data.settings.ready && ( - - )} -
- )} -
- ); -} - -function EmptyReport({ - data, - readOnly, - onConnect, -}: { - data: ObservedViewData; - readOnly: boolean; - onConnect: () => void; -}) { - function title() { - if (data.status.running) return "Reading repository activity"; - return data.settings.ready ? "Ready for your first report" : "Connect your repositories"; - } - return ( -
-

{title()}

-

- {data.status.running - ? "Your report will appear here when the first sync finishes" - : "Compare merged changes, issue trends, and recorded AI spend across your team"} -

- {!readOnly && !data.status.running && ( - - )} -
- ); -} - -export default function ObservedROIView({ - accessToken, - isViewOnly = false, -}: { - accessToken: string; - isViewOnly?: boolean; -}) { - const { data, error, refresh } = useObservedReport(accessToken); - const [{ demo, connected, connection_cancelled, connection_failed }, setQueryParams] = - useQueryStates(OBSERVED_ROI_QUERY_PARSERS); - const [sample, setSample] = useState(() => (demo === true ? createObservedDemo(28) : null)); - const [connections, setConnections] = useState( - ["github", "gitlab"].includes(connected ?? "") || connection_cancelled !== null || connection_failed !== null, - ); - const [connectionError, setConnectionError] = useState(() => { - if (connection_failed !== null) return "Connection failed or expired. Try again or use a token"; - if (connection_cancelled !== null) return "Connection cancelled. Choose an app or token to try again"; - return ""; - }); - const [afterAuthorization, setAfterAuthorization] = useState(Boolean(connected)); - const [busy, setBusy] = useState(false); - const [actionError, setActionError] = useState(""); - useEffect(() => { - setQueryParams({ connected: null, connection_cancelled: null, connection_failed: null }); - }, [setQueryParams]); - function previewSample(enabled: boolean) { - setQueryParams({ demo: enabled ? true : null }); - setSample(enabled ? createObservedDemo(28) : null); - } - function closeConnections() { - setConnections(false); - setConnectionError(""); - setAfterAuthorization(false); - refresh(); - } - async function sync(cancel: boolean, days?: number) { - setBusy(true); - setActionError(""); - try { - if (cancel) await apiClient.delete("/roi-calculator/observed/sync", { accessToken }); - else await apiClient.post("/roi-calculator/observed/sync", { accessToken, query: { days } }); - refresh(); - } catch (reason) { - setActionError(extractProxyErrorMessage(reason)); - } finally { - setBusy(false); - } - } - function retry() { - if (error || isViewOnly || !data?.settings.ready) refresh(); - else void sync(data.status.running); - } - if (sample) { - return ( - setConnections(true)} - actions={null} - notice={ previewSample(false)} />} - syncing={false} - onPeriod={(days) => setSample(createObservedDemo(days))} - /> - ); - } - const previewButton = ( - - ); - const actions = ( - <> - {previewButton} - - - ); - const content = data?.report ? ( - setConnections(true)} - actions={actions} - syncing={busy || data.status.running} - onPeriod={isViewOnly ? undefined : (days) => void sync(false, days)} - /> - ) : ( - - -
- ROI Calculator - {previewButton} -
- Are we shipping more, with fewer bugs, at a better cost? -
- - {!data && !error && ( - <> - - - - )} - {data && setConnections(true)} />} -
- ); - const showConnections = connections && data && !isViewOnly; - return ( - <> - {content} - {showConnections && ( - - )} - - ); -} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedReport.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedReport.tsx deleted file mode 100644 index 630aafa555d..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedReport.tsx +++ /dev/null @@ -1,631 +0,0 @@ -"use client"; - -import { useMemo, useState } from "react"; -import { ArrowDown, ArrowUp, CalendarDays, ChevronDown, ChevronRight, Link2, Search, Users } from "lucide-react"; -import { Page, PageTabsList, PageTabsTrigger } from "@/components/shared/Page"; -import { PageHeader, PageHeaderTitle } from "@/components/shared/PageHeader"; -import { Button } from "@/components/ui/button"; -import { Input } from "@/components/ui/input"; -import { Popover, PopoverContent, PopoverTitle, PopoverTrigger } from "@/components/ui/popover"; -import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; -import { Tabs, TabsContent } from "@/components/ui/tabs"; -import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; -import ObservedAccounts from "./ObservedAccounts"; -import { MatchedPeopleToggle } from "./MatchedPeopleToggle"; -import { BranchSpend, PersonDetails, PullList } from "./ObservedDetails"; -import { - change, - dateRange, - changeTerms, - duration, - money, - number, - visiblePeople, - reportPeople, - filterObservedPulls, - weeklyMerges, - type Comparison, - type ReportPerson, - type ObservedSnapshot, - type PeopleSort, -} from "./observedData"; - -function Delta({ - current, - baseline, - neutral = false, -}: { - current: number | null; - baseline: number | null; - neutral?: boolean; -}) { - const delta = current === null || baseline === null ? null : change(current, baseline); - if (delta === null) return No baseline; - const Icon = delta >= 0 ? ArrowUp : ArrowDown; - return ( - = 0 ? "increase" : "decrease"}`} - className={`inline-flex items-center gap-1 text-xs tabular-nums ${neutral ? "text-muted-foreground" : "text-foreground"}`} - > - - {number(Math.abs(delta))}% - - ); -} - -function Metric({ - label, - value, - detail, - current, - baseline, -}: { - label: string; - value: string; - detail: string; - current?: number | null; - baseline?: number | null; -}) { - return ( -
-
{label}
-
- {value} - {current !== undefined && baseline !== undefined && } -
-

{detail}

-
- ); -} - -function ShippingTrend({ snapshot, comparison }: { snapshot: ObservedSnapshot; comparison: Comparison }) { - const terms = changeTerms(snapshot.source_provider); - const current = weeklyMerges(snapshot, "current"); - const baseline = weeklyMerges(snapshot, comparison); - const max = Math.max(1, ...current, ...baseline); - return ( -
-
-

Repository shipping activity

-
- - - Current period - - - - {comparison === "previous" ? "Previous period" : "Last year"} - -
-
-
- {current.map((value, week) => ( -
-
-
- - {baseline[week]} - -
-
- {value} -
-
-

W{week + 1}

-
- ))} -
-
- ); -} - -function PeopleTable({ - rows, - provider, - matchedOnly, - comparison, - onSelect, -}: { - rows: ReportPerson[]; - provider: ObservedSnapshot["source_provider"]; - matchedOnly: boolean; - comparison: Comparison; - onSelect: (person: ReportPerson) => void; -}) { - const terms = changeTerms(provider); - const [query, setQuery] = useState(""); - const [sort, setSort] = useState("merged"); - const people = visiblePeople(rows, query, sort); - const emptyMessage = matchedOnly - ? "No matched people. Link accounts or turn off the filter to see all contributors" - : "No contributors in these periods"; - return ( -
-
-
- - setQuery(event.target.value)} - className="pl-9" - /> -
-
- - {people.length} {people.length === 1 ? "engineer" : "engineers"} - - -
-
-
- - - - Engineer - Merged {terms.plural} - Authored / agent - - {comparison === "previous" ? "vs. previous" : "vs. last year"} - - Median merge - Recorded spend - Spend / {terms.singular} - - Details - - - - - {people.map((person) => { - const current = person.periods.current; - const baseline = person.periods[comparison]; - return ( - - - - - {number(current.merged_prs)} - -
-
- - -
- - {current.direct_authored}/{current.declared_agent_owned} - -
-
- - - ({baseline.merged_prs}) - - {duration(current.median_merge_hours)} - - {money(current.spend_observation === "no_records" ? null : current.gateway_recorded_spend)} - - - {money(current.recorded_spend_per_attributed_pr)} - - - - -
- ); - })} -
-
- {people.length === 0 && ( -
- {query ? `No engineers match โ€œ${query}โ€` : emptyMessage} -
- )} -
-

- - - Authored - - - - Agent, explicit requester - - Spend / {terms.singular} is recorded period spend divided by matched {terms.plural} -

-
- ); -} - -function Quality({ snapshot, comparison }: { snapshot: ObservedSnapshot; comparison: Comparison }) { - const terms = changeTerms(snapshot.source_provider); - const current = snapshot.periods.current; - const baseline = snapshot.periods[comparison]; - const rows = [ - { - label: "New bug-labeled issues", - current: current.new_bug_labeled_issues, - baseline: baseline.new_bug_labeled_issues, - detail: "Opened during the period, with bug or kind:bug labels at collection", - }, - { - label: "New regression-labeled issues", - current: current.new_regression_labeled_issues, - baseline: baseline.new_regression_labeled_issues, - detail: "Opened during the period and labeled as regressions", - }, - { - label: `Revert-titled ${terms.plural}`, - current: current.explicitly_titled_revert_prs, - baseline: baseline.explicitly_titled_revert_prs, - detail: `Merged ${terms.plural} whose titles explicitly indicate a revert`, - }, - ]; - return ( -
-
- - - - Repository signal - Current - Comparison - Change - - - - {rows.map((row) => ( - - -

{row.label}

-

{row.detail}

-
- {number(row.current)} - {number(row.baseline)} - - - -
- ))} -
-
-
-

- These signals help check whether more shipping comes with more bugs. Labels and revert titles are incomplete - proxies; they do not establish a change-failure rate or attribute bugs to an engineer. -

-
- ); -} - -function costPerChange(period: ObservedSnapshot["periods"]["current"]) { - if (period.spend_observation !== "records_present" || period.matched_internal_prs === 0) return null; - return period.matched_users_recorded_spend / period.matched_internal_prs; -} - -export default function ObservedReport({ - snapshot, - accessToken, - readOnly, - onRefresh, - onConnect, - actions, - notice, - syncing, - onPeriod, -}: { - snapshot: ObservedSnapshot; - accessToken: string; - readOnly: boolean; - onRefresh: () => void; - onConnect: () => void; - actions: React.ReactNode; - notice?: React.ReactNode; - syncing: boolean; - onPeriod?: (days: number) => void; -}) { - const [comparison, setComparison] = useState("previous"); - const [activeTab, setActiveTab] = useState( - snapshot.people.length && snapshot.periods.current.merged_prs > 0 ? "people" : "pulls", - ); - const [accountEmail, setAccountEmail] = useState(null); - const [matchedOnly, setMatchedOnly] = useState(true); - const people = useMemo(() => reportPeople(snapshot, matchedOnly), [snapshot, matchedOnly]); - const pulls = useMemo(() => filterObservedPulls(snapshot, "current", matchedOnly), [snapshot, matchedOnly]); - const [personId, setPersonId] = useState(null); - const person = people.find((entry) => entry.id === personId) ?? null; - const terms = changeTerms(snapshot.source_provider); - const current = snapshot.periods.current; - const baseline = snapshot.periods[comparison]; - const days = Math.round((Date.parse(current.window.end) - Date.parse(current.window.start)) / 86400000) + 1; - const rangeOptions = [...new Set([7, 28, 90, days])].sort((a, b) => a - b); - const cost = costPerChange(current); - const baselineCost = costPerChange(baseline); - return ( - - - ROI Calculator -
- {actions} - {!readOnly && ( - - )} -
-
- {notice} -
- - }> - {number(snapshot.repos.length)} {snapshot.repos.length === 1 ? "repository" : "repositories"} - - - - Repositories -
    - {snapshot.repos.map((repo) => ( -
  • {repo}
  • - ))} -
-
-
-
- - - - {dateRange(current.window)} - - -
-
-
- - - - -
- -
- - - Engineers {people.length} - - {terms.requests} - Quality - Branch spend - - {activeTab !== "quality" && } - {!readOnly && ( - - )} -
- - setPersonId(selected.id)} - /> - - - {current.merged_prs > 0 || baseline.merged_prs > 0 ? ( -
- -
-
-

Behind the numbers

-

- {number(current.agent_authored)} of {number(current.merged_prs)} {terms.plural} were authored by - agents or bots. -

-

- Human-authored median merge time:{" "} - - {duration(current.human_summary.median_merge_hours)} - - , compared with {duration(baseline.human_summary.median_merge_hours)}. -

-
-
- - {number(current.agents_without_requester)} agent {terms.plural} have no requester - -
-
-
- ) : ( -
-

No merged changes yet

-

- Your repositories are connected. New activity will appear after the next sync -

-
- )} - -

- {matchedOnly - ? "Changes from people linked to internal accounts" - : `All repository ${terms.plural}, including agent work without a known requester`} -

- -
- - - - - - -
-
- - {number(current.matched_internal_prs)} {terms.plural} matched to {snapshot.people.length} engineers ยท{" "} - {number(current.agents_without_requester)} agent {terms.plural} without a requester - - Comparing with {dateRange(baseline.window)} ยท All dates UTC - Spend recorded by this gateway ยท Merge time is elapsed time, not effort -
- {accountEmail !== null && ( - setAccountEmail(null)} - onSaved={onRefresh} - /> - )} - {person && ( - setPersonId(null)} - onEdit={ - readOnly || !person.matched - ? undefined - : () => { - setAccountEmail(person.email); - setPersonId(null); - } - } - /> - )} -
- ); -} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorDialogs.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorDialogs.tsx deleted file mode 100644 index 39d7ef582ef..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorDialogs.tsx +++ /dev/null @@ -1,217 +0,0 @@ -"use client"; - -import React from "react"; - -import { extractErrorMessage } from "@/utils/errorUtils"; -import { Button, buttonVariants } from "@/components/ui/button"; -import { - Dialog, - DialogContent, - DialogDescription, - DialogFooter, - DialogHeader, - DialogTitle, -} from "@/components/ui/dialog"; -import { Input } from "@/components/ui/input"; -import { Label } from "@/components/ui/label"; -import { effortNote, estimateLabel, branchCostLabel } from "./roiCalculatorData"; -import type { ROIIdentityMapUpdate, ROIPull, ROISummary } from "./roiCalculatorData"; -import type { ROIPerson } from "./roiCalculatorData"; - -export type PersonMatchSelection = { person: ROIPerson; login: string }; - -export function PullReasoningDialog({ - pull, - summary, - onClose, -}: { - pull: ROIPull | null; - summary: ROISummary | null; - onClose: () => void; -}) { - const titleRef = React.useRef(null); - return ( - !open && onClose()}> - - {pull && ( - <> - - - {pull.title} - - - {pull.repo} #{pull.number} ยท {pull.login} - - -
-
-
-
Estimated effort
-
{estimateLabel(pull.estimate)}
-
-
-
Recorded AI cost
-
{branchCostLabel(pull)}
- {pull.branch_cost?.status === "matched" && ( -
{pull.branch_cost.requests} requests
- )} -
-
-

- {effortNote(pull.estimate.effort_basis ?? summary?.effort_basis)} -

-
-
-

Reasoning

-

- {pull.estimate.reasoning || "No estimate available."} -

-
-
-
Model
-
{pull.estimate.model || summary?.estimator_model}
-
Merged
-
{new Date(pull.merged_at).toLocaleDateString(undefined, { timeZone: "UTC" })}
-
Email match
-
{pull.email || "Not matched"}
-
-
-

Track costs for this branch

- {pull.branch_cost?.status === "matched" && ( -

- {pull.branch_cost.spend?.toFixed(8)} USD across {pull.branch_cost.requests} requests -

- )} - {pull.source_repo && pull.source_branch ? ( - <> -

Send both tags with each gateway request from this branch:

-
-                    {JSON.stringify(
-                      { metadata: { tags: [`repo:${pull.source_repo}`, `branch:${pull.source_branch}`] } },
-                      null,
-                      2,
-                    )}
-                  
-

- Retained requests in the reportโ€™s UTC period. Branch names are case-sensitive. Reused branches - cannot be split between changes. -

- - ) : ( -

- The source repository or branch is unavailable. Sync again to refresh its metadata. -

- )} -
- {summary?.estimator_prompt && ( -
- Estimator prompt -

{summary.estimator_prompt}

-
- )} - - {pull.url && ( - - View on {summary?.source_provider === "gitlab" ? "GitLab" : "GitHub"} - - )} - - - )} -
-
- ); -} - -export function IdentityMatchDialog({ - selection, - identityMap, - gatewayEmails, - onClose, - onSave, -}: { - selection: PersonMatchSelection | null; - identityMap: Record; - gatewayEmails: string[]; - onClose: () => void; - onSave: (payload: ROIIdentityMapUpdate) => Promise; -}) { - const [email, setEmail] = React.useState(() => - selection ? identityMap[selection.login.toLowerCase()] ?? selection.person.email ?? "" : "", - ); - const [error, setError] = React.useState(null); - const [busy, setBusy] = React.useState(false); - const person = selection?.person ?? null; - const login = selection?.login ?? ""; - const existingEmail = identityMap[login.toLowerCase()]; - - const save = async (value: string | null) => { - if (!login) return; - try { - setBusy(true); - await onSave({ github_login: login, email: value }); - setError(null); - onClose(); - } catch (reason) { - setError(extractErrorMessage(reason)); - } finally { - setBusy(false); - } - }; - - return ( - !open && onClose()}> - - - Match email - Link {login} to their gateway email. Manual matches take priority. - -
{ - event.preventDefault(); - void save(email.trim()); - }} - > -
- - setEmail(event.target.value)} - required - /> -
- - {Array.from(new Set(gatewayEmails)).map((address) => ( - - {error && ( -

- {error} -

- )} - - {existingEmail && ( - - )} - - -
-
-
- ); -} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.integration.test.tsx deleted file mode 100644 index 30cf5e09eee..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.integration.test.tsx +++ /dev/null @@ -1,814 +0,0 @@ -import userEvent from "@testing-library/user-event"; -import { act, fireEvent, render as renderWithoutNuqs, screen, waitFor } from "@testing-library/react"; -import { beforeEach, describe, expect, it, vi } from "vitest"; -import { NuqsAdapter } from "nuqs/adapters/react"; - -import { apiClient } from "@/components/networking"; -import ROICalculatorView from "./ROICalculatorView"; - -const render = (ui: Parameters[0]) => renderWithoutNuqs(ui, { wrapper: NuqsAdapter }); - -vi.mock("@/components/networking", () => ({ - apiClient: { - delete: vi.fn(), - get: vi.fn(), - post: vi.fn(), - put: vi.fn(), - }, -})); - -const summary = { - id: null, - mode: "live", - start: "2026-09-01", - end: "2026-09-30", - synced_at: "2026-09-30T12:00:00Z", - repos: ["org/repo"], - estimator_model: "estimator", - estimator_prompt: "Estimate hours.", - warnings: [], - effort_basis: "without_ai", - metrics: { - matched_spend: 12, - output_hours: 4, - total_spend: 20, - total_output_hours: 4, - excluded_spend: 8, - cost_per_hour: 3, - hours_per_dollar: 1 / 3, - merged_prs: 1, - estimated_prs: 1, - matched_prs: 1, - cohort_people: 1, - people_with_prs: 1, - pending_prs: 0, - }, - people: [ - { - id: "alice@example.com", - email: "alice@example.com", - logins: ["alice", "alice-work"], - spend: 12, - hours: 4, - prs: 1, - estimated_prs: 1, - pending_prs: 0, - match_methods: ["profile email"], - eligible: true, - cost_per_hour: 3, - }, - ], - pulls: [ - { - repo: "org/repo", - source_repo: "github.com/org/repo", - source_branch: "feature/routing", - branch_cost: { - repo: "github.com/org/repo", - branch: "feature/routing", - status: "matched", - spend: 8, - requests: 12, - }, - number: 42, - title: "Improve request routing", - url: "https://github.com/org/repo/pull/42", - login: "alice", - emails: ["alice@example.com"], - profile_email: "alice@example.com", - merged_at: "2026-09-12T00:00:00Z", - head_sha: "abc", - additions: 10, - deletions: 2, - changed_files: 1, - commit_count: 1, - incomplete_metadata: false, - estimate: { - status: "estimated", - hours: 4, - reasoning: "Updated routing and added a regression test.", - model: "estimator", - evidence_source: "pr_metadata", - effort_basis: "without_ai", - cached: false, - }, - cache_key: "cache", - email: "alice@example.com", - match_method: "profile email", - matched: true, - }, - ], - trend: [{ date: "2026-09-12", spend: 12, hours: 4, prs: 1 }], -} as const; - -const settings = { - github_api_url: "https://api.github.com", - repos: ["org/repo"], - estimator_model: "estimator", - estimator_prompt: "Estimate hours.", - backfill_days: 30, - identity_map: {}, - has_github_token: true, - default_prompt: "Estimate hours.", - available_models: ["estimator"], - ready: true, -}; - -const idleStatus = { - running: false, - phase: "idle", - stage: "Idle", - done: 0, - total: 0, - estimated: 0, - reused: 0, - needs_attention: 0, - error: null, -}; - -describe("ROICalculatorView", () => { - beforeEach(() => { - window.history.replaceState(null, "", "/roi-calculator/"); - vi.mocked(apiClient.get).mockReset(); - vi.mocked(apiClient.put).mockReset(); - vi.mocked(apiClient.post).mockReset(); - vi.mocked(apiClient.get).mockImplementation((path: string) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); - if (path === "/roi-calculator/report") return Promise.resolve({ report: summary }); - return Promise.resolve(idleStatus); - }); - vi.mocked(apiClient.put).mockResolvedValue({ report: summary, identity_map: { alice: "alice@example.com" } }); - }); - - it.each(["github", "gitlab"])("filters people and branches by linked accounts for %s", async (provider) => { - const outsidePerson = { - ...summary.people[0], - id: "outside", - email: "", - logins: ["outside"], - match_methods: ["no gateway match"], - spend: null, - }; - const outsidePull = { - ...summary.pulls[0], - number: 43, - title: "External contribution", - login: "outside", - email: "", - match_method: "no gateway match", - matched: false, - }; - const linkedPerson = { ...summary.people[0], spend: null, match_methods: ["manual"], eligible: false }; - const spendOnlyPerson = { - ...summary.people[0], - id: "internal@example.test", - email: "internal@example.test", - logins: [], - spend: 8.5, - prs: 0, - match_methods: [], - estimated_prs: 0, - hours: 0, - eligible: false, - cost_per_hour: null, - }; - const linkedPull = { ...summary.pulls[0], matched: false, match_method: "manual" }; - const report = { - ...summary, - source_provider: provider, - people: [linkedPerson, spendOnlyPerson, outsidePerson], - pulls: [linkedPull, outsidePull], - }; - vi.mocked(apiClient.get).mockImplementation((path: string) => { - if (path.endsWith("/settings")) return Promise.resolve(settings); - if (path.endsWith("/report")) return Promise.resolve({ report }); - return Promise.resolve(idleStatus); - }); - const user = userEvent.setup(); - render(); - await user.click(await screen.findByRole("tab", { name: "People" })); - expect(screen.getByRole("switch", { name: "Matched people only" })).toBeChecked(); - expect(screen.getByRole("button", { name: "alice" })).toBeInTheDocument(); - expect(screen.getByRole("row", { name: /internal@example.test/ })).toHaveTextContent("$8.50"); - expect(screen.queryByRole("button", { name: "outside" })).not.toBeInTheDocument(); - await user.click(screen.getByRole("switch", { name: "Matched people only" })); - expect(screen.getByRole("button", { name: "outside" })).toBeInTheDocument(); - await user.click(screen.getByRole("tab", { name: "Branches" })); - expect(screen.getByRole("switch", { name: "Matched people only" })).not.toBeChecked(); - expect(screen.getByText("External contribution")).toBeInTheDocument(); - await user.click(screen.getByRole("switch", { name: "Matched people only" })); - expect(screen.queryByText("External contribution")).not.toBeInTheDocument(); - expect(screen.getByText("Improve request routing")).toBeInTheDocument(); - expect(apiClient.put).not.toHaveBeenCalled(); - }); - - it("shows the spend summary and opens an accessible pull reasoning dialog", async () => { - render(); - - expect(await screen.findByText("Gateway AI cost")).toBeInTheDocument(); - expect(screen.getByText("$20.00")).toBeInTheDocument(); - fireEvent.click(screen.getByRole("button", { name: "Open estimate for org/repo pull request 42" })); - - expect(await screen.findByRole("dialog")).toBeInTheDocument(); - expect(screen.getByText("Updated routing and added a regression test.")).toBeInTheDocument(); - expect(screen.getByRole("link", { name: "View on GitHub" })).toHaveAttribute( - "href", - "https://github.com/org/repo/pull/42", - ); - }); - - it("separates the overview, people, and branch reports into three tabs", async () => { - render(); - expect(await screen.findByRole("heading", { name: "Where AI costs are matched" })).toBeVisible(); - fireEvent.click(screen.getByRole("tab", { name: "Branches" })); - expect(screen.getByRole("heading", { name: "Costs by branch" })).toBeVisible(); - expect(screen.getByRole("cell", { name: "$8.00" })).toBeVisible(); - fireEvent.click(screen.getByRole("tab", { name: "People" })); - expect(screen.getByRole("heading", { name: "People and account matches" })).toBeVisible(); - fireEvent.click(screen.getByRole("tab", { name: "Overview" })); - expect(screen.getByRole("heading", { name: "Highest-cost changes" })).toBeVisible(); - expect(screen.queryByRole("radio")).not.toBeInTheDocument(); - }); - - it("shows incomplete repository results without a spend-per-hour figure", async () => { - const warning = "Incomplete report: could not read org/unavailable. Spend-per-hour figures are unavailable."; - vi.mocked(apiClient.get).mockImplementation((path: string) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); - if (path === "/roi-calculator/report") { - return Promise.resolve({ - report: { - ...summary, - warnings: [warning], - metrics: { ...summary.metrics, cost_per_hour: null, hours_per_dollar: null }, - people: summary.people.map((person) => ({ ...person, cost_per_hour: null })), - }, - }); - } - return Promise.resolve(idleStatus); - }); - - render(); - - expect(await screen.findByRole("alert")).toHaveTextContent(warning); - expect(screen.getByRole("button", { name: "Open estimate for org/repo pull request 42" })).toBeInTheDocument(); - expect(screen.queryByText("$3.00")).not.toBeInTheDocument(); - fireEvent.click(screen.getByRole("tab", { name: "People" })); - fireEvent.click(screen.getByText("How this is calculated")); - expect( - screen.getByText("Spend per estimated hour is unavailable until all selected repositories can be read."), - ).toBeVisible(); - }); - - it("lets a view-only admin read the report without write controls", async () => { - const runningStatus = { - ...idleStatus, - running: true, - phase: "estimating", - stage: "Estimating pull requests", - total: 1, - }; - vi.mocked(apiClient.get).mockImplementation((path: string) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); - if (path === "/roi-calculator/report") return Promise.resolve({ report: summary }); - return Promise.resolve(runningStatus); - }); - - render(); - - expect(await screen.findByText("Gateway AI cost")).toBeInTheDocument(); - expect(screen.getByRole("note")).toHaveTextContent("Read-only access"); - expect(screen.queryByRole("button", { name: "Run analysis" })).not.toBeInTheDocument(); - expect(screen.queryByRole("button", { name: "Cancel sync" })).not.toBeInTheDocument(); - - fireEvent.click(screen.getByRole("tab", { name: "People" })); - expect(screen.getByText("alice-work")).toBeInTheDocument(); - expect(screen.queryByRole("button", { name: "alice-work" })).not.toBeInTheDocument(); - - fireEvent.click(screen.getByRole("button", { name: "Settings" })); - expect(screen.getByLabelText("GitHub token (optional for public repositories)")).toBeDisabled(); - expect(screen.queryByRole("button", { name: "Save settings" })).not.toBeInTheDocument(); - expect(screen.queryByRole("button", { name: "Run analysis" })).not.toBeInTheDocument(); - }); - - it("lets an admin open the people view and save a manual email match", async () => { - render(); - - fireEvent.click(await screen.findByRole("tab", { name: "People" })); - fireEvent.click(await screen.findByRole("button", { name: "alice-work" })); - fireEvent.change(screen.getByLabelText("Gateway email"), { - target: { value: "alice+work@example.com" }, - }); - fireEvent.click(screen.getByRole("button", { name: "Save match" })); - - await waitFor(() => - expect(apiClient.put).toHaveBeenCalledWith("/roi-calculator/identity-map", { - accessToken: "token", - body: { github_login: "alice-work", email: "alice+work@example.com" }, - }), - ); - }); - - it("presents onboarding settings once when no report exists", async () => { - const emptySettings = { ...settings, has_github_token: false, ready: false, repos: [], estimator_model: "" }; - vi.mocked(apiClient.get).mockImplementation((path: string) => { - if (path === "/roi-calculator/settings") return Promise.resolve(emptySettings); - if (path === "/roi-calculator/report") return Promise.resolve({ report: null }); - return Promise.resolve(idleStatus); - }); - - render(); - - expect(await screen.findByRole("heading", { name: "Connect your repositories" })).toBeInTheDocument(); - expect(screen.getByLabelText("GitHub token (optional for public repositories)")).toHaveAttribute( - "type", - "password", - ); - expect(screen.getAllByText("Connect your repositories")).toHaveLength(1); - }); - - it.each(["github", "gitlab"])("only permits unauthenticated repository browsing for GitLab: %s", async (provider) => { - const publicSettings = { - ...settings, - source_provider: provider, - gitlab_api_url: "https://gitlab.com/api/v4", - has_github_token: false, - has_gitlab_token: false, - }; - vi.mocked(apiClient.get).mockImplementation((path: string) => { - if (path === "/roi-calculator/settings") { - return Promise.resolve(publicSettings); - } - if (path === "/roi-calculator/report") return Promise.resolve({ report: summary }); - return Promise.resolve(idleStatus); - }); - render(); - fireEvent.click(await screen.findByRole("button", { name: "Settings" })); - const load = screen.getByRole("button", { name: "Load repositories" }); - if (provider === "github") expect(load).toBeDisabled(); - else expect(load).toBeEnabled(); - expect(screen.getByRole("textbox", { name: "Repository name" })).toBeEnabled(); - }); - - it("searches by the real model name and saves the selected gateway alias", async () => { - const user = userEvent.setup(); - const modelSettings = { - ...settings, - available_models: ["estimator", "fast-estimator"], - estimator_models: [ - { model_name: "estimator", provider_models: ["custom-model"] }, - { model_name: "fast-estimator", provider_models: ["openai/gpt-6-luna"] }, - ], - }; - vi.mocked(apiClient.get).mockImplementation((path: string) => { - if (path === "/roi-calculator/settings") return Promise.resolve(modelSettings); - if (path === "/roi-calculator/report") return Promise.resolve({ report: summary }); - return Promise.resolve(idleStatus); - }); - vi.mocked(apiClient.put).mockResolvedValue({ ...modelSettings, estimator_model: "fast-estimator" }); - render(); - await user.click(await screen.findByRole("button", { name: "Settings" })); - const search = screen.getByRole("combobox", { name: "Estimator model" }); - await user.clear(search); - await user.type(search, "Luna"); - expect(screen.queryByRole("option", { name: /custom-model/ })).not.toBeInTheDocument(); - await user.click(await screen.findByRole("option", { name: /GPT-6 Luna.*Recommended/ })); - expect(search).toHaveValue("GPT-6 Luna"); - await user.click(screen.getByRole("button", { name: "Save settings" })); - expect(apiClient.put).toHaveBeenCalledWith( - "/roi-calculator/settings", - expect.objectContaining({ - body: expect.objectContaining({ estimator_model: "fast-estimator" }), - }), - ); - }); - - it("closes the old settings dialog when saving a different source", async () => { - const gitlabSettings = { - ...settings, - source_provider: "gitlab", - gitlab_api_url: "https://gitlab.com/api/v4", - has_gitlab_token: false, - repos: [], - ready: false, - }; - vi.mocked(apiClient.put).mockResolvedValue(gitlabSettings); - render(); - fireEvent.click(await screen.findByRole("button", { name: "Settings" })); - fireEvent.change(screen.getByLabelText("Repository source"), { target: { value: "gitlab" } }); - fireEvent.click(screen.getByRole("button", { name: "Save settings" })); - expect(await screen.findByRole("heading", { name: "Connect your repositories" })).toBeVisible(); - expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); - expect(screen.getAllByLabelText("Repository source")).toHaveLength(1); - expect(screen.getByLabelText("Repository source")).toHaveValue("gitlab"); - }); - - it("keeps a running analysis visible when a source change finishes saving", async () => { - const gitlabSettings = { ...settings, source_provider: "gitlab", repos: ["group/project"] }; - const runningStatus = { ...idleStatus, running: true, phase: "estimating", total: 1 }; - const saveRequest = Promise.withResolvers(); - vi.mocked(apiClient.put).mockReturnValue(saveRequest.promise); - render(); - fireEvent.click(await screen.findByRole("button", { name: "Settings" })); - fireEvent.change(screen.getByLabelText("Repository source"), { target: { value: "gitlab" } }); - fireEvent.click(screen.getByRole("button", { name: "Save settings" })); - vi.mocked(apiClient.get).mockResolvedValue(runningStatus); - expect(await screen.findByRole("progressbar", { hidden: true }, { timeout: 3000 })).toBeInTheDocument(); - - saveRequest.resolve(gitlabSettings); - await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument()); - expect(screen.getByRole("progressbar", { name: "Sync progress" })).toBeVisible(); - expect(screen.getByRole("button", { name: "Cancel sync" })).toBeEnabled(); - expect(screen.queryByRole("button", { name: "Run analysis" })).not.toBeInTheDocument(); - expect(screen.queryByRole("heading", { name: "Connect your repositories" })).not.toBeInTheDocument(); - expect(apiClient.post).not.toHaveBeenCalled(); - }); - - it.each(["success", "failure"])("ignores an old report refresh %s after switching sources", async (outcome) => { - const gitlabSettings = { ...settings, source_provider: "gitlab", repos: [], ready: false }; - const oldRequest = Promise.withResolvers<{ report: typeof summary }>(); - const complete = { ...idleStatus, phase: "complete", finished_at: "2026-09-30T12:00:00Z" }; - vi.mocked(apiClient.put).mockResolvedValue(gitlabSettings); - render(); - fireEvent.click(await screen.findByRole("button", { name: "Settings" })); - vi.mocked(apiClient.get) - .mockClear() - .mockImplementation((path: string) => - path === "/roi-calculator/report" ? oldRequest.promise : Promise.resolve(complete), - ); - await waitFor( - () => expect(apiClient.get).toHaveBeenCalledWith("/roi-calculator/report", { accessToken: "token" }), - { - timeout: 3000, - }, - ); - fireEvent.change(screen.getByLabelText("Repository source"), { target: { value: "gitlab" } }); - fireEvent.click(screen.getByRole("button", { name: "Save settings" })); - expect(await screen.findByRole("heading", { name: "Connect your repositories" })).toBeVisible(); - - await act(async () => { - if (outcome === "success") oldRequest.resolve({ report: summary }); - else oldRequest.reject(new Error("The old source is unavailable")); - }); - expect(screen.getByRole("heading", { name: "Connect your repositories" })).toBeVisible(); - expect(screen.getByLabelText("Repository source")).toHaveValue("gitlab"); - expect(screen.queryByText("Improve request routing")).not.toBeInTheDocument(); - expect(screen.queryByText("The old source is unavailable")).not.toBeInTheDocument(); - - const nextReport = { ...summary, pulls: [{ ...summary.pulls[0], title: "New source merge request" }] }; - vi.mocked(apiClient.get).mockImplementation((path: string) => - Promise.resolve( - path === "/roi-calculator/report" - ? { report: nextReport } - : { ...complete, finished_at: "2026-09-30T13:00:00Z" }, - ), - ); - expect(await screen.findByText("New source merge request", {}, { timeout: 3000 })).toBeVisible(); - expect(screen.queryByText("Improve request routing")).not.toBeInTheDocument(); - }); - - it("clearly identifies the sample report and returns to setup when exiting", async () => { - const emptySettings = { ...settings, has_github_token: false, ready: false, repos: [], estimator_model: "" }; - vi.mocked(apiClient.get).mockImplementation((path: string, options) => { - if (path === "/roi-calculator/settings") return Promise.resolve(emptySettings); - if (path === "/roi-calculator/report") { - return Promise.resolve({ report: options?.query?.mode === "demo" ? { ...summary, mode: "demo" } : null }); - } - return Promise.resolve(idleStatus); - }); - - render(); - fireEvent.click(await screen.findByRole("button", { name: "Preview sample report" })); - - expect(await screen.findByText("Youโ€™re viewing demo data")).toBeVisible(); - expect(screen.getByRole("tab", { name: "Branches" })).toHaveAttribute("aria-selected", "true"); - expect(screen.getByText("Cost / estimated hour")).toBeVisible(); - expect(screen.queryByRole("button", { name: "Run analysis" })).not.toBeInTheDocument(); - expect(screen.queryByRole("tab", { name: "Settings" })).not.toBeInTheDocument(); - - fireEvent.click(screen.getByRole("button", { name: "Exit demo" })); - - expect(screen.getByRole("heading", { name: "Connect your repositories" })).toBeVisible(); - expect(screen.queryByText("Youโ€™re viewing demo data")).not.toBeInTheDocument(); - expect(apiClient.post).not.toHaveBeenCalled(); - expect(apiClient.put).not.toHaveBeenCalled(); - }); - - it("opens sample PR costs from a live report and restores the live data on exit", async () => { - const samplePull = { - ...summary.pulls[0], - title: "Sample usage breakdown", - source_repo: "github.com/org/repo", - source_branch: "feature/usage", - branch_cost: { - status: "matched", - spend: 9.1, - requests: 75, - repo: "github.com/org/repo", - branch: "feature/usage", - }, - }; - vi.mocked(apiClient.get).mockImplementation((path: string, options) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); - if (path === "/roi-calculator/report") { - return Promise.resolve({ - report: options?.query?.mode === "demo" ? { ...summary, mode: "demo", pulls: [samplePull] } : summary, - }); - } - return Promise.resolve(idleStatus); - }); - - render(); - fireEvent.click(await screen.findByRole("tab", { name: "Branches" })); - fireEvent.change(screen.getByRole("searchbox"), { target: { value: "no matching PR" } }); - fireEvent.click(screen.getByRole("button", { name: "Preview sample report" })); - - expect(await screen.findByText("Youโ€™re viewing demo data")).toBeVisible(); - await waitFor(() => expect(window.location.search).toBe("?demo=1")); - expect(screen.getByRole("tab", { name: "Branches" })).toHaveAttribute("aria-selected", "true"); - expect(screen.getByRole("searchbox")).toHaveValue(""); - expect(screen.getByRole("cell", { name: "$9.10" })).toBeVisible(); - expect(screen.queryByText("Improve request routing")).not.toBeInTheDocument(); - expect(screen.queryByRole("button", { name: "Run analysis" })).not.toBeInTheDocument(); - - const runningStatus = { ...idleStatus, running: true, phase: "estimating", total: 1 }; - vi.mocked(apiClient.get).mockClear().mockResolvedValue(runningStatus); - await waitFor(() => expect(apiClient.get).toHaveBeenCalledWith("/roi-calculator/sync", { accessToken: "token" }), { - timeout: 3000, - }); - expect(screen.queryByRole("progressbar")).not.toBeInTheDocument(); - - fireEvent.click(screen.getByRole("button", { name: "Open estimate for org/repo pull request 42" })); - expect(await screen.findByRole("dialog")).toBeVisible(); - expect(screen.getByText("75 requests")).toBeVisible(); - expect(screen.getByText(/branch:feature\/usage/)).toBeVisible(); - fireEvent.click(screen.getByRole("button", { name: "Close" })); - fireEvent.click(screen.getByRole("button", { name: "Exit demo" })); - - expect(screen.getByText("Improve request routing")).toBeVisible(); - expect(screen.queryByText("Sample usage breakdown")).not.toBeInTheDocument(); - await waitFor(() => expect(window.location.search).toBe("")); - expect(screen.getByRole("button", { name: "Syncingโ€ฆ" })).toBeDisabled(); - expect(apiClient.post).not.toHaveBeenCalled(); - expect(apiClient.put).not.toHaveBeenCalled(); - }); - - it("opens a demo link with sample data even while live analysis is running", async () => { - window.history.replaceState(null, "", "/roi-calculator/?demo=1"); - const demoSummary = { ...summary, mode: "demo", metrics: { ...summary.metrics, total_spend: 38.4 } }; - const runningStatus = { ...idleStatus, running: true, phase: "estimating", total: 1 }; - vi.mocked(apiClient.get).mockImplementation((path: string, options) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); - if (path === "/roi-calculator/report") { - return Promise.resolve({ report: options?.query?.mode === "demo" ? demoSummary : summary }); - } - return Promise.resolve(runningStatus); - }); - render(); - expect(await screen.findByText("Youโ€™re viewing demo data")).toBeVisible(); - expect(screen.getByText("$38.40")).toBeVisible(); - expect(screen.queryByRole("progressbar")).not.toBeInTheDocument(); - expect(screen.queryByRole("button", { name: "Run analysis" })).not.toBeInTheDocument(); - expect(apiClient.post).not.toHaveBeenCalled(); - fireEvent.click(screen.getByRole("button", { name: "Exit demo" })); - expect(screen.getByRole("progressbar")).toBeVisible(); - expect(screen.getByText("$20.00")).toBeVisible(); - await waitFor(() => expect(window.location.search).toBe("")); - }); - - it.each(["report", "sync"])("loads a demo link when the live %s request fails", async (failedRequest) => { - window.history.replaceState(null, "", "/roi-calculator/?demo=1"); - vi.mocked(apiClient.get).mockImplementation((path: string, options) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); - if (options?.query?.mode === "demo") return Promise.resolve({ report: { ...summary, mode: "demo" } }); - if (path === `/roi-calculator/${failedRequest}`) return Promise.reject(new Error("Live data unavailable")); - if (path === "/roi-calculator/report") return Promise.resolve({ report: summary }); - return Promise.resolve(idleStatus); - }); - render(); - expect(await screen.findByText("Youโ€™re viewing demo data")).toBeVisible(); - expect(screen.getByText("Gateway AI cost")).toBeVisible(); - expect(screen.queryByRole("alert")).not.toBeInTheDocument(); - expect(apiClient.post).not.toHaveBeenCalled(); - }); - - it("shows the demo without waiting for a stalled live request", async () => { - window.history.replaceState(null, "", "/roi-calculator/?demo=1"); - vi.mocked(apiClient.get).mockImplementation((path: string, options) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); - if (options?.query?.mode === "demo") return Promise.resolve({ report: { ...summary, mode: "demo" } }); - return new Promise(() => {}); - }); - render(); - expect(await screen.findByText("Youโ€™re viewing demo data")).toBeVisible(); - expect(screen.getByText("Gateway AI cost")).toBeVisible(); - }); - - it("keeps the live calculator usable when a demo link cannot load sample data", async () => { - window.history.replaceState(null, "", "/roi-calculator/?demo=1&from=review#overview"); - vi.mocked(apiClient.get).mockImplementation((path: string, options) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); - if (options?.query?.mode === "demo") return Promise.reject(new Error("Sample data unavailable")); - if (path === "/roi-calculator/report") return Promise.resolve({ report: summary }); - return Promise.resolve(idleStatus); - }); - render(); - expect(await screen.findByText("Gateway AI cost")).toBeVisible(); - expect(screen.getByText("$20.00")).toBeVisible(); - expect(screen.getByRole("alert")).toHaveTextContent("Sample data unavailable"); - expect(screen.getByRole("button", { name: "Settings" })).toBeEnabled(); - expect(screen.queryByText("Youโ€™re viewing demo data")).not.toBeInTheDocument(); - expect(apiClient.post).not.toHaveBeenCalled(); - await waitFor(() => expect(window.location.search).toBe("?from=review")); - expect(window.location.hash).toBe("#overview"); - }); - - it.each(["report", "sync"])("waits for the live %s when exiting a demo", async (pendingRequest) => { - window.history.replaceState(null, "", "/roi-calculator/?demo=1"); - const pending = Promise.withResolvers<{ report: typeof summary } | typeof idleStatus>(); - vi.mocked(apiClient.get).mockImplementation((path: string, options) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); - if (options?.query?.mode === "demo") return Promise.resolve({ report: { ...summary, mode: "demo" } }); - if (path === `/roi-calculator/${pendingRequest}`) return pending.promise; - if (path === "/roi-calculator/report") return Promise.resolve({ report: summary }); - return Promise.resolve(idleStatus); - }); - render(); - fireEvent.click(await screen.findByRole("button", { name: "Exit demo" })); - expect(screen.getByText("Loading ROI Calculatorโ€ฆ")).toBeVisible(); - expect(screen.queryByRole("heading", { name: "Connect your repositories" })).not.toBeInTheDocument(); - expect(screen.queryByRole("button", { name: "Run analysis" })).not.toBeInTheDocument(); - await waitFor(() => expect(window.location.search).toBe("")); - - pending.resolve(pendingRequest === "report" ? { report: summary } : idleStatus); - expect(await screen.findByText("Gateway AI cost")).toBeVisible(); - expect(screen.getByText("$20.00")).toBeVisible(); - expect(screen.queryByText("Loading ROI Calculatorโ€ฆ")).not.toBeInTheDocument(); - expect(apiClient.post).not.toHaveBeenCalled(); - }); - - it.each(["report", "demo"])("retains a failed %s load after a successful sync poll", async (failedRequest) => { - if (failedRequest === "demo") window.history.replaceState(null, "", "/roi-calculator/?demo=1"); - const message = `The ${failedRequest} is unavailable`; - vi.mocked(apiClient.get).mockImplementation((path: string, options) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); - if (options?.query?.mode === "demo") return Promise.reject(new Error(message)); - if (path === "/roi-calculator/report") { - return failedRequest === "report" ? Promise.reject(new Error(message)) : Promise.resolve({ report: summary }); - } - return Promise.resolve(idleStatus); - }); - render(); - expect(await screen.findByRole("alert")).toHaveTextContent(message); - const running = { ...idleStatus, running: true, phase: "estimating", total: 1 }; - vi.mocked(apiClient.get).mockImplementation((path: string) => { - if (path === "/roi-calculator/report") return Promise.reject(new Error(message)); - return Promise.resolve(running); - }); - expect(await screen.findByRole("progressbar", { name: "Sync progress" }, { timeout: 3000 })).toBeVisible(); - expect(screen.getByRole("alert")).toHaveTextContent(message); - }); - - it("retries a failed initial report and clears its error only when the report recovers", async () => { - vi.mocked(apiClient.get).mockImplementation((path: string) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); - if (path === "/roi-calculator/report") return Promise.reject(new Error("Report unavailable")); - return Promise.resolve(idleStatus); - }); - render(); - expect(await screen.findByRole("alert")).toHaveTextContent("Report unavailable"); - expect(screen.getByRole("button", { name: "Settings" })).toBeEnabled(); - expect(screen.getByRole("button", { name: "Run analysis" })).toBeEnabled(); - vi.mocked(apiClient.get).mockImplementation((path: string) => - Promise.resolve(path === "/roi-calculator/report" ? { report: summary } : idleStatus), - ); - expect(await screen.findByText("Gateway AI cost", {}, { timeout: 3000 })).toBeVisible(); - expect(screen.queryByRole("alert")).not.toBeInTheDocument(); - expect(screen.queryByRole("heading", { name: "Connect your repositories" })).not.toBeInTheDocument(); - }); - - it("returns to Overview and shows the last sync time when completion is polled from Settings", async () => { - const runningStatus = { - ...idleStatus, - running: true, - phase: "estimating", - stage: "Estimating pull requests", - total: 1, - }; - const completedStatus = { ...idleStatus, phase: "complete", done: 57, total: 57, reused: 57 }; - vi.mocked(apiClient.get) - .mockResolvedValueOnce(settings) - .mockResolvedValueOnce({ report: null }) - .mockResolvedValueOnce(runningStatus) - .mockResolvedValueOnce(completedStatus) - .mockImplementationOnce( - () => - new Promise((resolve) => { - window.setTimeout(() => resolve({ report: summary }), 25); - }), - ); - - render(); - - expect(await screen.findByRole("progressbar", { name: "Sync progress" })).toBeInTheDocument(); - expect(await screen.findByText("Gateway AI cost", {}, { timeout: 5000 })).toBeInTheDocument(); - expect(screen.queryByRole("heading", { name: "Connect your repositories" })).not.toBeInTheDocument(); - expect(screen.getByRole("status")).toHaveTextContent("Last synced Sep 30, 2026, 12:00 PM UTC"); - }); - - it("shows the sync error returned by the status endpoint", async () => { - const runningStatus = { - ...idleStatus, - running: true, - phase: "estimating", - stage: "Estimating pull requests", - total: 1, - }; - const errorStatus = { - ...idleStatus, - phase: "error", - error: "The estimator could not score a pull request.", - }; - vi.mocked(apiClient.get) - .mockResolvedValueOnce(settings) - .mockResolvedValueOnce({ report: null }) - .mockResolvedValueOnce(runningStatus) - .mockResolvedValueOnce(errorStatus); - - render(); - - expect(await screen.findByRole("alert", {}, { timeout: 5000 })).toHaveTextContent( - "The estimator could not score a pull request.", - ); - expect(screen.getByText("Sync failed")).toBeInTheDocument(); - }); - - it("shows a report error and ends progress when the completed report cannot load", async () => { - const runningStatus = { - ...idleStatus, - running: true, - phase: "estimating", - stage: "Estimating pull requests", - total: 1, - }; - const completedStatus = { ...idleStatus, phase: "complete", done: 1, total: 1 }; - vi.mocked(apiClient.get) - .mockResolvedValueOnce(settings) - .mockResolvedValueOnce({ report: null }) - .mockResolvedValueOnce(runningStatus) - .mockResolvedValueOnce(completedStatus) - .mockRejectedValueOnce(new Error("The report could not be loaded.")); - - render(); - - expect(await screen.findByRole("progressbar", { name: "Sync progress" })).toBeInTheDocument(); - expect(await screen.findByRole("alert", {}, { timeout: 5000 })).toHaveTextContent( - "The report could not be loaded.", - ); - expect(screen.queryByRole("progressbar", { name: "Sync progress" })).not.toBeInTheDocument(); - }); - - it("clears a transient poll error when the next poll completes and loads the report", async () => { - const runningStatus = { - ...idleStatus, - running: true, - phase: "estimating", - stage: "Estimating pull requests", - total: 1, - }; - const completedStatus = { ...idleStatus, phase: "complete", done: 1, total: 1 }; - vi.mocked(apiClient.get) - .mockResolvedValueOnce(settings) - .mockResolvedValueOnce({ report: null }) - .mockResolvedValueOnce(runningStatus) - .mockRejectedValueOnce(new Error("The sync status could not be loaded.")) - .mockResolvedValueOnce(completedStatus) - .mockResolvedValueOnce({ report: summary }); - - render(); - - expect(await screen.findByRole("progressbar", { name: "Sync progress" })).toBeInTheDocument(); - expect(await screen.findByRole("alert", {}, { timeout: 5000 })).toHaveTextContent( - "The sync status could not be loaded.", - ); - expect(await screen.findByText("Gateway AI cost", {}, { timeout: 7000 })).toBeInTheDocument(); - expect(screen.queryByText("The sync status could not be loaded.")).not.toBeInTheDocument(); - }); - it("saves the edited schedule before running from Settings", async () => { - vi.mocked(apiClient.put).mockResolvedValue(settings); - vi.mocked(apiClient.post).mockResolvedValue({ ...idleStatus, running: true }); - render(); - fireEvent.click(await screen.findByRole("button", { name: "Settings" })); - fireEvent.change(screen.getByLabelText("Update interval (hours)"), { target: { value: "6" } }); - fireEvent.click(screen.getByRole("button", { name: "Save and run analysis" })); - await waitFor(() => expect(apiClient.post).toHaveBeenCalledWith("/roi-calculator/sync", { accessToken: "token" })); - expect(apiClient.put).toHaveBeenCalledWith( - "/roi-calculator/settings", - expect.objectContaining({ - body: expect.objectContaining({ update_interval_minutes: 360, estimator_model: "estimator" }), - }), - ); - expect(vi.mocked(apiClient.put).mock.invocationCallOrder[0]).toBeLessThan( - vi.mocked(apiClient.post).mock.invocationCallOrder[0], - ); - }); -}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.tsx deleted file mode 100644 index 36cb3d8d736..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.tsx +++ /dev/null @@ -1,508 +0,0 @@ -"use client"; - -import { Page, PageTabsList, PageTabsTrigger } from "@/components/shared/Page"; -import React from "react"; -import { Calculator, RefreshCw, Settings2 } from "lucide-react"; - -import { apiClient } from "@/components/networking"; -import { DemoNotice } from "@/components/shared/DemoNotice"; -import { PageHeader, PageHeaderDescription, PageHeaderTitle } from "@/components/shared/PageHeader"; -import { Alert, AlertDescription, AlertTitle } from "@/components/ui/alert"; -import { Button } from "@/components/ui/button"; -import { Card, CardContent } from "@/components/ui/card"; -import { Skeleton } from "@/components/ui/skeleton"; -import { Tabs } from "@/components/ui/tabs"; -import { Dialog, DialogContent, DialogHeader, DialogTitle, DialogDescription } from "@/components/ui/dialog"; -import { extractErrorMessage } from "@/utils/errorUtils"; -import { isProxyAdminTierRole } from "@/utils/roles"; -import { useQueryState } from "nuqs"; -import ROISettingsPanel from "./ROISettingsPanel"; -import { MatchedPeopleToggle } from "./MatchedPeopleToggle"; -import { IdentityMatchDialog, type PersonMatchSelection, PullReasoningDialog } from "./ROICalculatorDialogs"; -import { ROIBranches, ROIOverview, ROIPeopleView } from "./ROICalculatorViews"; -import { filterPulls, formatSyncedAt } from "./roiCalculatorData"; -import type { - ROIIdentityMapResponse, - ROIIdentityMapUpdate, - ROIPull, - ROIReportResponse, - ROISettings, - ROISummary, - ROISyncStatus, -} from "./roiCalculatorData"; -import { parseAsDemoFlag } from "./demoUrlState"; - -type View = "overview" | "people" | "branches"; - -const IDLE_STATUS: ROISyncStatus = { - running: false, - elapsed_seconds: 0, - phase: "idle", - stage: "Idle", - done: 0, - total: 0, - estimated: 0, - reused: 0, - needs_attention: 0, - error: null, -}; - -export default function ROICalculatorView({ - accessToken, - userRole = null, - isViewOnly = false, -}: { - accessToken: string | null; - userRole?: string | null; - isViewOnly?: boolean; -}) { - const [demo, setDemo] = useQueryState("demo", parseAsDemoFlag); - const [demoRequestedOnLoad] = React.useState(demo === true); - const [sampleSummary, setSampleSummary] = React.useState(null); - const adminReadOnly = isViewOnly && isProxyAdminTierRole(userRole ?? ""); - const readOnly = adminReadOnly || sampleSummary !== null; - const [settingsOpen, setSettingsOpen] = React.useState(false); - const [view, setView] = React.useState("overview"); - const [settings, setSettings] = React.useState(null); - const [liveSummary, setSummary] = React.useState(null); - const summary = sampleSummary ?? liveSummary; - const [status, setStatus] = React.useState(IDLE_STATUS); - const [selectedPull, setSelectedPull] = React.useState(null); - const [matchingPerson, setMatchingPerson] = React.useState(null); - const [error, setError] = React.useState(null); - const [demoError, setDemoError] = React.useState(null); - const [reportError, setReportError] = React.useState(null); - const [syncError, setSyncError] = React.useState(null); - const [loadingInitialData, setLoadingInitialData] = React.useState(true); - const [loadingLiveData, setLoadingLiveData] = React.useState(true); - const statusRef = React.useRef(IDLE_STATUS); - const reportNeedsRefresh = React.useRef(false); - const sourceRevision = React.useRef(0); - const settingsLoaded = settings !== null && !loadingInitialData && !loadingLiveData; - const requestError = [error, demoError, reportError, syncError].filter(Boolean).join(" "); - const [query, setQuery] = React.useState(""); - const [matchedOnly, setMatchedOnly] = React.useState(true); - - const loadReport = React.useCallback(async () => { - if (!accessToken) return null; - const response: ROIReportResponse = await apiClient.get("/roi-calculator/report", { accessToken }); - return response.report; - }, [accessToken]); - - React.useEffect(() => { - if (!accessToken) return; - let cancelled = false; - const settingsRequest = apiClient.get("/roi-calculator/settings", { accessToken }); - const reportRequest = apiClient - .get("/roi-calculator/report", { accessToken }) - .then((response) => { - if (cancelled) return; - setSummary(response.report); - setReportError(null); - reportNeedsRefresh.current = false; - }) - .catch((reason: unknown) => { - if (cancelled) return; - setReportError(extractErrorMessage(reason)); - reportNeedsRefresh.current = true; - }); - const statusRequest = apiClient - .get("/roi-calculator/sync", { accessToken }) - .then((syncStatus) => { - if (cancelled) return; - setStatus(syncStatus); - statusRef.current = syncStatus; - setSyncError(null); - }) - .catch((reason: unknown) => { - if (!cancelled) setSyncError(extractErrorMessage(reason)); - }); - const liveData = Promise.all([reportRequest, statusRequest]) - .then(() => null) - .finally(() => { - if (!cancelled) setLoadingLiveData(false); - }); - Promise.all([ - settingsRequest, - demoRequestedOnLoad - ? apiClient - .get("/roi-calculator/report", { accessToken, query: { mode: "demo" } }) - .catch((reason: unknown) => { - if (!cancelled) { - setDemoError(`Could not load demo data: ${extractErrorMessage(reason)}`); - setDemo(null); - } - return liveData; - }) - : liveData, - ]) - .then(([nextSettings, sampleResponse]) => { - if (cancelled) return; - setSettings(nextSettings); - setSampleSummary(sampleResponse?.report ?? null); - setError(null); - if (sampleResponse) setDemoError(null); - }) - .catch((reason: unknown) => { - if (!cancelled) setError(extractErrorMessage(reason)); - }) - .finally(() => { - if (!cancelled) setLoadingInitialData(false); - }); - return () => { - cancelled = true; - }; - }, [accessToken, demoRequestedOnLoad, setDemo]); - - React.useEffect(() => { - if (!accessToken || !settingsLoaded) return; - let cancelled = false; - let requestInFlight = false; - const interval = window.setInterval(() => { - if (requestInFlight) return; - requestInFlight = true; - const revision = sourceRevision.current; - const isCurrent = () => !cancelled && revision === sourceRevision.current; - apiClient - .get("/roi-calculator/sync", { accessToken }) - .then(async (nextStatus) => { - if (!isCurrent()) return; - const previousStatus = statusRef.current; - statusRef.current = nextStatus; - setStatus(nextStatus); - setSyncError(null); - const finished = !nextStatus.running && nextStatus.phase === "complete"; - const reportChanged = previousStatus.running || nextStatus.finished_at !== previousStatus.finished_at; - if (reportNeedsRefresh.current || (finished && reportChanged)) { - reportNeedsRefresh.current = true; - try { - const report = await loadReport(); - if (!isCurrent()) return; - setSummary(report); - setReportError(null); - reportNeedsRefresh.current = false; - } catch (reason: unknown) { - if (isCurrent()) setReportError(extractErrorMessage(reason)); - } - } - }) - .catch((reason: unknown) => { - if (isCurrent()) setSyncError(extractErrorMessage(reason)); - }) - .finally(() => { - requestInFlight = false; - }); - }, 1500); - return () => { - cancelled = true; - window.clearInterval(interval); - }; - }, [accessToken, loadReport, settingsLoaded]); - - const startSync = React.useCallback(async () => { - if (!accessToken || readOnly) return; - try { - setError(null); - const nextStatus = await apiClient.post("/roi-calculator/sync", { accessToken }); - statusRef.current = nextStatus; - setStatus(nextStatus); - setSyncError(null); - setSettingsOpen(false); - } catch (reason) { - setError(extractErrorMessage(reason)); - } - }, [accessToken, readOnly]); - - const cancelSync = React.useCallback(async () => { - if (!accessToken || readOnly) return; - try { - const nextStatus = await apiClient.delete("/roi-calculator/sync", { accessToken }); - setStatus(nextStatus); - statusRef.current = nextStatus; - setSyncError(null); - setError(null); - } catch (reason) { - setError(extractErrorMessage(reason)); - } - }, [accessToken, readOnly]); - - const updateIdentity = React.useCallback( - async (payload: ROIIdentityMapUpdate) => { - if (!accessToken || readOnly) return; - const response: ROIIdentityMapResponse = await apiClient.put("/roi-calculator/identity-map", { - accessToken, - body: payload, - }); - setSummary(response.report); - setReportError(null); - reportNeedsRefresh.current = false; - setSettings((current) => (current ? { ...current, identity_map: response.identity_map } : current)); - }, - [accessToken, readOnly], - ); - - const filteredPulls = React.useMemo( - () => (summary ? filterPulls(summary.pulls, query, matchedOnly) : []), - [query, summary, matchedOnly], - ); - - if (error && !settings && !loadingInitialData) { - return ( -
- - Could not load ROI Calculator - {error} - -
- ); - } - - const awaitingLiveData = !sampleSummary && loadingLiveData; - if (!settings || loadingInitialData || awaitingLiveData) { - return ( -
-

Loading ROI Calculatorโ€ฆ

- - -
- ); - } - - const previewSample = async () => { - try { - const response = await apiClient.get("/roi-calculator/report", { - accessToken, - query: { mode: "demo" }, - }); - setSampleSummary(response.report); - setDemoError(null); - setDemo(true); - setView("branches"); - setQuery(""); - } catch (reason) { - setDemoError(`Could not load demo data: ${extractErrorMessage(reason)}`); - } - }; - const resetView = (updated: ROISettings, resetSyncStatus = true) => { - sourceRevision.current += 1; - setSettings(updated); - setSummary(null); - setReportError(null); - reportNeedsRefresh.current = false; - setSettingsOpen(false); - setView("overview"); - if (resetSyncStatus) { - setStatus(IDLE_STATUS); - statusRef.current = IDLE_STATUS; - setSyncError(null); - setError(null); - } - }; - const showLiveStatus = !sampleSummary && !status.running; - const showReportActions = summary !== null || reportError !== null; - const progress = status.total > 0 ? Math.min(100, (status.done / status.total) * 100) : 0; - const statusIsIdleOrComplete = status.phase === "idle" || status.phase === "complete"; - const syncIsUpToDate = !status.running && statusIsIdleOrComplete; - const syncedAt = sampleSummary?.synced_at ?? (syncIsUpToDate ? summary?.synced_at : null); - - return ( - - -
- - - ROI Calculator - - {!sampleSummary && ( -
- {showLiveStatus && ( - - )} - {showReportActions && ( - - )} - {showReportActions && !readOnly && ( - - )} -
- )} -
- - - {summary - ? `${summary.start} through ${summary.end} ยท UTC` - : "Compare AI costs with estimated engineering effort"} - - {syncedAt && ( - - Last synced {formatSyncedAt(syncedAt)} - - )} - -
- {sampleSummary && ( - { - setDemo(null); - setSampleSummary(null); - }} - /> - )} - {adminReadOnly && ( -

- Read-only access. Settings, analysis runs, and email matches are unavailable. -

- )} - - {summary && ( - setView(value as View)}> - - Overview - People - Branches - - {view !== "overview" && ( -
- -
- )} -
- )} - - {!sampleSummary && requestError && ( - - ROI Calculator request failed - {requestError} - - )} - {!sampleSummary && status.error && ( - - Sync failed - {status.error} - - )} - {summary?.warnings.map((warning) => ( - - Sync note - {warning} - - ))} - {!sampleSummary && status.running && ( - - -
-

{status.stage}

-
-
-
-

- {status.done} of {status.total} changes processed ยท {status.reused} reused - {` ยท ${status.elapsed_seconds ?? 0}s elapsed`} - {status.remaining_seconds != null ? ` ยท about ${status.remaining_seconds}s remaining` : ""} -

-
- {!readOnly && ( - - )} - - - )} - - {!summary && !status.running && !reportError ? ( - - ) : null} - {view === "overview" && summary && ( - setView("people")} - onViewBranches={() => setView("branches")} - /> - )} - {view === "branches" && summary && ( - - )} - {view === "people" && summary && ( - setMatchingPerson({ person, login })} - readOnly={readOnly} - /> - )} - - - - Calculator settings - Connect repositories and choose how to estimate effort. - - { - if ( - updated.source_provider !== settings.source_provider || - updated.github_api_url !== settings.github_api_url || - updated.gitlab_api_url !== settings.gitlab_api_url - ) { - resetView(updated, false); - return; - } - setSettings(updated); - }} - onReset={resetView} - onStartSync={startSync} - readOnly={readOnly} - syncDisabled={status.running} - /> - - - setSelectedPull(null)} /> - {!readOnly && ( - (person.email ? [person.email] : [])) ?? []} - onClose={() => setMatchingPerson(null)} - onSave={updateIdentity} - /> - )} - - ); -} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorViews.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorViews.tsx deleted file mode 100644 index b581e080e6c..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorViews.tsx +++ /dev/null @@ -1,569 +0,0 @@ -"use client"; - -import React from "react"; -import { ChevronDown, Download, GitBranch, Search } from "lucide-react"; -import { Button } from "@/components/ui/button"; -import { Input } from "@/components/ui/input"; -import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; -import { - peopleCsv, - isMatchedPerson, - effortNote, - estimateLabel, - formatMoney, - formatNumber, - branchCostLabel, - highestCostPulls, -} from "./roiCalculatorData"; -import type { ROIPerson, ROIPull, ROISummary } from "./roiCalculatorData"; - -type PullSelection = { summary: ROISummary; onSelectPull: (pull: ROIPull) => void }; - -export function ROIOverview({ - summary, - onSelectPull, - onViewPeople, - onViewBranches, -}: PullSelection & { - onViewPeople: () => void; - onViewBranches: () => void; -}) { - const metrics = summary.metrics; - const branches = summary.branch_metrics; - const topPulls = highestCostPulls(summary.pulls); - return ( -
-
-
- - - - -
-
-
-
-

Where AI costs are matched

-

- People use gateway account costs. Branches use tagged requests for these repositories. -

-
-
- - - - View - Matched cost - Unmatched cost - Changes matched - - - - - - - - - {formatMoney(metrics.matched_spend)} - - - {formatMoney(metrics.excluded_spend)} - - - {metrics.matched_prs} / {metrics.merged_prs} - - - - - - - - {formatMoney(branches?.spend)} - - - {formatMoney(branches?.unlinked_spend)} - - - {branches?.matched_pulls ?? 0} / {metrics.merged_prs} - - - -
-
-
- -
- ); -} - -export function ROIBranches({ - summary, - onSelectPull, - pulls, - query, - onQueryChange, - matchedOnly, -}: PullSelection & { - pulls: ROIPull[]; - query: string; - matchedOnly: boolean; - onQueryChange: (query: string) => void; -}) { - return ( -
-
-
- - -
-
- -
- ); -} - -function ROIPulls({ - summary, - onSelectPull, - pulls, - query = "", - onQueryChange, - compact = false, - matchedOnly = false, - onViewBranches, -}: PullSelection & { - pulls: ROIPull[]; - query?: string; - onQueryChange?: (query: string) => void; - compact?: boolean; - matchedOnly?: boolean; - onViewBranches?: () => void; -}) { - const changeName = summary.source_provider === "gitlab" ? "merge request" : "pull request"; - const [pagination, setPagination] = React.useState({ query, visibleCount: 10 }); - const visibleCount = pagination.query === query ? pagination.visibleCount : 10; - const metrics = summary.metrics; - const emptyMessage = compact - ? "No merged changes with tagged costs yet. Open Branches to see how to add tags." - : `No merged ${changeName}s ${matchedOnly ? "from matched people " : ""}in this period.`; - return ( -
-
-
-

{compact ? "Highest-cost changes" : "Costs by branch"}

-

- {compact - ? "Merged work ranked by tagged AI cost" - : `${pulls.length} of ${metrics.merged_prs} ${changeName}s`} -

-
- {compact ? ( - - ) : ( -
-
- )} -
-
- - - - - {changeName === "merge request" ? "Merge request" : "Pull request"} - - AI cost - Estimated effort - - - - {pulls.slice(0, visibleCount).map((pull) => ( - - - - - - {branchCostLabel(pull)} - - {estimateLabel(pull.estimate)} - - ))} - {pulls.length === 0 && ( - - - {query ? `No matching ${changeName}s. Try another search.` : emptyMessage} - - - )} - -
- {pulls.length > visibleCount && ( -
-

- Showing {Math.min(visibleCount, pulls.length)} of {pulls.length} -

- -
- )} -
-
- ); -} - -function ROIMetrics({ summary, branchMode }: { summary: ROISummary; branchMode: boolean }) { - const branches = summary.branch_metrics; - const metrics = summary.metrics; - return ( -
- - - - -
- ); -} - -function MetricCard({ - title, - value, - description, - primary = false, -}: { - title: string; - value: string; - description: string; - primary?: boolean; -}) { - return ( -
-
{title}
-
- {value} -
-
{description}
-
- ); -} - -function ROIComparison({ - summary, - branchMode, - onViewPeople, -}: { - summary: ROISummary; - branchMode: boolean; - onViewPeople?: () => void; -}) { - const branches = summary.branch_metrics; - const metrics = summary.metrics; - const unavailableRate = - metrics.output_hours > 0 - ? "Spend per estimated hour is unavailable until all selected repositories can be read." - : "Match gateway accounts to calculate costs per estimated hour."; - return ( -
- {!branchMode && metrics.cohort_people === 0 && onViewPeople && ( -
-

Match people to gateway accounts to see their AI costs.

- -
- )} -
- - - - - {formatMoney(branchMode ? branches?.unlinked_spend : metrics.excluded_spend)}{" "} - {branchMode ? "in unmatched costs" : "excluded from calculation"} - - -
-

{effortNote(summary.effort_basis)}

- {branchMode ? ( - <> -

- Only branches with matched request costs and complete effort estimates enter the calculation. Costs - cover retained requests in this reportโ€™s UTC dates, not the branchโ€™s lifetime. -

-

- Open a change below to find its repository and branch tags. Send both with each gateway request. Email - matching is not required. Shared branches stay ambiguous so their costs are not counted twice. -

-

- {formatMoney(branches?.total_tagged_spend)} in tagged costs was found for these repositories. Costs - without a unique, fully estimated change stay unmatched. -

- {(summary.unlinked_branches?.length ?? 0) > 0 && ( -
-

Unmatched branches

-
    - {summary.unlinked_branches?.map((row) => ( -
  • - - {row.repo} -
    - {row.branch} -
    - {formatMoney(row.spend)} -
  • - ))} -
-
- )} - - ) : ( - <> -

- {metrics.cost_per_hour != null - ? `${formatMoney(metrics.matched_spend)} AI costs รท ${formatNumber(metrics.output_hours)} estimated hours = ${formatMoney(metrics.cost_per_hour)} per estimated hour.` - : unavailableRate} -

-

- Includes {metrics.cohort_people} matched {metrics.cohort_people === 1 ? "person" : "people"} with - complete estimates. Costs include each personโ€™s full gateway usage across repositories during this UTC - period. -

- {onViewPeople && ( - - )} - - )} -
-
-
- ); -} - -export function ROIPeopleView({ - summary, - identityMap, - onMatch, - readOnly = false, - matchedOnly = true, -}: { - summary: ROISummary; - identityMap: Record; - onMatch: (person: ROIPerson, login: string) => void; - readOnly?: boolean; - matchedOnly?: boolean; -}) { - const people = matchedOnly ? summary.people.filter(isMatchedPerson) : summary.people; - const exportCsv = () => { - const url = URL.createObjectURL(new Blob([peopleCsv({ ...summary, people })], { type: "text/csv;charset=utf-8" })); - const link = document.createElement("a"); - link.href = url; - link.download = "litellm-roi.csv"; - link.click(); - window.setTimeout(() => URL.revokeObjectURL(url), 1000); - }; - return ( -
-
- - -
-
-
-

People and account matches

-

Select a person to match their gateway email.

-
- -
-
- - - - Person - AI cost - Estimated effort - Cost / est. hour - - - - {people.map((person) => ( - - -
- {person.logins.length ? ( - person.logins.map((login) => - readOnly ? ( - {login} - ) : ( - - ), - ) - ) : ( - Unassigned gateway spend - )} - - {isMatchedPerson(person) ? "Matched" : "Unmatched"} - -
-

- {person.email || "No public email"} - {person.logins.some((login) => identityMap[login.toLowerCase()]) ? " ยท Manual match" : ""} -

-
- {formatMoney(person.spend)} - - {person.estimated_prs > 0 ? `${formatNumber(person.hours)} hrs` : "โ€”"} -

- {person.prs} {person.prs === 1 ? "change" : "changes"} - {person.pending_prs > 0 ? ` ยท ${person.pending_prs} pending` : ""} -

-
- - {formatMoney(person.cost_per_hour)} - {!person.eligible &&

Not included

} -
-
- ))} - {people.length === 0 && ( - - - {matchedOnly - ? "No matched people in this period. Turn off the filter to see all contributors" - : "No people in this period"} - - - )} -
-
-
-
- - -
-

- Matches use the authorโ€™s public profile email - {summary.source_provider === "gitlab" - ? "." - : " or commit emails associated with their GitHub account."}{" "} - Private, noreply, and ambiguous emails stay unmatched. Manual matches take priority. -

-

- Costs include each personโ€™s full gateway usage for this period. People without a spend record or with - incomplete estimates are not included in the calculation. -

-

{effortNote(summary.effort_basis)}

-
-
-
- ); -} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROISettingsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROISettingsPanel.tsx deleted file mode 100644 index 535655ed574..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROISettingsPanel.tsx +++ /dev/null @@ -1,566 +0,0 @@ -"use client"; - -import React from "react"; - -import { SearchSelect } from "@/components/shared/SearchSelect"; -import { estimatorModelOptions } from "./roiCalculatorData"; - -import { apiClient } from "@/components/networking"; -import { extractErrorMessage } from "@/utils/errorUtils"; -import { Button } from "@/components/ui/button"; -import { Card, CardContent, CardDescription, CardHeader } from "@/components/ui/card"; -import { Input } from "@/components/ui/input"; -import { Label } from "@/components/ui/label"; -import { - Dialog, - DialogContent, - DialogHeader, - DialogTitle, - DialogDescription, - DialogFooter, -} from "@/components/ui/dialog"; -import { Textarea } from "@/components/ui/textarea"; -import type { ROIRepository, ROIRepositoriesResponse, ROISettings, ROISettingsUpdate } from "./roiCalculatorData"; - -export default function ROISettingsPanel({ - accessToken, - initialSettings, - onboarding, - onSaved, - onReset, - onStartSync, - readOnly, - syncDisabled, -}: { - accessToken: string | null; - initialSettings: ROISettings; - onboarding: boolean; - onSaved: (settings: ROISettings) => void; - onReset: (settings: ROISettings) => void; - onStartSync: () => Promise; - readOnly: boolean; - syncDisabled: boolean; -}) { - const [provider, setProvider] = React.useState<"github" | "gitlab">(initialSettings.source_provider ?? "github"); - const sourceName = provider === "gitlab" ? "GitLab" : "GitHub"; - const savedToken = provider === "gitlab" ? initialSettings.has_gitlab_token : initialSettings.has_github_token; - const savedUrl = provider === "gitlab" ? initialSettings.gitlab_api_url : initialSettings.github_api_url; - const initialStep = savedToken ? 1 : 0; - const [step, setStep] = React.useState(initialSettings.ready ? 2 : initialStep); - const [apiUrl, setApiUrl] = React.useState(savedUrl ?? "https://gitlab.com/api/v4"); - const [token, setToken] = React.useState(""); - const [clearToken, setClearToken] = React.useState(false); - const [repos, setRepos] = React.useState(initialSettings.repos); - const [model, setModel] = React.useState(initialSettings.estimator_model); - const [prompt, setPrompt] = React.useState(initialSettings.estimator_prompt); - const [backfillDays, setBackfillDays] = React.useState(String(initialSettings.backfill_days)); - const [intervalHours, setIntervalHours] = React.useState( - String((initialSettings.update_interval_minutes ?? 1440) / 60), - ); - const [estimatorKey, setEstimatorKey] = React.useState(""); - const [clearEstimatorKey, setClearEstimatorKey] = React.useState(false); - const [repositoryName, setRepositoryName] = React.useState(""); - const [resetOpen, setResetOpen] = React.useState(false); - const [repositoryQuery, setRepositoryQuery] = React.useState(""); - const [repositoryPage, setRepositoryPage] = React.useState(1); - const [availableRepos, setAvailableRepos] = React.useState([]); - const [hasMoreRepos, setHasMoreRepos] = React.useState(false); - const [busy, setBusy] = React.useState(false); - const [error, setError] = React.useState(null); - const [message, setMessage] = React.useState(null); - - const sourceUnchanged = provider === (initialSettings.source_provider ?? "github") && apiUrl === savedUrl; - const credentialsSaved = sourceUnchanged && !token.trim() && !clearToken; - const canLoadRepositories = credentialsSaved && (provider === "gitlab" || savedToken); - const tokenHelp = - provider === "gitlab" - ? "For private projects, use a token with read_api scope and project access." - : "For private repositories, use a token with read access to contents and pull requests."; - - const changeProvider = (next: "github" | "gitlab") => { - setProvider(next); - setApiUrl( - next === "gitlab" - ? initialSettings.gitlab_api_url ?? "https://gitlab.com/api/v4" - : initialSettings.github_api_url, - ); - setToken(""); - setClearToken(false); - setRepos([]); - setAvailableRepos([]); - setHasMoreRepos(false); - }; - - const loadRepositories = async (page: number) => { - if (!accessToken || !canLoadRepositories) return; - try { - setBusy(true); - const response: ROIRepositoriesResponse = await apiClient.get("/roi-calculator/repositories", { - accessToken, - query: { query: repositoryQuery, page }, - }); - setAvailableRepos((current) => (page === 1 ? response.repositories : [...current, ...response.repositories])); - setHasMoreRepos(response.has_more); - setRepositoryPage(page); - setError(null); - } catch (reason) { - setError(extractErrorMessage(reason)); - } finally { - setBusy(false); - } - }; - - const saveSettings = async () => { - if (!accessToken || readOnly) return false; - const tokenValue = clearToken ? null : token.trim() || undefined; - const body: ROISettingsUpdate = { - source_provider: provider, - ...(provider === "gitlab" ? { gitlab_api_url: apiUrl } : { github_api_url: apiUrl }), - repos, - estimator_model: model, - estimator_prompt: prompt, - backfill_days: Number(backfillDays), - update_interval_minutes: Number(intervalHours) * 60, - ...(clearEstimatorKey ? { estimator_key: null } : {}), - ...(estimatorKey.trim() ? { estimator_key: estimatorKey.trim() } : {}), - ...(provider === "gitlab" ? { gitlab_token: tokenValue } : { github_token: tokenValue }), - }; - try { - setBusy(true); - const updated: ROISettings = await apiClient.put("/roi-calculator/settings", { accessToken, body }); - onSaved(updated); - setToken(""); - setEstimatorKey(""); - setClearEstimatorKey(false); - setClearToken(false); - setMessage("Settings saved."); - setError(null); - return true; - } catch (reason) { - setError(extractErrorMessage(reason)); - setMessage(null); - return false; - } finally { - setBusy(false); - } - }; - - const submit = async (event: React.FormEvent) => { - event.preventDefault(); - if (!(await saveSettings())) return; - if (onboarding && step === 0) { - setStep(1); - if (!(token.trim() || savedToken) || clearToken) return; - try { - const result = await apiClient.get("/roi-calculator/repositories", { accessToken }); - setAvailableRepos(result.repositories); - setHasMoreRepos(result.has_more); - setRepositoryPage(1); - setStep(1); - } catch (reason) { - setError(extractErrorMessage(reason)); - } - } else if (onboarding && step === 1) setStep(2); - else if (onboarding) await onStartSync(); - }; - - const saveAndRun = async () => { - if (await saveSettings()) await onStartSync(); - }; - - const testConnections = async () => { - if (!(await saveSettings())) return; - setBusy(true); - try { - await apiClient.post("/roi-calculator/connections/test", { accessToken }); - setMessage("Gateway model and selected repositories are available."); - } catch (reason) { - setError(extractErrorMessage(reason)); - } finally { - setBusy(false); - } - }; - - const resetSetup = async () => { - setBusy(true); - try { - const updated = await apiClient.post("/roi-calculator/setup/reset", { accessToken }); - setRepos([]); - setStep(0); - setResetOpen(false); - onReset(updated); - } catch (reason) { - setError(extractErrorMessage(reason)); - } finally { - setBusy(false); - } - }; - - const toggleRepository = (name: string) => { - setRepos((current) => (current.includes(name) ? current.filter((repo) => repo !== name) : [...current, name])); - }; - - const formDisabled = busy || syncDisabled; - const runDisabled = formDisabled || !repos.length || !model; - const sourceUrlChanged = apiUrl !== savedUrl; - const missingReplacementToken = savedToken && sourceUrlChanged && !token.trim(); - const stepReady = [true, repos.length > 0, Boolean(model)][step]; - const onboardingLabel = step < 2 ? "Continue" : "Start backfill"; - const submitLabel = onboarding ? onboardingLabel : "Save settings"; - - return ( - - {onboarding && ( - -

- {["Connect your repositories", "Choose repositories", "Choose an estimator"][step]} -

- - Your gateway is already connected. Choose a source and an estimator for your first report. - -
- )} - - {error && ( -

- {error} -

- )} - {message && ( -

- {message} -

- )} - {onboarding && ( -

Step {step + 1} of 3 ยท Source / Repositories / Estimator

- )} -
void submit(event)}> -
- {(!onboarding || step === 0) && ( -
- {!onboarding &&

Connection

} -
- - - {provider !== (initialSettings.source_provider ?? "github") && !onboarding && ( -

- Switching source starts a new report and resets email matches. -

- )} -
-
- Self-hosted {sourceName} -
- - setApiUrl(event.target.value)} - /> -
-
-
- - { - setToken(event.target.value); - setClearToken(false); - }} - placeholder={savedToken ? "Token saved" : `Enter a ${sourceName} token`} - /> -

- {savedToken ? "A token is saved securely and is never shown here." : tokenHelp} -

- {missingReplacementToken && ( -

- Changing the API URL clears the saved token. Enter a replacement token to keep access. -

- )} - {savedToken && ( - - )} -
-
- )} - {(!onboarding || step === 1) && ( -
-

Repositories

- -
- setRepositoryQuery(event.target.value)} - placeholder="Search repositories" - /> - -
- {!canLoadRepositories && ( -

- {provider === "github" && !savedToken - ? "Save a GitHub token to browse repositories, or add a public repository by name." - : "Save the source and connection settings before loading repositories."} -

- )} - {repos.length > 0 && ( -
- {repos.map((repo) => ( - - ))} -
- )} -
- -
- setRepositoryName(e.target.value)} - /> - -
-
- {availableRepos.length > 0 && ( -
- {availableRepos.map((repository) => ( - - ))} -
- )} - {hasMoreRepos && ( - - )} -
- )} - {(!onboarding || step === 2) && ( -
- {!onboarding &&

Estimation and updates

} -
- - setModel(value ?? "")} - placeholder="Search estimator models" - emptyText="No matching models configured on this gateway" - disabled={readOnly} - className="h-9" - /> -

- We recommend GPT-6 Luna for estimating PR effort. Choose a model configured on your gateway. -

-
-
- Advanced estimator options -
- -