diff --git a/.circleci/scripts/classify_changes.sh b/.circleci/scripts/classify_changes.sh index ad265a5e39f..8c2ac019b99 100755 --- a/.circleci/scripts/classify_changes.sh +++ b/.circleci/scripts/classify_changes.sh @@ -14,12 +14,12 @@ while IFS= read -r file || [ -n "$file" ]; do [ -n "$file" ] || continue 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/test_litellm/test_circleci_path_filter.py | tests/test_litellm/test_detect_changes.py) + 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) has_mcp_dependencies=true ;; esac case "$file" in tests/e2e/*/*.py) : ;; - tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/test_litellm/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock) + tests/e2e/*.py | tests/code_coverage_tests/test_provider_cache.py | tests/code_coverage_tests/test_provider_replay_harness.py | tests/unit/test_circleci_path_filter.py | .circleci/* | pyproject.toml | uv.lock) has_provider_harness=true ;; esac case "$file" in diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh index f2ee7550df3..e9e5dd3d66b 100755 --- a/.circleci/scripts/unit_selection.sh +++ b/.circleci/scripts/unit_selection.sh @@ -7,7 +7,10 @@ legacy_flags=( caching-local enterprise-package enterprise-routing + llm-other-providers + llm-vertex-ai mcp-integration + misc proxy-db-auth-checks proxy-db-budgets proxy-db-custom-logging @@ -22,6 +25,7 @@ legacy_flags=( proxy-db-proxy-utils proxy-extras proxy-infra + responses-caching-types ) legacy_paths() { @@ -36,6 +40,7 @@ legacy_paths() { echo tests/unit/enterprise/proxy/test_audit_logging_endpoints.py echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;; enterprise-routing) + echo tests/unit/google_genai 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 @@ -47,10 +52,31 @@ legacy_paths() { 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 ;; + 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/images + echo tests/unit/interactions + echo tests/unit/messages + echo tests/unit/rag + echo tests/unit/rerank_api + 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 @@ -113,6 +139,7 @@ legacy_paths() { 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 ;; + responses-caching-types) echo tests/unit/types ;; *) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;; esac } diff --git a/.circleci/tests.yml b/.circleci/tests.yml index 264d7695a94..994d67da64d 100644 --- a/.circleci/tests.yml +++ b/.circleci/tests.yml @@ -341,6 +341,7 @@ workflows: flag: - enterprise-package - proxy-infra + - responses-caching-types - proxy-db-auth-checks - proxy-db-jwt-and-keys - proxy-db-proxy-server-core @@ -353,6 +354,28 @@ workflows: - 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-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 diff --git a/.github/merge-smoke-tests.json b/.github/merge-smoke-tests.json index 6088953b7eb..a563424c230 100644 --- a/.github/merge-smoke-tests.json +++ b/.github/merge-smoke-tests.json @@ -1,12 +1,12 @@ { "cases": { - "CHAT-JSON": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport", - "CHAT-TEXT-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport", - "CHAT-TOOL-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport", + "CHAT-JSON": "tests/unit/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport", + "CHAT-TEXT-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport", + "CHAT-TOOL-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport", "MODEL-ALLOW": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_allows_listed_model_for_key", "MODEL-DENY": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_denials_return_forbidden[key-key_model_access_denied]", - "COST-EXPLICIT": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones", - "COST-ZERO": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero", + "COST-EXPLICIT": "tests/unit/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones", + "COST-ZERO": "tests/unit/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero", "LOG-CONTENT-ON": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on", "LOG-CONTENT-OFF": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off", "CALLBACK-SUCCESS": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger", diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index ef1dc53b4a6..fac0d766535 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -13,12 +13,13 @@ on: have its path existence-checked like any other token. required: true type: string - fork-flag: + unit-flag: description: >- Codecov flag of the `.circleci/tests.yml` job that now owns part of - this shard. CircleCI does not run on pull requests from forks, so on - those events this shard also runs the files - `.circleci/scripts/unit_selection.sh` lists for the flag. + 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: "" @@ -175,8 +176,7 @@ jobs: timeout-minutes: ${{ inputs.timeout-minutes }} env: TEST_PATH: ${{ inputs.test-path }} - FORK_FLAG: ${{ inputs.fork-flag }} - IS_FORK: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.repo.full_name != github.repository }} + UNIT_FLAG: ${{ inputs.unit-flag }} MAX_FAILURES: ${{ inputs.max-failures }} WORKERS: ${{ inputs.workers }} RERUNS: ${{ inputs.reruns }} @@ -186,11 +186,11 @@ jobs: run: | echo "has-coverage=false" >> "$GITHUB_OUTPUT" selection="${TEST_PATH}" - if [ "${IS_FORK}" = "true" ] && [ -n "${FORK_FLAG}" ]; then - selection="${TEST_PATH} $(bash .circleci/scripts/unit_selection.sh "${FORK_FLAG}" | tr '\n' ' ')" + 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 on this event (CircleCI flag ${FORK_FLAG:-none} owns it); nothing to run" + echo "shard selection is empty; nothing to run" exit 0 fi pytest_args=() diff --git a/.github/workflows/test-redis-compat.yml b/.github/workflows/test-redis-compat.yml index 25fb8f8bce3..2f5ce4d441a 100644 --- a/.github/workflows/test-redis-compat.yml +++ b/.github/workflows/test-redis-compat.yml @@ -10,7 +10,7 @@ on: - "litellm/_redis_credential_provider.py" - "litellm/caching/redis_cache.py" - "litellm/caching/evicted_client_closer.py" - - "tests/test_litellm/test_redis.py" + - "tests/unit/test_redis.py" - "tests/local_testing/test_caching.py" - "tests/test_litellm/caching/test_redis_connection_pool.py" - "tests/test_litellm/caching/test_redis_cluster_cache.py" @@ -84,7 +84,7 @@ jobs: run: | redis-server --version uv run --no-sync pytest \ - tests/test_litellm/test_redis.py \ + tests/unit/test_redis.py \ tests/test_litellm/caching/test_redis_connection_pool.py \ tests/test_litellm/caching/test_redis_cluster_cache.py \ tests/test_litellm/caching/test_evicted_client_closer.py \ diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 86b385d91a7..da4477b6947 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -22,9 +22,10 @@ concurrency: # # `.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. CircleCI does not build pull requests from forks, so `fork-flag` -# makes the shard run that list there. `test-path` keeps the files that still -# reach real providers and never left tests/proxy_unit_tests. +# 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. @@ -78,7 +79,7 @@ jobs: # Must run serially — event-loop conflict with the logging worker. - test-group: key-generation test-path: "" - fork-flag: proxy-db-key-generation + unit-flag: proxy-db-key-generation workers: 0 dist: loadscope timeout: 20 @@ -86,13 +87,13 @@ jobs: # ---- auth: split into 2 shards ---- - test-group: auth-checks test-path: "" - fork-flag: proxy-db-auth-checks + unit-flag: proxy-db-auth-checks workers: 4 dist: loadscope timeout: 15 - test-group: jwt-and-keys test-path: "" - fork-flag: proxy-db-jwt-and-keys + unit-flag: proxy-db-jwt-and-keys workers: 4 dist: loadscope timeout: 15 @@ -100,7 +101,7 @@ jobs: # ---- test_proxy_utils.py, single shard, worksteal distribution ---- - test-group: proxy-utils test-path: "" - fork-flag: proxy-db-proxy-utils + unit-flag: proxy-db-proxy-utils workers: 4 dist: worksteal timeout: 15 @@ -108,13 +109,13 @@ jobs: # ---- proxy server: split into 2 shards ---- - test-group: proxy-server-core test-path: "tests/proxy_unit_tests/test_proxy_server_gemini_pass_through.py" - fork-flag: proxy-db-proxy-server-core + unit-flag: proxy-db-proxy-server-core workers: 4 dist: loadscope timeout: 15 - test-group: proxy-runtime test-path: "" - fork-flag: proxy-db-proxy-runtime + unit-flag: proxy-db-proxy-runtime workers: 4 dist: loadscope timeout: 15 @@ -122,20 +123,20 @@ jobs: # ---- logging: split into 2 shards ---- - test-group: custom-logging test-path: "tests/proxy_unit_tests/test_proxy_custom_logger.py" - fork-flag: proxy-db-custom-logging + unit-flag: proxy-db-custom-logging workers: 4 dist: loadscope timeout: 15 - test-group: logging-misc test-path: "" - fork-flag: proxy-db-logging-misc + unit-flag: proxy-db-logging-misc workers: 4 dist: loadscope timeout: 15 - test-group: db-and-spend test-path: "" - fork-flag: proxy-db-db-and-spend + unit-flag: proxy-db-db-and-spend workers: 4 dist: loadscope timeout: 15 @@ -143,27 +144,27 @@ jobs: # ---- guardrails + budget + hooks: split into 2 ---- - test-group: guardrails-hooks test-path: "" - fork-flag: proxy-db-guardrails-hooks + unit-flag: proxy-db-guardrails-hooks workers: 4 dist: loadscope timeout: 15 - test-group: budgets test-path: "" - fork-flag: proxy-db-budgets + unit-flag: proxy-db-budgets workers: 4 dist: loadscope timeout: 15 - test-group: endpoints-and-responses test-path: "tests/proxy_unit_tests/test_proxy_exception_mapping.py" - fork-flag: proxy-db-endpoints-and-responses + unit-flag: proxy-db-endpoints-and-responses workers: 4 dist: loadscope timeout: 15 uses: ./.github/workflows/_test-unit-base.yml with: test-path: ${{ matrix.test-path }} - fork-flag: ${{ matrix.fork-flag }} + 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 126a6e26e6f..2fa05879350 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -36,9 +36,9 @@ concurrency: # Folding it in here is a follow-up, together with generalising that guard into # assert_ci_coverage.py. # -# `fork-flag` names the `.circleci/tests.yml` job that now runs part of the -# shard under the same Codecov flag. CircleCI does not build pull requests from -# forks, so the shard still runs those files there and skips them elsewhere. +# `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 }} @@ -52,8 +52,8 @@ jobs: include: - shard: mcp-integration artifact-name: mcp-integration - test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client" - fork-flag: mcp-integration + test-path: "tests/mcp_tests" + unit-flag: mcp-integration workers: 2 reruns: 0 timeout-minutes: 20 @@ -70,10 +70,9 @@ jobs: - shard: enterprise-routing artifact-name: enterprise-routing test-path: >- - tests/test_litellm/google_genai tests/test_litellm/router_utils tests/test_litellm/router_strategy - fork-flag: enterprise-routing + unit-flag: enterprise-routing workers: 2 reruns: 2 timeout-minutes: 20 @@ -90,6 +89,7 @@ jobs: - shard: Vertex AI artifact-name: llm-vertex-ai test-path: "tests/test_litellm/llms/vertex_ai" + unit-flag: llm-vertex-ai workers: 1 reruns: 2 timeout-minutes: 20 @@ -98,6 +98,7 @@ jobs: - shard: All Other Providers artifact-name: llm-other-providers test-path: "tests/test_litellm/llms --ignore=tests/test_litellm/llms/vertex_ai" + unit-flag: llm-other-providers workers: 2 reruns: 2 timeout-minutes: 20 @@ -106,26 +107,13 @@ jobs: - shard: misc artifact-name: misc test-path: >- - tests/test_litellm/batches tests/test_litellm/secret_managers - tests/test_litellm/a2a_protocol - tests/test_litellm/chat_completions - tests/test_litellm/completion_extras - tests/test_litellm/containers - tests/test_litellm/endpoints - tests/test_litellm/files - tests/test_litellm/images tests/test_litellm/interactions - tests/test_litellm/messages - tests/test_litellm/embeddings tests/test_litellm/ocr tests/test_litellm/passthrough - tests/test_litellm/rag - tests/test_litellm/rerank_api tests/test_litellm/rust_bridge - tests/test_litellm/vector_stores - tests/test_litellm/videos tests/test_litellm/test_*.py + unit-flag: misc workers: 2 reruns: 2 timeout-minutes: 20 @@ -205,7 +193,7 @@ jobs: tests/test_litellm/proxy/types_utils tests/test_litellm/proxy/logging_endpoints tests/test_litellm/proxy/test_*.py - fork-flag: proxy-infra + unit-flag: proxy-infra workers: 4 reruns: 2 timeout-minutes: 20 @@ -214,7 +202,7 @@ jobs: - shard: caching-local artifact-name: caching-local test-path: "" - fork-flag: caching-local + unit-flag: caching-local workers: 2 reruns: 2 timeout-minutes: 20 @@ -223,7 +211,7 @@ jobs: - shard: proxy-extras artifact-name: proxy-extras test-path: "" - fork-flag: proxy-extras + unit-flag: proxy-extras workers: 2 reruns: 2 timeout-minutes: 20 @@ -232,7 +220,7 @@ jobs: - shard: enterprise-package artifact-name: enterprise-package test-path: "" - fork-flag: enterprise-package + unit-flag: enterprise-package workers: 4 reruns: 2 timeout-minutes: 20 @@ -243,7 +231,7 @@ jobs: test-path: >- tests/test_litellm/responses tests/test_litellm/caching - tests/test_litellm/types + unit-flag: responses-caching-types workers: 2 reruns: 2 timeout-minutes: 20 @@ -251,7 +239,7 @@ jobs: uses: ./.github/workflows/_test-unit-base.yml with: test-path: ${{ matrix.test-path }} - fork-flag: ${{ matrix.fork-flag || '' }} + unit-flag: ${{ matrix.unit-flag || '' }} workers: ${{ matrix.workers }} reruns: ${{ matrix.reruns }} timeout-minutes: ${{ matrix.timeout-minutes }} diff --git a/Makefile b/Makefile index 28daf589a23..e86047b1987 100644 --- a/Makefile +++ b/Makefile @@ -314,7 +314,7 @@ test-unit: install-test-deps # Matrix test targets (matching CI workflow groups) test-unit-llms: install-test-deps - $(UV_RUN) pytest tests/test_litellm/llms --tb=short -vv -n 4 --durations=20 + $(UV_RUN) pytest tests/unit/llms --tb=short -vv -n 4 --durations=20 test-unit-proxy-guardrails: install-test-deps $(UV_RUN) pytest tests/test_litellm/proxy/guardrails tests/test_litellm/proxy/management_endpoints tests/test_litellm/proxy/management_helpers --tb=short -vv -n 4 --durations=20 @@ -332,10 +332,10 @@ test-unit-core-utils: install-test-deps $(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20 test-unit-other: install-test-deps - $(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/test_litellm/vector_stores tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface tests/test_litellm/completion_extras tests/test_litellm/containers tests/unit/enterprise tests/test_litellm/experimental_mcp_client tests/test_litellm/google_genai tests/test_litellm/images tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/test_litellm/types --tb=short -vv -n 4 --durations=20 + $(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20 test-unit-root: install-test-deps - $(UV_RUN) pytest tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20 + $(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20 # Proxy unit tests (tests/unit/proxy split alphabetically) test-proxy-unit-a: install-test-deps diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index e7d911f5fd9..3677d1d654f 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3166,10 +3166,11 @@ dependencies = [ "aws-smithy-types", "bytes", "futures-util", + "proptest", "rstest", - "sse-stream", "thiserror 2.0.19", "tokio", + "tokio-util", ] [[package]] @@ -5468,19 +5469,6 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "sse-stream" -version = "0.2.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c25ac7aff0abd1dbc474536e40416e1102c7dd9bfba0b9861c6d357f835dcfb4" -dependencies = [ - "bytes", - "futures-util", - "http-body 1.1.0", - "http-body-util", - "pin-project-lite", -] - [[package]] name = "stable_deref_trait" version = "1.2.1" diff --git a/litellm-rust/crates/core/tests/messages/request.rs b/litellm-rust/crates/core/tests/messages/request.rs index 2927356b773..d37910d4ac4 100644 --- a/litellm-rust/crates/core/tests/messages/request.rs +++ b/litellm-rust/crates/core/tests/messages/request.rs @@ -398,7 +398,7 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i api_key: Some("sk".into()), api_base: Some(upstream.uri()), shaping: MessagesShaping { - capabilities: capabilities.clone(), + capabilities, drop_params, ..MessagesShaping::default() }, diff --git a/litellm-rust/crates/framer/Cargo.toml b/litellm-rust/crates/framer/Cargo.toml index 62bfcc7da3d..e11f2c02a97 100644 --- a/litellm-rust/crates/framer/Cargo.toml +++ b/litellm-rust/crates/framer/Cargo.toml @@ -8,16 +8,17 @@ repository.workspace = true [features] default = ["aws", "sse"] aws = ["dep:aws-smithy-eventstream", "dep:aws-smithy-types"] -sse = ["dep:sse-stream"] +sse = [] [dependencies] aws-smithy-eventstream = { version = "=0.61.4", optional = true } aws-smithy-types = { version = "1.6.1", optional = true } bytes = "1" futures-util.workspace = true -sse-stream = { version = "=0.2.6", optional = true } thiserror.workspace = true +tokio-util = { version = "0.7", features = ["codec", "io"] } [dev-dependencies] +proptest.workspace = true rstest.workspace = true tokio.workspace = true diff --git a/litellm-rust/crates/framer/src/aws_event_stream.rs b/litellm-rust/crates/framer/src/aws_event_stream.rs index efd7adeb64b..405ec2d5ad2 100644 --- a/litellm-rust/crates/framer/src/aws_event_stream.rs +++ b/litellm-rust/crates/framer/src/aws_event_stream.rs @@ -1,66 +1,47 @@ -use bytes::{Buf, Bytes, BytesMut}; -use futures_util::{Stream, StreamExt}; +use aws_smithy_eventstream::frame::{read_message_from, write_message_to}; +pub use aws_smithy_types::event_stream::{Header, HeaderValue, Message}; +use bytes::BytesMut; +use tokio_util::codec::{Decoder, Encoder}; -use aws_smithy_eventstream::frame::read_message_from; -use aws_smithy_types::event_stream::Header; - -use crate::{Error, Framer}; +use crate::EventStreamError; +const MIN_FRAME_BYTES: usize = 16; const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024; -#[derive(Clone, Debug, PartialEq)] -pub struct AwsEventStreamFrame { - pub headers: Vec
, - pub payload: Bytes, -} - #[derive(Clone, Copy, Debug, Default)] -pub struct AwsEventStreamFramer; +pub struct AwsEventStreamCodec; -impl Framer for AwsEventStreamFramer { - type Frame = AwsEventStreamFrame; +impl Decoder for AwsEventStreamCodec { + type Item = Message; + type Error = EventStreamError; - fn frame(self, input: S) -> impl Stream> + Send - where - S: Stream> + Send, - B: Buf + Send, - E: std::error::Error + Send + Sync + 'static, - { - futures_util::stream::try_unfold( - (Box::pin(input), BytesMut::new()), - |(mut input, mut buffer)| async move { - loop { - if buffer.len() >= 4 { - let length = (&buffer[..4]).get_u32() as usize; - if !(16..=MAX_FRAME_BYTES).contains(&length) { - return Err(Error::InvalidLength(length)); - } - if buffer.len() >= length { - let raw = buffer.split_to(length).freeze(); - let message = read_message_from(raw)?; - let frame = AwsEventStreamFrame { - headers: message.headers().to_vec(), - payload: message.payload().clone(), - }; - return Ok(Some((frame, (input, buffer)))); - } - } - match input.next().await { - Some(Ok(mut chunk)) => { - while chunk.has_remaining() { - let bytes = chunk.chunk(); - buffer.extend_from_slice(bytes); - let length = bytes.len(); - chunk.advance(length); - } - } - Some(Err(error)) => return Err(Error::Body(Box::new(error))), - None if buffer.is_empty() => return Ok(None), - None => return Err(Error::Truncated), - } - } - }, - ) - .fuse() + fn decode(&mut self, src: &mut BytesMut) -> Result, EventStreamError> { + let Some(prefix) = src.first_chunk::<4>() else { + return Ok(None); + }; + let length = u32::from_be_bytes(*prefix) as usize; + if !(MIN_FRAME_BYTES..=MAX_FRAME_BYTES).contains(&length) { + return Err(EventStreamError::InvalidLength(length)); + } + if src.len() < length { + return Ok(None); + } + Ok(Some(read_message_from(src.split_to(length).freeze())?)) + } + + fn decode_eof(&mut self, src: &mut BytesMut) -> Result, EventStreamError> { + match self.decode(src)? { + Some(message) => Ok(Some(message)), + None if src.is_empty() => Ok(None), + None => Err(EventStreamError::Truncated), + } + } +} + +impl Encoder for AwsEventStreamCodec { + type Error = EventStreamError; + + fn encode(&mut self, message: Message, dst: &mut BytesMut) -> Result<(), EventStreamError> { + Ok(write_message_to(&message, dst)?) } } diff --git a/litellm-rust/crates/framer/src/error.rs b/litellm-rust/crates/framer/src/error.rs index b1f7ed96c5a..879d7557671 100644 --- a/litellm-rust/crates/framer/src/error.rs +++ b/litellm-rust/crates/framer/src/error.rs @@ -1,17 +1,21 @@ +#[cfg(feature = "sse")] #[derive(Debug, thiserror::Error)] -pub enum Error { - #[cfg(feature = "sse")] - #[error("SSE framing failed: {0}")] - Sse(#[from] sse_stream::Error), - #[cfg(feature = "aws")] - #[error("AWS EventStream framing failed: {0}")] - Aws(#[from] aws_smithy_eventstream::error::Error), +pub enum SseError { #[error("body stream failed: {0}")] - Body(#[source] Box), - #[cfg(feature = "aws")] + Body(#[from] std::io::Error), + #[error("SSE field is not UTF-8: {0}")] + InvalidUtf8(#[from] std::str::Utf8Error), +} + +#[cfg(feature = "aws")] +#[derive(Debug, thiserror::Error)] +pub enum EventStreamError { + #[error("body stream failed: {0}")] + Body(#[from] std::io::Error), #[error("invalid AWS EventStream frame length: {0}")] InvalidLength(usize), - #[cfg(feature = "aws")] #[error("truncated AWS EventStream frame")] Truncated, + #[error("malformed AWS EventStream frame: {0}")] + Malformed(#[from] aws_smithy_eventstream::error::Error), } diff --git a/litellm-rust/crates/framer/src/framed.rs b/litellm-rust/crates/framer/src/framed.rs new file mode 100644 index 00000000000..7a19dd40e13 --- /dev/null +++ b/litellm-rust/crates/framer/src/framed.rs @@ -0,0 +1,21 @@ +use std::io; + +use bytes::Buf; +use futures_util::{Stream, StreamExt, TryStreamExt}; +use tokio_util::{ + codec::{Decoder, FramedRead}, + io::StreamReader, +}; + +pub fn frames( + input: S, + codec: D, +) -> impl Stream> + Send +where + S: Stream> + Send, + B: Buf + Send, + E: std::error::Error + Send + Sync + 'static, + D: Decoder + Send, +{ + FramedRead::new(StreamReader::new(input.map_err(io::Error::other)), codec).fuse() +} diff --git a/litellm-rust/crates/framer/src/lib.rs b/litellm-rust/crates/framer/src/lib.rs index 552de419984..223f2f64120 100644 --- a/litellm-rust/crates/framer/src/lib.rs +++ b/litellm-rust/crates/framer/src/lib.rs @@ -1,8 +1,8 @@ mod error; -mod framer; +mod framed; pub use error::*; -pub use framer::*; +pub use framed::frames; #[cfg(feature = "aws")] pub mod aws_event_stream; diff --git a/litellm-rust/crates/framer/src/sse.rs b/litellm-rust/crates/framer/src/sse.rs index 79659f6ce13..6fee1cfab7f 100644 --- a/litellm-rust/crates/framer/src/sse.rs +++ b/litellm-rust/crates/framer/src/sse.rs @@ -1,43 +1,170 @@ -use futures_util::{Stream, StreamExt}; +use std::str; -use crate::{Error, Framer}; +use bytes::{Buf, BufMut, BytesMut}; +use tokio_util::codec::{Decoder, Encoder}; -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct SseFrame { +use crate::SseError; + +const BOM: &[u8] = b"\xEF\xBB\xBF"; + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct SseEvent { pub event: Option, - pub data: Option, + pub data: String, pub id: Option, pub retry: Option, } #[derive(Clone, Copy, Debug, Default)] -pub struct SseFramer; +pub struct SseCodec { + past_bom: bool, +} -impl Framer for SseFramer { - type Frame = SseFrame; +impl Decoder for SseCodec { + type Item = SseEvent; + type Error = SseError; - fn frame(self, input: S) -> impl Stream> + Send - where - S: Stream> + Send, - B: bytes::Buf + Send, - E: std::error::Error + Send + Sync + 'static, - { - let frames = Box::pin(sse_stream::SseStream::from_bytes_stream(input)); - futures_util::stream::try_unfold(frames, |mut frames| async move { - let Some(frame) = frames.next().await else { - return Ok(None); - }; - let frame = frame?; - Ok(Some(( - SseFrame { - event: frame.event, - data: frame.data, - id: frame.id, - retry: frame.retry, - }, - frames, - ))) - }) - .fuse() + fn decode(&mut self, src: &mut BytesMut) -> Result, SseError> { + if !self.skip_bom(src) { + return Ok(None); + } + while let Some(end) = block_end(src) { + let block = src.split_to(end); + let pending = lines(&block) + .map(|(line, _)| line) + .take_while(|line| !line.is_empty()) + .try_fold(Pending::default(), Pending::apply)?; + if let Some(event) = pending.dispatch() { + return Ok(Some(event)); + } + } + Ok(None) + } + + fn decode_eof(&mut self, _pending: &mut BytesMut) -> Result, SseError> { + Ok(None) + } +} + +impl SseCodec { + fn skip_bom(&mut self, src: &mut BytesMut) -> bool { + if self.past_bom { + return true; + } + if src.starts_with(BOM) { + src.advance(BOM.len()); + } else if BOM.starts_with(src) { + return false; + } + self.past_bom = true; + true + } +} + +fn block_end(bytes: &[u8]) -> Option { + lines(bytes) + .find(|(line, _)| line.is_empty()) + .map(|(_, end)| end) +} + +fn lines(bytes: &[u8]) -> impl Iterator { + let mut cursor: usize = 0; + std::iter::from_fn(move || { + let rest = &bytes[cursor..]; + let end = rest.iter().position(|byte| matches!(byte, b'\n' | b'\r'))?; + cursor += end + terminator_len(&rest[end..]); + Some((&rest[..end], cursor)) + }) +} + +fn terminator_len(terminated: &[u8]) -> usize { + match terminated { + [b'\r', b'\n', ..] => 2, + _ => 1, + } +} + +#[derive(Default)] +struct Pending { + event: Option, + data: Option, + id: Option, + retry: Option, +} + +impl Pending { + fn apply(self, line: &[u8]) -> Result { + let (name, value) = split_field(line); + Ok(match name { + b"event" => Self { + event: Some(str::from_utf8(value)?.to_owned()), + ..self + }, + b"data" => Self { + data: Some(append_data(self.data, str::from_utf8(value)?)), + ..self + }, + b"id" if !value.contains(&0) => Self { + id: Some(str::from_utf8(value)?.to_owned()), + ..self + }, + b"retry" => Self { + retry: parse_retry(value).or(self.retry), + ..self + }, + _ => self, + }) + } + + fn dispatch(self) -> Option { + Some(SseEvent { + event: self.event, + data: self.data?, + id: self.id, + retry: self.retry, + }) + } +} + +fn split_field(line: &[u8]) -> (&[u8], &[u8]) { + let Some(colon) = line.iter().position(|byte| *byte == b':') else { + return (line, &[]); + }; + let value = &line[colon + 1..]; + (&line[..colon], value.strip_prefix(b" ").unwrap_or(value)) +} + +fn append_data(buffer: Option, line: &str) -> String { + match buffer { + Some(existing) => format!("{existing}\n{line}"), + None => line.to_owned(), + } +} + +fn parse_retry(value: &[u8]) -> Option { + if !value.iter().all(u8::is_ascii_digit) { + return None; + } + str::from_utf8(value).ok()?.parse().ok() +} + +impl Encoder for SseCodec { + type Error = SseError; + + fn encode(&mut self, event: SseEvent, dst: &mut BytesMut) -> Result<(), SseError> { + if let Some(name) = event.event { + dst.put_slice(format!("event: {name}\n").as_bytes()); + } + for line in event.data.split('\n') { + dst.put_slice(format!("data: {line}\n").as_bytes()); + } + if let Some(id) = event.id { + dst.put_slice(format!("id: {id}\n").as_bytes()); + } + if let Some(retry) = event.retry { + dst.put_slice(format!("retry: {retry}\n").as_bytes()); + } + dst.put_u8(b'\n'); + Ok(()) } } diff --git a/litellm-rust/crates/framer/tests/aws_event_stream.rs b/litellm-rust/crates/framer/tests/aws_event_stream.rs index c90a15a2b0e..d16caa39948 100644 --- a/litellm-rust/crates/framer/tests/aws_event_stream.rs +++ b/litellm-rust/crates/framer/tests/aws_event_stream.rs @@ -4,89 +4,174 @@ mod support; use std::io; -use futures_util::TryStreamExt; -use litellm_framing::aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer}; -use litellm_framing::{Error, Framer}; +use bytes::Bytes; +use futures_util::{StreamExt, TryStreamExt, stream}; +use litellm_framing::{ + EventStreamError, + aws_event_stream::{AwsEventStreamCodec, Header, HeaderValue, Message}, + frames, +}; +use proptest::prelude::*; use rstest::{fixture, rstest}; +use support::{body_cause, cut_at, encode_all, every, input, runtime}; -use support::encode; - -async fn collect_aws(bytes: &[u8], chunk_size: usize) -> Result, Error> { - AwsEventStreamFramer - .frame(futures_util::stream::iter( - bytes.chunks(chunk_size).map(Ok::<_, io::Error>), - )) +async fn collect(pieces: Vec) -> Result, EventStreamError> { + frames(input(pieces), AwsEventStreamCodec) .try_collect() .await } -#[fixture] -fn two_frames() -> Vec { - [encode(b"\xff\x00"), encode(b"second")].concat() +fn message(payload: &[u8]) -> Message { + Message::new(Bytes::copy_from_slice(payload)) + .add_header(Header::new( + ":event-type", + HeaderValue::String("payload".into()), + )) + .add_header(Header::new("sequence", HeaderValue::Int32(7))) } #[fixture] fn payload_frame() -> Vec { - encode(b"payload") + encode_all(AwsEventStreamCodec, [message(b"payload")]) +} + +fn header_value() -> impl Strategy { + prop_oneof![ + "[a-z]{0,8}".prop_map(|text| HeaderValue::String(text.into())), + any::().prop_map(HeaderValue::Int32), + any::().prop_map(HeaderValue::Bool), + proptest::collection::vec(any::(), 0..8) + .prop_map(|bytes| HeaderValue::ByteArray(bytes.into())), + ] +} + +fn arbitrary_message() -> impl Strategy { + ( + proptest::collection::vec(("[a-z:-]{1,12}", header_value()), 0..3), + proptest::collection::vec(any::(), 0..32), + ) + .prop_map(|(headers, payload)| { + headers.into_iter().fold( + Message::new(Bytes::from(payload)), + |message, (name, value)| message.add_header(Header::new(name, value)), + ) + }) +} + +proptest! { + #[test] + fn any_messages_survive_a_round_trip_through_any_cuts( + messages in proptest::collection::vec(arbitrary_message(), 1..4), + cuts in proptest::collection::vec(0_usize..512, 0..4), + ) { + let wire = encode_all(AwsEventStreamCodec, messages.clone()); + let decoded = runtime().block_on(collect(cut_at(&wire, cuts))).unwrap(); + prop_assert_eq!(decoded, messages); + } } #[rstest] -#[case(1)] -#[case(3)] -#[case(12)] -#[case(usize::MAX)] +#[case::prelude_crc(8)] +#[case::message_crc(usize::MAX)] #[tokio::test] -async fn fragmented_and_coalesced_frames_preserve_typed_headers_and_binary_payloads( - two_frames: Vec, - #[case] chunk_size: usize, -) { - let chunk_size = chunk_size.min(two_frames.len()); - let frames = collect_aws(&two_frames, chunk_size).await.unwrap(); - assert_eq!(frames.len(), 2); - assert_eq!(frames[0].payload, &b"\xff\x00"[..]); - assert_eq!(frames[1].payload, "second"); - assert_eq!( - frames[0].headers[0].value().as_string().unwrap().as_str(), - "payload" - ); - assert_eq!(frames[0].headers[1].value().as_int32(), Ok(7)); -} - -#[rstest] -#[case(8)] -#[case(0)] -#[tokio::test] -async fn rejects_corrupt_crcs(payload_frame: Vec, #[case] index: usize) { - let corrupt_index = if index == 0 { - payload_frame.len() - 1 - } else { - index - }; +async fn a_corrupt_crc_is_malformed(payload_frame: Vec, #[case] index: usize) { let mut corrupt = payload_frame; - corrupt[corrupt_index] ^= 1; - assert!(matches!(collect_aws(&corrupt, 3).await, Err(Error::Aws(_)))); -} - -#[rstest] -#[case(0_u32)] -#[case(15)] -#[case(u32::MAX)] -#[tokio::test] -async fn rejects_invalid_lengths(#[case] length: u32) { + let flipped = index.min(corrupt.len() - 1); + corrupt[flipped] ^= 1; assert!(matches!( - collect_aws(&length.to_be_bytes(), 1).await, - Err(Error::InvalidLength(_)) + collect(every(&corrupt, 3)).await, + Err(EventStreamError::Malformed(_)) )); } #[rstest] -#[case(1)] -#[case(3)] -#[case(5)] +#[case::zero(0)] +#[case::below_minimum(15)] +#[case::above_maximum(16 * 1024 * 1024 + 1)] +#[case::u32_max(u32::MAX)] #[tokio::test] -async fn rejects_truncation(payload_frame: Vec, #[case] end: usize) { +async fn a_length_outside_the_frame_bounds_fails_before_buffering(#[case] length: u32) { assert!(matches!( - collect_aws(&payload_frame[..end], 1).await, - Err(Error::Truncated) + collect(every(&length.to_be_bytes(), 1)).await, + Err(EventStreamError::InvalidLength(seen)) if seen == length as usize )); } + +#[rstest] +#[case::before_the_length(1)] +#[case::inside_the_prelude(5)] +#[case::one_byte_short(usize::MAX)] +#[tokio::test] +async fn eof_inside_a_frame_is_truncation(payload_frame: Vec, #[case] end: usize) { + let end = end.min(payload_frame.len() - 1); + assert!(matches!( + collect(every(&payload_frame[..end], 1)).await, + Err(EventStreamError::Truncated) + )); +} + +const FRAME_OVERHEAD_BYTES: usize = 16; +const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024; + +#[tokio::test] +async fn a_frame_at_exactly_the_maximum_length_decodes() { + let largest = Message::new(vec![0xAB; MAX_FRAME_BYTES - FRAME_OVERHEAD_BYTES]); + let wire = encode_all(AwsEventStreamCodec, [largest.clone()]); + assert_eq!(wire.len(), MAX_FRAME_BYTES); + assert_eq!(collect(every(&wire, 1 << 20)).await.unwrap(), vec![largest]); +} + +#[tokio::test] +async fn a_frame_one_byte_over_the_maximum_length_is_rejected_by_its_prelude() { + let oversized = Message::new(vec![0xAB; MAX_FRAME_BYTES - FRAME_OVERHEAD_BYTES + 1]); + let wire = encode_all(AwsEventStreamCodec, [oversized]); + assert!(matches!( + collect(every(&wire[..4], 1)).await, + Err(EventStreamError::InvalidLength(length)) if length == MAX_FRAME_BYTES + 1 + )); +} + +#[tokio::test] +async fn an_empty_body_yields_nothing() { + assert_eq!(collect(vec![]).await.unwrap(), vec![]); +} + +#[tokio::test] +async fn a_complete_frame_precedes_a_truncated_following_frame() { + let wire = encode_all(AwsEventStreamCodec, [message(b"first"), message(b"second")]); + let mut messages = Box::pin(frames( + input(every(&wire[..wire.len() - 1], 3)), + AwsEventStreamCodec, + )); + + assert_eq!(messages.next().await.unwrap().unwrap(), message(b"first")); + assert!(matches!( + messages.next().await, + Some(Err(EventStreamError::Truncated)) + )); + assert!(messages.next().await.is_none()); +} + +#[tokio::test] +async fn a_body_error_after_a_complete_frame_preserves_its_cause() { + let first = encode_all(AwsEventStreamCodec, [message(b"first")]); + let mut messages = Box::pin(frames( + stream::iter([ + Ok(cut_at(&first, [5])[0].clone()), + Ok(cut_at(&first, [5])[1].clone()), + Ok(Bytes::from_static(b"\0\0\0")), + Err(io::Error::new(io::ErrorKind::ConnectionReset, "reset")), + ]), + AwsEventStreamCodec, + )); + + assert_eq!(messages.next().await.unwrap().unwrap(), message(b"first")); + let Some(Err(EventStreamError::Body(body))) = messages.next().await else { + panic!("the body error surfaces"); + }; + assert_eq!( + body_cause::(&body).unwrap().kind(), + io::ErrorKind::ConnectionReset + ); + assert!(messages.next().await.is_none()); +} diff --git a/litellm-rust/crates/framer/tests/chaining.rs b/litellm-rust/crates/framer/tests/chaining.rs index afd24a90704..81884d58ba1 100644 --- a/litellm-rust/crates/framer/tests/chaining.rs +++ b/litellm-rust/crates/framer/tests/chaining.rs @@ -2,28 +2,64 @@ mod support; -use std::io; +use bytes::Bytes; +use futures_util::{StreamExt, TryStreamExt}; +use litellm_framing::{ + EventStreamError, SseError, + aws_event_stream::{AwsEventStreamCodec, Message}, + frames, + sse::{SseCodec, SseEvent}, +}; +use proptest::prelude::*; +use support::{body_cause, cut_at, encode_all, every, input, runtime}; -use futures_util::TryStreamExt; -use litellm_framing::Framer; -use litellm_framing::aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer}; -use litellm_framing::sse::SseFramer; +fn delta(data: &str) -> SseEvent { + SseEvent { + event: Some("delta".into()), + data: data.into(), + id: Some("7".into()), + retry: None, + } +} -use support::encode; +fn envelopes(payloads: Vec) -> Vec { + encode_all(AwsEventStreamCodec, payloads.into_iter().map(Message::new)) +} + +proptest! { + #[test] + fn an_sse_event_cut_anywhere_across_envelopes_is_reassembled(cut in 0_usize..64, chunk in 1_usize..8) { + let sse = encode_all(SseCodec::default(), [delta("hello")]); + let wire = envelopes(cut_at(&sse, [cut.min(sse.len())])); + let events = runtime().block_on(async { + let payloads = frames(input(every(&wire, chunk)), AwsEventStreamCodec) + .map_ok(|message| message.payload().clone()); + frames(payloads, SseCodec::default()).try_collect::>().await + }) + .unwrap(); + prop_assert_eq!(events, vec![delta("hello")]); + } +} #[tokio::test] -async fn hosting_payloads_feed_the_same_sse_framer_across_envelope_boundaries() { - let bytes = [encode(b"event: delta\ndata: hel"), encode(b"lo\nid: 7\n\n")].concat(); - let envelopes = AwsEventStreamFramer.frame(futures_util::stream::iter( - bytes.chunks(3).map(Ok::<_, io::Error>), +async fn a_truncated_envelope_after_an_sse_event_keeps_the_event_and_its_cause() { + let complete = encode_all(SseCodec::default(), [delta("complete")]); + let incomplete = encode_all(SseCodec::default(), [delta("incomplete")]); + let wire = envelopes(vec![complete.into(), incomplete.into()]); + let payloads = frames( + input(every(&wire[..wire.len() - 1], 3)), + AwsEventStreamCodec, + ) + .map_ok(|message| message.payload().clone()); + let mut events = Box::pin(frames(payloads, SseCodec::default())); + + assert_eq!(events.next().await.unwrap().unwrap(), delta("complete")); + let Some(Err(SseError::Body(body))) = events.next().await else { + panic!("the envelope error surfaces through the SSE layer"); + }; + assert!(matches!( + body_cause::(&body), + Some(EventStreamError::Truncated) )); - let frames = SseFramer - .frame(envelopes.map_ok(|frame: AwsEventStreamFrame| frame.payload)) - .try_collect::>() - .await - .unwrap(); - assert_eq!(frames.len(), 1); - assert_eq!(frames[0].event.as_deref(), Some("delta")); - assert_eq!(frames[0].data.as_deref(), Some("hello")); - assert_eq!(frames[0].id.as_deref(), Some("7")); + assert!(events.next().await.is_none()); } diff --git a/litellm-rust/crates/framer/tests/sse.rs b/litellm-rust/crates/framer/tests/sse.rs index 66339dfbfd2..2fa064653a6 100644 --- a/litellm-rust/crates/framer/tests/sse.rs +++ b/litellm-rust/crates/framer/tests/sse.rs @@ -1,67 +1,169 @@ #![cfg(feature = "sse")] +mod support; + use std::io; -use futures_util::{StreamExt, TryStreamExt}; -use litellm_framing::sse::{SseFrame, SseFramer}; -use litellm_framing::{Error, Framer}; +use bytes::Bytes; +use futures_util::{StreamExt, TryStreamExt, stream}; +use litellm_framing::{ + SseError, frames, + sse::{SseCodec, SseEvent}, +}; +use proptest::prelude::*; use rstest::rstest; +use support::{body_cause, cut_at, encode_all, every, input, runtime}; -async fn collect_sse(chunks: &[&[u8]]) -> Result, Error> { - SseFramer - .frame(futures_util::stream::iter( - chunks.iter().copied().map(Ok::<_, io::Error>), - )) +async fn collect(pieces: Vec) -> Result, SseError> { + frames(input(pieces), SseCodec::default()) .try_collect() .await } +fn event(name: Option<&str>, data: &str) -> SseEvent { + SseEvent { + event: name.map(str::to_owned), + data: data.to_owned(), + id: None, + retry: None, + } +} + +fn sse_event() -> impl Strategy { + ( + proptest::option::of("[^\r\n\0]{0,8}"), + "[^\r\0]{0,16}", + proptest::option::of("[^\r\n\0]{0,8}"), + proptest::option::of(any::()), + ) + .prop_map(|(event, data, id, retry)| SseEvent { + event, + data, + id, + retry, + }) +} + +fn terminators() -> impl Strategy { + prop_oneof![Just(&b"\n"[..]), Just(&b"\r\n"[..]), Just(&b"\r"[..])] +} + +proptest! { + #[test] + fn any_events_survive_a_round_trip_through_any_terminator_and_any_cuts( + events in proptest::collection::vec(sse_event(), 1..4), + terminator in terminators(), + cuts in proptest::collection::vec(0_usize..256, 0..4), + bom in any::(), + ) { + let lf_wire = encode_all(SseCodec::default(), events.clone()); + let body: Vec = lf_wire + .iter() + .flat_map(|byte| if *byte == b'\n' { terminator.to_vec() } else { vec![*byte] }) + .collect(); + let wire = if bom { [&b"\xEF\xBB\xBF"[..], &body].concat() } else { body }; + let decoded = runtime().block_on(collect(cut_at(&wire, cuts))).unwrap(); + prop_assert_eq!(decoded, events); + } +} + #[rstest] -#[case( - &[&b":ping\r\nevent: delta\r\nid: 7\r\nretry: 10\r\ndata: \xe2"[..], &b"\x82"[..], &b"\xac\r"[..], &b"\ndata: next\r\n\r"[..], &b"\ndata: [DONE]\n\n"[..]], - vec![ - SseFrame { - event: Some("delta".into()), - data: Some("€\nnext".into()), - id: Some("7".into()), - retry: Some(10), - }, - SseFrame { - event: None, - data: Some("[DONE]".into()), - id: None, - retry: None, - }, - ] -)] +#[case::comment(b":ping\ndata: x\n\n")] +#[case::unknown_field(b"vendor: 1\ndata: x\n\n")] +#[case::field_without_colon(b"garbage\ndata: x\n\n")] +#[case::retry_with_non_digits(b"retry: soon\ndata: x\n\n")] +#[case::retry_with_a_sign(b"retry: +5\ndata: x\n\n")] +#[case::retry_without_a_value(b"retry:\ndata: x\n\n")] +#[case::id_with_nul(b"id: a\0b\ndata: x\n\n")] #[tokio::test] -async fn fragmented_utf8_crlf_and_multiline_data_retain_metadata_and_sentinel( - #[case] chunks: &[&[u8]], - #[case] expected: Vec, +async fn lines_the_spec_ignores_do_not_change_the_event(#[case] wire: &[u8]) { + assert_eq!( + collect(every(wire, 1)).await.unwrap(), + vec![event(None, "x")] + ); +} + +#[rstest] +#[case::no_data_at_all(b"event: ping\nid: 1\n\ndata: x\n\n", vec![event(None, "x")])] +#[case::empty_data_field(b"data:\n\n", vec![event(None, "")])] +#[case::one_leading_space_stripped(b"data: x\n\n", vec![event(None, " x")])] +#[case::multiline_data(b"data: a\ndata: b\ndata:\n\n", vec![event(None, "a\nb\n")])] +#[case::last_event_name_wins(b"event: a\nevent: b\ndata: x\n\n", vec![event(Some("b"), "x")])] +#[case::last_retry_wins(b"retry: 1\nretry: 2\ndata: x\n\n", vec![SseEvent { retry: Some(2), ..event(None, "x") }])] +#[case::split_utf8_across_lines_is_not_joined(b"data: \xe2\x82\xac\ndata: \xe2\x82\xac\n\n", vec![event(None, "€\n€")])] +#[tokio::test] +async fn dispatch_follows_the_data_buffer(#[case] wire: &[u8], #[case] expected: Vec) { + assert_eq!(collect(every(wire, 1)).await.unwrap(), expected); +} + +#[rstest] +#[case::unterminated_single(b"data: partial\n", vec![])] +#[case::unterminated_tail_after_complete(b"data: complete\n\ndata: unfinished\n", vec![event(None, "complete")])] +#[case::lone_cr_terminates_at_eof(b"data: x\r\r", vec![event(None, "x")])] +#[case::lone_cr_line_then_eof(b"data: x\r", vec![])] +#[tokio::test] +async fn eof_dispatches_only_terminated_events( + #[case] wire: &[u8], + #[case] expected: Vec, ) { - assert_eq!(collect_sse(chunks).await.unwrap(), expected); + assert_eq!( + collect(vec![Bytes::copy_from_slice(wire)]).await.unwrap(), + expected + ); +} + +#[rstest] +#[case::inside_the_first_line(vec![&b"data: a\r"[..], &b"\ndata: b\r\n\r\n"[..]])] +#[case::inside_the_blank_line(vec![&b"data: a\r\ndata: b\r\n\r"[..], &b"\n"[..]])] +#[tokio::test] +async fn a_crlf_split_across_chunks_is_one_terminator(#[case] pieces: Vec<&[u8]>) { + let pieces = pieces.into_iter().map(Bytes::copy_from_slice).collect(); + assert_eq!(collect(pieces).await.unwrap(), vec![event(None, "a\nb")]); } #[tokio::test] -async fn eof_does_not_dispatch_an_unterminated_frame() { - assert!(collect_sse(&[b"data: partial\n"]).await.unwrap().is_empty()); +async fn a_bom_is_stripped_only_at_the_start_of_the_stream() { + let wire = b"\xEF\xBB\xBFdata: a\n\n\xEF\xBB\xBFdata: b\ndata: c\n\n"; + let decoded = collect(every(wire, 2)).await.unwrap(); + assert_eq!(decoded, vec![event(None, "a"), event(None, "c")]); +} + +#[tokio::test] +async fn invalid_utf8_in_a_field_fails_after_earlier_events_and_terminates() { + let mut events = Box::pin(frames( + input(every(b"data: ok\n\ndata: \xff\n\n", 3)), + SseCodec::default(), + )); + + assert_eq!(events.next().await.unwrap().unwrap(), event(None, "ok")); + assert!(matches!( + events.next().await, + Some(Err(SseError::InvalidUtf8(_))) + )); + assert!(events.next().await.is_none()); } #[rstest] #[case(io::ErrorKind::ConnectionReset)] #[case(io::ErrorKind::UnexpectedEof)] #[tokio::test] -async fn framing_errors_terminate_and_preserve_input_error_causes(#[case] kind: io::ErrorKind) { - let mut frames = Box::pin(SseFramer.frame(futures_util::stream::iter([ - Err(io::Error::new(kind, "reset")), - Ok(&b"data: later\n\n"[..]), - ]))); - let error = frames.next().await.unwrap().unwrap_err(); - assert!(matches!( - error, - Error::Sse(sse_stream::Error::Body(ref cause)) - if cause.downcast_ref::().unwrap().kind() == kind +async fn a_body_error_keeps_earlier_events_and_its_cause_then_terminates( + #[case] kind: io::ErrorKind, +) { + let mut events = Box::pin(frames( + stream::iter([ + Ok(&b"data: first\n\ndata: partial"[..]), + Err(io::Error::new(kind, "reset")), + Ok(&b"\n\n"[..]), + ]), + SseCodec::default(), )); - assert!(frames.next().await.is_none()); - assert!(frames.next().await.is_none()); + + assert_eq!(events.next().await.unwrap().unwrap(), event(None, "first")); + let Some(Err(SseError::Body(body))) = events.next().await else { + panic!("the body error surfaces"); + }; + assert_eq!(body_cause::(&body).unwrap().kind(), kind); + assert!(events.next().await.is_none()); + assert!(events.next().await.is_none()); } diff --git a/litellm-rust/crates/framer/tests/support/mod.rs b/litellm-rust/crates/framer/tests/support/mod.rs index 9db305af073..9ff67aef149 100644 --- a/litellm-rust/crates/framer/tests/support/mod.rs +++ b/litellm-rust/crates/framer/tests/support/mod.rs @@ -1,15 +1,57 @@ -use aws_smithy_eventstream::frame::write_message_to; -use aws_smithy_types::event_stream::{Header, HeaderValue, Message}; -use bytes::Bytes; +#![allow(dead_code)] -pub fn encode(payload: &'static [u8]) -> Vec { - let message = Message::new(Bytes::from_static(payload)) - .add_header(Header::new( - ":event-type", - HeaderValue::String("payload".into()), - )) - .add_header(Header::new("sequence", HeaderValue::Int32(7))); - let mut bytes = Vec::new(); - write_message_to(&message, &mut bytes).unwrap(); - bytes +use std::{error::Error, io}; + +use bytes::{Bytes, BytesMut}; +use futures_util::{Stream, stream}; +use tokio_util::codec::Encoder; + +pub fn encode_all(mut codec: C, items: impl IntoIterator) -> Vec +where + C: Encoder, + C::Error: std::fmt::Debug, +{ + let mut wire = BytesMut::new(); + for item in items { + codec.encode(item, &mut wire).unwrap(); + } + wire.to_vec() +} + +pub fn cut_at(bytes: &[u8], offsets: impl IntoIterator) -> Vec { + let mut sorted: Vec = offsets + .into_iter() + .filter(|offset| *offset <= bytes.len()) + .collect(); + sorted.sort_unstable(); + sorted.dedup(); + let bounds = std::iter::once(0) + .chain(sorted) + .chain(std::iter::once(bytes.len())) + .collect::>(); + bounds + .windows(2) + .map(|pair| Bytes::copy_from_slice(&bytes[pair[0]..pair[1]])) + .collect() +} + +pub fn every(bytes: &[u8], size: usize) -> Vec { + bytes + .chunks(size.max(1)) + .map(Bytes::copy_from_slice) + .collect() +} + +pub fn input(pieces: Vec) -> impl Stream> + Send { + stream::iter(pieces.into_iter().map(Ok)) +} + +pub fn body_cause(body: &io::Error) -> Option<&T> { + body.get_ref()?.downcast_ref::() +} + +pub fn runtime() -> tokio::runtime::Runtime { + tokio::runtime::Builder::new_current_thread() + .build() + .unwrap() } diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs index 35e7d5820b0..3f1b7ed9bcc 100644 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs +++ b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs @@ -2,9 +2,9 @@ use base64::Engine; use bytes::Buf; use futures_util::{Stream, StreamExt}; use litellm_framing::{ - Framer, - aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer}, - sse::{SseFrame, SseFramer}, + aws_event_stream::{AwsEventStreamCodec, Message}, + frames, + sse::{SseCodec, SseEvent}, }; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; @@ -13,8 +13,6 @@ use serde_json::{Map, Value}; pub enum Error { #[error("stream framing failed: {0}")] StreamFraming(String), - #[error("Anthropic SSE frame has no data")] - MissingStreamData, #[error("Anthropic stream event is invalid: {0}")] InvalidStreamEvent(String), #[error("Bedrock event payload is invalid: {0}")] @@ -165,15 +163,14 @@ struct BedrockChunkPayload { bytes: String, } -pub fn decode_anthropic_sse_frame(frame: SseFrame) -> Result { - let data = frame.data.ok_or(Error::MissingStreamData)?; - serde_json::from_str(&data).map_err(|error| Error::InvalidStreamEvent(error.to_string())) +pub fn decode_anthropic_sse_frame(event: SseEvent) -> Result { + serde_json::from_str(&event.data).map_err(|error| Error::InvalidStreamEvent(error.to_string())) } pub fn decode_bedrock_anthropic_frame( - frame: AwsEventStreamFrame, + message: Message, ) -> Result { - let payload: BedrockChunkPayload = serde_json::from_slice(&frame.payload) + let payload: BedrockChunkPayload = serde_json::from_slice(message.payload()) .map_err(|error| Error::InvalidBedrockPayload(error.to_string()))?; let event = base64::engine::general_purpose::STANDARD .decode(payload.bytes) @@ -189,9 +186,8 @@ where B: Buf + Send, E: std::error::Error + Send + Sync + 'static, { - SseFramer.frame(input).map(|frame| { - let frame = frame.map_err(|error| Error::StreamFraming(error.to_string()))?; - decode_anthropic_sse_frame(frame) + frames(input, SseCodec::default()).map(|event| { + decode_anthropic_sse_frame(event.map_err(|error| Error::StreamFraming(error.to_string()))?) }) } @@ -203,9 +199,10 @@ where B: Buf + Send, E: std::error::Error + Send + Sync + 'static, { - AwsEventStreamFramer.frame(input).map(|frame| { - let frame = frame.map_err(|error| Error::StreamFraming(error.to_string()))?; - decode_bedrock_anthropic_frame(frame) + frames(input, AwsEventStreamCodec).map(|message| { + decode_bedrock_anthropic_frame( + message.map_err(|error| Error::StreamFraming(error.to_string()))?, + ) }) } @@ -247,12 +244,10 @@ mod tests { #[test] fn decodes_citations_delta_events() { - let event = decode_anthropic_sse_frame(SseFrame { + let event = decode_anthropic_sse_frame(SseEvent { event: Some("content_block_delta".into()), - data: Some( - r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"# - .into(), - ), + data: r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"# + .into(), id: None, retry: None, }) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py index 0bd46382fef..5550590d0c0 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py @@ -16,6 +16,7 @@ from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.llms.anthropic.common_utils import ANTHROPIC_ERROR_STATUS_CODE_MAP +from litellm.llms.anthropic.experimental_pass_through.messages.utils import INCOMPLETE_STREAM_ERROR_MESSAGE from litellm.proxy.pass_through_endpoints.success_handler import ( PassThroughEndpointLogging, ) @@ -28,11 +29,6 @@ GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ: Final = PassThroughEndpointLogging() _UPSTREAM_PUMP_TASKS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: stdlib strong-ref set for pump tasks _DETACHED_STREAM_DRAINS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: bounded strong-ref set, detached drains -INCOMPLETE_STREAM_ERROR_MESSAGE: Final = ( - "Provider stream ended before emitting a message_stop event; " - "the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated." -) - def _is_message_stop_chunk(chunk: object) -> bool: if isinstance(chunk, dict): diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/utils.py b/litellm/llms/anthropic/experimental_pass_through/messages/utils.py index 89105c00428..fe8ac2cd7a2 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/utils.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/utils.py @@ -15,6 +15,12 @@ if TYPE_CHECKING: from litellm.exceptions import ContentPolicyViolationError +INCOMPLETE_STREAM_ERROR_MESSAGE: Final = ( + "Provider stream ended before emitting a message_stop event; " + "the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated." +) + + def get_safeguard_refusal_stop_details(response: object) -> Mapping[str, Any] | None: """ Return the ``stop_details`` of an Anthropic Messages response refused by a diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py index f753e87fee3..59ccde872fc 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py @@ -2,20 +2,25 @@ ## Translates OpenAI call to Anthropic `/v1/messages` format import asyncio import json -import traceback from collections import deque from collections.abc import AsyncIterator, Iterator, Mapping from typing import TYPE_CHECKING, Any, Final +from pydantic import BaseModel, ConfigDict, field_validator + from litellm import verbose_logger +from litellm._logging import redact_internal_details_from_client_message from litellm._uuid import uuid +from litellm.exceptions import MidStreamFallbackError from litellm.litellm_core_utils.prompt_templates.common_utils import ( encrypted_reasoning_signature, ) from litellm.llms.anthropic.experimental_pass_through.messages.utils import ( + INCOMPLETE_STREAM_ERROR_MESSAGE, refusal_stop_details, responses_output_refusal_text, ) +from litellm.responses.streaming_iterator import stream_error_status_and_message from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage from .transformation import ( @@ -27,6 +32,72 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject +class _UpstreamFailure(BaseModel): + model_config = ConfigDict(frozen=True) + + status_code: int | None = None + message: str | None = None + + @field_validator("status_code", mode="before") + @classmethod + def http_error_status_or_none(cls, value: object) -> int | None: + candidate: Final = ( + value + if isinstance(value, int) and not isinstance(value, bool) + else int(value) + if isinstance(value, str) and value.isdecimal() + else None + ) + return candidate if candidate is not None and 400 <= candidate <= 599 else None + + @field_validator("message", mode="before") + @classmethod + def str_or_none(cls, value: object) -> str | None: + return value if isinstance(value, str) else None + + +class _FailedResponse(BaseModel): + model_config = ConfigDict(frozen=True, from_attributes=True) + + error: object | None = None + + +class _FailedResponseEvent(BaseModel): + model_config = ConfigDict(frozen=True, from_attributes=True) + + response: _FailedResponse | None = None + + +def _original_failure(exception: Exception) -> Exception: + failure = exception # rebind-ok: walks the MidStreamFallbackError chain down to the provider failure + while isinstance(failure, MidStreamFallbackError) and failure.original_exception is not None: + failure = failure.original_exception + return failure + + +def _failure_status_and_message(exception: Exception) -> tuple[int, str]: + original: Final = _original_failure(exception) + failure: Final = _UpstreamFailure.model_validate( + {"status_code": getattr(original, "status_code", None), "message": getattr(original, "message", None)} + ) + status_code: Final = failure.status_code if failure.status_code is not None else 500 + message: Final = failure.message or str(original) or INCOMPLETE_STREAM_ERROR_MESSAGE + return status_code, message + + +def _anthropic_error_chunk(status_code: int, message: str) -> dict[str, object]: + from litellm.anthropic_interface.exceptions.exception_mapping_utils import ( + AnthropicExceptionMapping, + ) + + return dict( + AnthropicExceptionMapping.transform_to_anthropic_error( + status_code=status_code, + raw_message=redact_internal_details_from_client_message(message), + ) + ) + + class AnthropicResponsesStreamWrapper: """ Wraps a Responses API streaming iterator and re-emits events in Anthropic SSE format. @@ -40,6 +111,7 @@ class AnthropicResponsesStreamWrapper: response.function_call_arguments.delta -> content_block_delta (input_json_delta) response.output_item.done -> content_block_delta (signature_delta) + content_block_stop response.completed -> message_delta + message_stop + response.failed -> error (the stream ends without message_stop) """ def __init__( @@ -60,6 +132,7 @@ class AnthropicResponsesStreamWrapper: self._pending_tool_ids: dict[str, str] = {} # item_id -> call_id / name accumulator self._sent_message_start = False self._sent_message_stop = False + self._stream_failed = False self._chunk_queue: deque[dict[str, object]] = deque() self._refusal_text: str = "" self._sync_responses_iterator: Iterator[object] | None = None @@ -293,10 +366,23 @@ class AnthropicResponsesStreamWrapper: ) return + if event_type == "response.failed": + failed: Final = _FailedResponseEvent.model_validate(event) + status_code, message = stream_error_status_and_message( + failed.response.error if failed.response is not None else None + ) + verbose_logger.error( + "AnthropicResponsesStreamWrapper: upstream Responses stream for %s failed (%s): %s", + self.model, + status_code, + message, + ) + self._fail_stream(status_code, message) + return + # ---- response completed -> message_delta + message_stop ---- if event_type in ( "response.completed", - "response.failed", "response.incomplete", ): response_obj: Final = getattr(event, "response", None) or ( @@ -350,21 +436,24 @@ class AnthropicResponsesStreamWrapper: self._sent_message_stop = True return + def _fail_stream(self, status_code: int, message: str) -> None: + self._stream_failed = True + self._chunk_queue.append(_anthropic_error_chunk(status_code, message)) + def __aiter__(self) -> "AnthropicResponsesStreamWrapper": return self async def __anext__(self) -> dict[str, object]: - # Return any queued chunks first if self._chunk_queue: return self._chunk_queue.popleft() + if self._stream_failed: + raise StopAsyncIteration - # Emit message_start if not yet done (fallback if response.created wasn't fired) if not self._sent_message_start: self._sent_message_start = True self._chunk_queue.append(self._make_message_start()) return self._chunk_queue.popleft() - # Consume the upstream stream try: if hasattr(self.responses_stream, "__aiter__"): async for event in self.responses_stream: @@ -382,10 +471,19 @@ class AnthropicResponsesStreamWrapper: return self._chunk_queue.popleft() except StopAsyncIteration: pass - except Exception as e: - verbose_logger.error("AnthropicResponsesStreamWrapper error: %s\n%s", e, traceback.format_exc()) + except Exception as e: # noqa: BLE001 # every upstream failure becomes a client error event + verbose_logger.exception( + "AnthropicResponsesStreamWrapper: upstream Responses stream for %s failed", self.model + ) + self._fail_stream(*_failure_status_and_message(e)) + + if not self._chunk_queue and not self._sent_message_stop and not self._stream_failed: + verbose_logger.error( + "AnthropicResponsesStreamWrapper: upstream Responses stream for %s ended without a terminal event", + self.model, + ) + self._fail_stream(500, INCOMPLETE_STREAM_ERROR_MESSAGE) - # Drain any remaining queued chunks if self._chunk_queue: return self._chunk_queue.popleft() diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 7859c678c07..618b200a14c 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -4618,6 +4618,13 @@ class GoogleSSOHandler: return result or {} +def _raise_if_sso_debug_disabled() -> None: + """The debug routes run the browser-redirect SSO flow, so they cannot carry a + bearer credential; an explicit opt-in flag is the only way to gate them.""" + if get_secret_bool("ENABLE_SSO_DEBUG") is not True: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Not Found") + + @router.get("/sso/debug/login", tags=["experimental"], include_in_schema=False) async def debug_sso_login(request: Request): """ @@ -4625,6 +4632,8 @@ async def debug_sso_login(request: Request): PROXY_BASE_URL should be the your deployed proxy endpoint, e.g. PROXY_BASE_URL="https://litellm-production-7002.up.railway.app/" Example: """ + _raise_if_sso_debug_disabled() + from litellm.proxy.proxy_server import premium_user microsoft_client_id: Final = os.getenv("MICROSOFT_CLIENT_ID", None) @@ -4670,6 +4679,8 @@ async def debug_sso_callback(request: Request): """ Returns the OpenID object returned by the SSO provider """ + _raise_if_sso_debug_disabled() + import json from fastapi.responses import HTMLResponse diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 70f2a7db6da..fdc702af005 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -230,6 +230,11 @@ def _status_code_for_error_fields(error_type: str | None, error_code: str | None return next((status for status in map(_status_code_for_error_field, fields) if status is not None), 500) +def stream_error_status_and_message(error_obj: object) -> tuple[int, str]: + message, error_type, error_code = _error_event_fields(error_obj) + return _status_code_for_error_fields(error_type, error_code), message + + def _map_stream_error_to_exception(error_obj: object, model: str, custom_llm_provider: str) -> Exception: from litellm.llms.base_llm.chat.transformation import BaseLLMException diff --git a/tests/_vcr_conftest_common.py b/tests/_vcr_conftest_common.py index ab046674eb6..3adc671021b 100644 --- a/tests/_vcr_conftest_common.py +++ b/tests/_vcr_conftest_common.py @@ -52,7 +52,7 @@ from tests._vcr_redis_persister import ( # network call entirely, so skip tests record nothing (NOOP) and passing tests # stop carrying a volatile github episode. This matches the established idiom in # the unit-test suite, which sets the same flag (see e.g. -# tests/test_litellm/test_cost_calculator.py). ``setdefault`` so an explicit +# tests/unit/test_cost_calculator.py). ``setdefault`` so an explicit # override still wins. os.environ.setdefault("LITELLM_LOCAL_MODEL_COST_MAP", "True") diff --git a/tests/code_coverage_tests/code_qa_check_tests.py b/tests/code_coverage_tests/code_qa_check_tests.py index 025f836511c..6c620a02522 100644 --- a/tests/code_coverage_tests/code_qa_check_tests.py +++ b/tests/code_coverage_tests/code_qa_check_tests.py @@ -13,15 +13,16 @@ def check_for_litellm_module_deletion(base_dir): del sys.modules[module] """ problematic_files = [] - test_dir = os.path.join(base_dir, "test_litellm") + candidate_dirs = [os.path.join(base_dir, name) for name in ("test_litellm", "unit")] + test_dirs = [test_dir for test_dir in candidate_dirs if os.path.exists(test_dir)] - if not os.path.exists(test_dir): - print(f"Warning: Directory {test_dir} does not exist.") + if not test_dirs: + print(f"Warning: None of {candidate_dirs} exist.") return [] - print(f"Checking directory: {test_dir}") + print(f"Checking directories: {test_dirs}") - for root, _, files in os.walk(test_dir): + for root, _, files in (entry for test_dir in test_dirs for entry in os.walk(test_dir)): for file in files: if file.endswith(".py"): file_path = os.path.join(root, file) @@ -173,7 +174,7 @@ def main(): f"This can cause import issues and test failures. Files: {problematic_files}" ) else: - print("✓ No litellm module deletion patterns found in test_litellm directory.") + print("✓ No litellm module deletion patterns found in tests/test_litellm or tests/unit.") if __name__ == "__main__": diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index 7332a533872..06e5b020836 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -31,7 +31,7 @@ def get_all_functions_called_in_tests(base_dir): specifically in files containing the word 'router'. """ called_functions = set() - test_dirs = ["local_testing", "router_unit_tests", "test_litellm"] + test_dirs = ["local_testing", "router_unit_tests", "test_litellm", "unit"] for test_dir in test_dirs: dir_path = os.path.join(base_dir, test_dir) diff --git a/tests/llm_translation/test_bedrock_gpt_oss.py b/tests/llm_translation/test_bedrock_gpt_oss.py index 4af81ee81f7..b264c16601f 100644 --- a/tests/llm_translation/test_bedrock_gpt_oss.py +++ b/tests/llm_translation/test_bedrock_gpt_oss.py @@ -22,7 +22,7 @@ class TestBedrockGPTOSS(BaseLLMChatTest): """Bedrock GPT-OSS intermittently emits truncated toolUse.input deltas on the live endpoint, which makes the inherited live integration test flaky. The accumulation side is covered deterministically by - tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py::test_transform_tool_calls_index; + tests/unit/llms/bedrock/chat/test_invoke_handler.py::test_transform_tool_calls_index; the GPT-OSS-specific request-body transformation is covered by test_function_calling_request_body_gpt_oss below. """ diff --git a/tests/llm_translation/test_skills_api.py b/tests/llm_translation/test_skills_api.py index aeab5f0da3e..d21e7376ea7 100644 --- a/tests/llm_translation/test_skills_api.py +++ b/tests/llm_translation/test_skills_api.py @@ -277,4 +277,4 @@ class BaseSkillsAPITest(ABC): # # Transformation logic (URL construction, headers, request/response parsing) is # covered by unit tests in: -# tests/test_litellm/test_anthropic_skills_transformation.py +# tests/unit/test_anthropic_skills_transformation.py diff --git a/tests/local_testing/test_function_calling.py b/tests/local_testing/test_function_calling.py index 3a5e2209f1e..2d79f8a6af6 100644 --- a/tests/local_testing/test_function_calling.py +++ b/tests/local_testing/test_function_calling.py @@ -324,7 +324,7 @@ def test_parallel_function_call_anthropic_error_msg(model, messages): Anthropic (and Bedrock Invoke via ``AnthropicConfig.transform_request``) inject a dummy tool so CLIs work with ``modify_params`` left off. Bedrock Converse's no-raise behavior is covered offline in - ``tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py`` + ``tests/unit/llms/bedrock/chat/test_converse_transformation.py`` (see #24158, #27138), which needs no live credentials. """ # Force modify_params off as a clean baseline: it exercises the Anthropic diff --git a/tests/local_testing/test_handler_gc_does_not_close_client.py b/tests/local_testing/test_handler_gc_does_not_close_client.py index 1a6ab1b1827..63c5694dd89 100644 --- a/tests/local_testing/test_handler_gc_does_not_close_client.py +++ b/tests/local_testing/test_handler_gc_does_not_close_client.py @@ -17,7 +17,7 @@ body can still arrive, released once the caller is done with the response. Nothing here re-tests the shapes ``_handler_may_close_client`` covers -- a borrowed ``handler.client``, a caller-supplied client, an evicted-but-held -client. Those are pinned in ``tests/test_litellm/llms/custom_httpx/ +client. Those are pinned in ``tests/unit/llms/custom_httpx/ test_http_handler.py``. What is uncovered there is the in-flight response, so no test here may keep the client in a local: that inflates the very refcount under test, and the test then passes on a broken handler. They hold weak references diff --git a/tests/local_testing/test_sagemaker_nova_integration.py b/tests/local_testing/test_sagemaker_nova_integration.py index beeb1fa2db3..95f28fe9892 100644 --- a/tests/local_testing/test_sagemaker_nova_integration.py +++ b/tests/local_testing/test_sagemaker_nova_integration.py @@ -4,7 +4,7 @@ Integration tests for SageMaker Nova provider. These tests require a live SageMaker Nova endpoint and AWS credentials. They are skipped by default — run manually with: - pytest tests/test_litellm/llms/sagemaker/test_sagemaker_nova_integration.py -v --no-header -rN + pytest tests/local_testing/test_sagemaker_nova_integration.py -v --no-header -rN Prerequisites: export AWS_PROFILE= # or set AWS_ACCESS_KEY_ID / AWS_SECRET_ACCESS_KEY @@ -251,7 +251,7 @@ class TestSagemakerNova2LiteIntegration: Run with: export SAGEMAKER_NOVA2_LITE_ENDPOINT= - pytest tests/test_litellm/llms/sagemaker/test_sagemaker_nova_integration.py::TestSagemakerNova2LiteIntegration -v + pytest tests/local_testing/test_sagemaker_nova_integration.py::TestSagemakerNova2LiteIntegration -v """ def test_should_accept_reasoning_effort_low(self): diff --git a/tests/search_tests/test_bing_grounding_search.py b/tests/search_tests/test_bing_grounding_search.py index 3d1737477a1..f532158e462 100644 --- a/tests/search_tests/test_bing_grounding_search.py +++ b/tests/search_tests/test_bing_grounding_search.py @@ -85,7 +85,7 @@ class TestBingGroundingSearch(BaseSearchTest): class TestBingGroundingSearchTransformation: """ Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked. - Transformation details are unit-tested in tests/test_litellm/llms/azure/search/. + Transformation details are unit-tested in tests/unit/llms/azure/search/. """ @pytest.fixture(autouse=True) diff --git a/tests/search_tests/test_nimble_search.py b/tests/search_tests/test_nimble_search.py index df432f8ae84..3426fc712f4 100644 --- a/tests/search_tests/test_nimble_search.py +++ b/tests/search_tests/test_nimble_search.py @@ -58,7 +58,7 @@ class TestNimbleSearch(BaseSearchTest): class TestNimbleSearchTransformation: """ Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked. - Transformation details are unit-tested in tests/test_litellm/llms/nimble/search/. + Transformation details are unit-tested in tests/unit/llms/nimble/search/. """ @pytest.fixture(autouse=True) diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py deleted file mode 100644 index 0b2bfe9d266..00000000000 --- a/tests/test_litellm/batches/test_batch_utils.py +++ /dev/null @@ -1,387 +0,0 @@ -import json - -import pytest - -import litellm -import litellm.batches.batch_utils as bu -from litellm.types.llms.openai import Batch - -GROUNDED_USAGE_METADATA = { - "promptTokenCount": 19, - "candidatesTokenCount": 59, - "thoughtsTokenCount": 406, - "toolUsePromptTokenCount": 73, - "totalTokenCount": 557, - "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 19}], - "candidatesTokensDetails": [{"modality": "TEXT", "tokenCount": 59}], - "toolUsePromptTokensDetails": [{"modality": "TEXT", "tokenCount": 73}], - "trafficType": "ON_DEMAND", -} -PASSTHROUGH_OUTPUT_URI = ( - "gs://litellm-bucket/litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/u/" - "predictions.jsonl" -) -UNGROUNDED_USAGE_METADATA = { - "promptTokenCount": 20, - "candidatesTokenCount": 48, - "thoughtsTokenCount": 195, - "toolUsePromptTokenCount": 73, - "totalTokenCount": 336, - "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 20}], - "trafficType": "ON_DEMAND", -} - - -def _batch(output_file_id: str) -> Batch: - return Batch( - id="b", - completion_window="24h", - created_at=1, - endpoint="/v1/chat/completions", - input_file_id="f", - object="batch", - status="completed", - output_file_id=output_file_id, - ) - - -def _vertex_jsonl(rows: list[dict]) -> bytes: - return "\n".join(json.dumps(row) for row in rows).encode() - - -def _vertex_openai_row(custom_id: str, model: str, prompt_tokens: int, completion_tokens: int) -> dict: - return { - "id": f"batch_req_{custom_id}", - "custom_id": custom_id, - "response": { - "status_code": 200, - "request_id": custom_id, - "body": { - "id": f"chatcmpl-{custom_id}", - "object": "chat.completion", - "model": model, - "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], - "usage": { - "prompt_tokens": prompt_tokens, - "completion_tokens": completion_tokens, - "total_tokens": prompt_tokens + completion_tokens, - }, - }, - }, - "error": None, - } - - -def _native_vertex_row(usage_metadata: dict, *, grounded: bool, model_version: str | None = "gemini-2.5-flash"): - candidate = {"content": {"role": "model", "parts": [{"text": "ok"}]}, "finishReason": "STOP"} - grounding = {"groundingMetadata": {"webSearchQueries": ["q"]}} if grounded else {} - response = {"candidates": [{**candidate, **grounding}], "usageMetadata": usage_metadata} - return { - "request": {"contents": [{"role": "user", "parts": [{"text": "q"}]}], "tools": [{"googleSearch": {}}]}, - "status": "", - "response": {**response, **({"modelVersion": model_version} if model_version else {})}, - "processed_time": "2026-09-23T19:02:00.000+00:00", - } - - -def _capture_cost_calls(monkeypatch, prompt_cost=0.5, completion_cost=0.25) -> list: - import litellm.cost_calculator as cc - - calls: list = [] - - def _calc(**kw): - calls.append(kw) - return (prompt_cost, completion_cost) - - monkeypatch.setattr(cc, "batch_cost_calculator", _calc) - return calls - - -def test_vertex_native_cost_bills_embedding_rows(monkeypatch): - monkeypatch.setitem(litellm.model_cost, "vertex_ai/gemini-embedding-2", {"input_cost_per_token_batches": 1e-7}) - rows = [ - { - "key": "id_1", - "status": "", - "request": {"content": {"parts": [{"text": "hello world"}]}}, - "response": {"embedding": {"values": [0.1, 0.2]}, "usageMetadata": {"promptTokenCount": 2}}, - }, - { - "key": "id_2", - "status": "", - "request": {"content": {"parts": [{"text": "hello"}]}}, - "response": {"embedding": {"values": [0.3]}, "tokenCount": "3"}, - }, - {"key": "id_3", "status": "INVALID_ARGUMENT", "request": {"content": {"parts": [{"text": ""}]}}}, - ] - - result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-embedding-2") - - assert (result.successful_requests, result.failed_requests) == (2, 1) - assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (5, 0, 5) - assert result.cost == pytest.approx(5 * 1e-7) - assert result.models == ["gemini-embedding-2"] - - -@pytest.mark.asyncio -async def test_native_vertex_rows_route_to_vertex_cost_path_without_flag(monkeypatch): - monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False) - monkeypatch.setattr( - bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run") - ) - calls = _capture_cost_calls(monkeypatch) - rows = [ - _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True), - _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False), - ] - - result = await bu.calculate_batch_cost_and_usage( - file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash" - ) - - assert result.cost == pytest.approx(1.5) - assert (result.successful_requests, result.failed_requests) == (2, 0) - assert result.models == ["gemini-2.5-flash"] - assert {(call["model"], call["custom_llm_provider"]) for call in calls} == {("gemini-2.5-flash", "vertex_ai")} - - -@pytest.mark.asyncio -async def test_openai_shaped_vertex_rows_keep_the_generic_path_without_flag(monkeypatch): - monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False) - monkeypatch.setattr( - bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run") - ) - _capture_cost_calls(monkeypatch) - rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)] - - result = await bu.calculate_batch_cost_and_usage( - file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash" - ) - - assert result.successful_requests == 1 - - -@pytest.mark.asyncio -async def test_native_vertex_rows_on_another_provider_keep_the_generic_path(monkeypatch): - monkeypatch.setattr( - bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run") - ) - _capture_cost_calls(monkeypatch) - - result = await bu.calculate_batch_cost_and_usage( - file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)], - custom_llm_provider="openai", - ) - - assert result.successful_requests == 0 - - -@pytest.mark.asyncio -async def test_handle_completed_batch_routes_native_rows_without_flag(monkeypatch): - monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False) - raw_rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)] - - async def fake_fetch(batch, custom_llm_provider, litellm_params=None): - return _vertex_jsonl(raw_rows) - - monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch) - monkeypatch.setattr( - bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run") - ) - calls = _capture_cost_calls(monkeypatch, prompt_cost=0.7, completion_cost=0.3) - deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6} - - result = await bu._handle_completed_batch( - _batch(PASSTHROUGH_OUTPUT_URI), - custom_llm_provider="vertex_ai", - model_name="gemini-2.5-flash", - model_info=deployment_model_info, - ) - - assert result.cost == pytest.approx(1.0) - assert result.usage.total_tokens == 557 - assert [call["model_info"] for call in calls] == [deployment_model_info] - - -def test_native_vertex_usage_is_billed_like_the_online_path(monkeypatch): - calls = _capture_cost_calls(monkeypatch) - grounded = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True) - ungrounded = _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False) - - result = bu.calculate_vertex_ai_batch_cost_and_usage([grounded, ungrounded], "gemini-2.5-flash") - - grounded_usage, ungrounded_usage = (call["usage"] for call in calls) - assert grounded_usage.prompt_tokens == 19 - assert grounded_usage.completion_tokens == 59 + 406 - assert grounded_usage.completion_tokens_details.reasoning_tokens == 406 - assert ungrounded_usage.prompt_tokens == 20 + 73 - assert ungrounded_usage.completion_tokens == 48 + 195 - assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == ( - 19 + 93, - 465 + 243, - 557 + 336, - ) - - -def test_native_vertex_rows_are_priced_by_model_version_without_a_model_name(monkeypatch): - calls = _capture_cost_calls(monkeypatch) - rows = [ - _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"), - _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-pro"), - _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None), - ] - - result = bu.calculate_vertex_ai_batch_cost_and_usage(rows) - - assert [call["model"] for call in calls] == ["gemini-2.5-flash", "gemini-2.5-pro"] - assert result.models == ["gemini-2.5-flash", "gemini-2.5-pro"] - assert result.cost == pytest.approx(1.5) - assert result.successful_requests == 3 - assert result.usage.total_tokens == 557 + 336 + 336 - - -def test_native_vertex_rows_without_usage_metadata_count_as_failed(monkeypatch): - _capture_cost_calls(monkeypatch) - rows = [ - {"request": {"contents": []}, "status": "Error: bad request", "processed_time": "t"}, - {"request": {"contents": []}, "response": {"candidates": []}}, - _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True), - ] - - result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash") - - assert (result.successful_requests, result.failed_requests) == (1, 2) - assert result.usage.total_tokens == 557 - - -def test_native_vertex_batch_whose_rows_all_failed_still_names_the_deployment_model(monkeypatch): - calls = _capture_cost_calls(monkeypatch) - rows = [{"request": {"contents": []}, "status": "Error: quota exceeded", "processed_time": "t"}] * 2 - - result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash") - - assert result.models == ["gemini-2.5-flash"] - assert (result.successful_requests, result.failed_requests, result.cost) == (0, 2, 0.0) - assert calls == [] - - -def test_native_vertex_rows_are_priced_with_the_deployment_model_info(monkeypatch): - calls = _capture_cost_calls(monkeypatch) - deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6} - - bu.calculate_vertex_ai_batch_cost_and_usage( - [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)], - "gemini-2.5-flash", - model_info=deployment_model_info, - ) - - assert [call["model_info"] for call in calls] == [deployment_model_info] - - -@pytest.mark.asyncio -async def test_native_vertex_rows_keep_the_deployment_model_info_through_the_batch_entrypoint(monkeypatch): - calls = _capture_cost_calls(monkeypatch) - deployment_model_info = {"input_cost_per_token_batches": 1e-6} - - await bu.calculate_batch_cost_and_usage( - file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)], - custom_llm_provider="vertex_ai", - model_name="gemini-2.5-flash", - model_info=deployment_model_info, - ) - - assert [call["model_info"] for call in calls] == [deployment_model_info] - - -def test_native_vertex_rows_are_priced_by_the_deployment_model_over_model_version(monkeypatch): - calls = _capture_cost_calls(monkeypatch) - rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-pro")] - - result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash") - - assert [call["model"] for call in calls] == ["gemini-2.5-flash"] - assert result.models == ["gemini-2.5-flash"] - - -def test_native_vertex_rows_that_fail_response_validation_count_as_failed(monkeypatch): - calls = _capture_cost_calls(monkeypatch) - rows = [ - {"request": {"contents": []}, "response": {"candidates": "nope", "usageMetadata": GROUNDED_USAGE_METADATA}}, - _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True), - ] - - result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash") - - assert (result.successful_requests, result.failed_requests) == (1, 1) - assert result.usage.total_tokens == 557 - assert len(calls) == 1 - - -@pytest.mark.parametrize("wildcard_model", ["*", "vertex_ai/*"]) -def test_native_vertex_rows_under_a_wildcard_deployment_are_priced_by_model_version(monkeypatch, wildcard_model): - calls = _capture_cost_calls(monkeypatch) - rows = [ - _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"), - _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None), - ] - - result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, wildcard_model) - - assert [call["model"] for call in calls] == ["gemini-2.5-flash", wildcard_model] - assert result.cost == pytest.approx(1.5) - assert (result.successful_requests, result.failed_requests) == (2, 0) - assert result.usage.total_tokens == 557 + 336 - - -def test_native_vertex_row_without_model_version_under_a_wildcard_deployment_bills_its_explicit_prices(): - deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6} - with_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash") - without_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version=None) - - twin = bu.calculate_vertex_ai_batch_cost_and_usage([with_version], "vertex_ai/*", model_info=deployment_model_info) - both = bu.calculate_vertex_ai_batch_cost_and_usage( - [with_version, without_version], "vertex_ai/*", model_info=deployment_model_info - ) - - assert twin.cost > 0 - assert both.cost == pytest.approx(2 * twin.cost) - assert (both.successful_requests, both.failed_requests) == (2, 0) - - -def test_native_vertex_row_the_cost_map_cannot_price_is_billed_at_zero_and_the_rest_still_bills(monkeypatch): - import litellm.cost_calculator as cc - - def _calc(**kw): - if kw["model"] == "gemini-unpriced": - raise ValueError("no pricing") - return (0.5, 0.25) - - monkeypatch.setattr(cc, "batch_cost_calculator", _calc) - rows = [ - _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-unpriced"), - _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-flash"), - ] - - result = bu.calculate_vertex_ai_batch_cost_and_usage(rows) - - assert result.cost == pytest.approx(0.75) - assert (result.successful_requests, result.failed_requests) == (2, 0) - assert result.usage.total_tokens == 557 + 336 - assert result.models == ["gemini-unpriced", "gemini-2.5-flash"] - - -@pytest.mark.asyncio -async def test_flag_sends_every_vertex_row_down_the_native_path_when_a_model_is_known(monkeypatch): - monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False) - monkeypatch.setattr( - bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run") - ) - calls = _capture_cost_calls(monkeypatch) - rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)] - - result = await bu.calculate_batch_cost_and_usage( - file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash" - ) - - assert calls == [] - assert (result.successful_requests, result.failed_requests) == (0, 1) diff --git a/tests/test_litellm/chat_completions/test_dispatch.py b/tests/test_litellm/chat_completions/test_dispatch.py deleted file mode 100644 index ddb6e827309..00000000000 --- a/tests/test_litellm/chat_completions/test_dispatch.py +++ /dev/null @@ -1,117 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping -from typing import Final - -import pytest - -import litellm -from litellm.chat_completions import dispatch -from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.catalog import Route, RouteRule, Rules -from litellm.rust_bridge.chat_completions.entrypoints import ( - LiteLLMChatCompletionsRequest, - NativeAcompletion, - NativeCompletion, -) -from litellm.rust_bridge.configuration import Rollout -from litellm.types.utils import ModelResponse - -MESSAGES: Final = [{"role": "user", "content": "hi"}] - - -@pytest.mark.asyncio -async def test_public_completion_calls_keep_the_python_result() -> None: - sync_response: Final = litellm.completion(model="openai/test-model", messages=MESSAGES, mock_response="ok") - async_response: Final = await litellm.acompletion(model="openai/test-model", messages=MESSAGES, mock_response="ok") - - assert isinstance(sync_response, ModelResponse) - assert isinstance(async_response, ModelResponse) - assert sync_response.choices[0].message.content == "ok" - assert async_response.choices[0].message.content == "ok" - - -def test_sync_completion_request_projects_public_arguments() -> None: - rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),) - expected: Final = ModelResponse() - - def native( - request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> ModelResponse: - assert request.model == "test-model" - assert request.messages == MESSAGES - assert request.custom_llm_provider == "openai" - assert request.stream is True - return expected - - binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None) - binding.override(native) - response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision - ("test-model", MESSAGES), - {"custom_llm_provider": "openai", "stream": True}, - python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"), - binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), - rules=rules, - ) - - assert response is expected - - -@pytest.mark.asyncio -async def test_async_completion_falls_back_after_native_declines() -> None: - from litellm.rust_bridge.bindings import native_exception_types - - native_types: Final = native_exception_types() - if native_types is None: - pytest.skip("native bridge is unavailable") - declined, _ = native_types - expected: Final = ModelResponse() - rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_OPT_OUT),) - - async def native( - request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> ModelResponse: - raise declined("unsupported") - - async def python(*args: object, **kwargs: object) -> ModelResponse: - return expected - - binding: Final[NativeBinding[NativeAcompletion]] = NativeBinding("acompletion", validate=lambda _: None) - binding.override(native) - response: Final = await dispatch._ADISPATCH.arun( # pyright: ignore[reportPrivateUsage] # test an explicit route decision - ("test-model", MESSAGES), - {}, - python=python, - binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), - rules=rules, - ) - - assert response is expected - - -def test_internal_acompletion_marker_bypasses_native() -> None: - rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),) - expected: Final = ModelResponse() - - def python(*args: object, **kwargs: object) -> ModelResponse: - return expected - - def native( - request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> ModelResponse: - pytest.fail("acompletion's inner completion call must stay on Python") - - binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None) - binding.override(native) - response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision - ("test-model", MESSAGES), - {"custom_llm_provider": "openai", "acompletion": True}, - python=python, - binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), - rules=rules, - ) - - assert response is expected diff --git a/tests/test_litellm/conftest.py b/tests/test_litellm/conftest.py index beca10d5555..f8c7d5273d1 100644 --- a/tests/test_litellm/conftest.py +++ b/tests/test_litellm/conftest.py @@ -14,6 +14,7 @@ from pathlib import Path from types import SimpleNamespace import httpx import pytest +from pytest_socket import _remove_restrictions import asyncio @@ -509,6 +510,14 @@ def setup_and_teardown(): print(f"[conftest] Module teardown complete (worker: {worker_id or 'master'})") +def pytest_collectstart(): + _remove_restrictions() + + +def pytest_runtest_setup(): + _remove_restrictions() + + def pytest_collection_modifyitems(config, items): """ Customize test collection order. diff --git a/tests/test_litellm/integrations/test_helicone.py b/tests/test_litellm/integrations/test_helicone.py index 64960de050a..99cb1380dd7 100644 --- a/tests/test_litellm/integrations/test_helicone.py +++ b/tests/test_litellm/integrations/test_helicone.py @@ -13,7 +13,7 @@ def _claude_mapping(messages, response_obj): def test_claude_mapping_serializes_custom_tool_calls(monkeypatch): """ Stub the anthropic module unconditionally: the SDK may be absent (it lives in the - proxy-runtime extra), and the tests/test_litellm/llms/anthropic test package can + proxy-runtime extra), and the tests/unit/llms/anthropic test package can shadow it on sys.path, so an import probe proves nothing about the real SDK. """ stub = types.ModuleType("anthropic") diff --git a/tests/test_litellm/interactions/test_litellm_responses_bridge.py b/tests/test_litellm/interactions/test_litellm_responses_bridge.py index 8400f2c4840..17e7f9fc4ff 100644 --- a/tests/test_litellm/interactions/test_litellm_responses_bridge.py +++ b/tests/test_litellm/interactions/test_litellm_responses_bridge.py @@ -7,10 +7,6 @@ the litellm_responses bridge provider, which calls litellm.responses() internall import os -from litellm.interactions.litellm_responses_transformation.transformation import ( - LiteLLMResponsesInteractionsConfig, -) -from litellm.types.interactions import Turn from tests.test_litellm.interactions.base_interactions_test import ( BaseInteractionsTest, ) @@ -30,71 +26,3 @@ class TestLiteLLMResponsesBridge(BaseInteractionsTest): def get_api_key(self) -> str: """Return the OpenAI API key from environment.""" return os.getenv("OPENAI_API_KEY", "") - - -class TestBridgeInputTransformation: - """Regression tests for translating Interactions input into Responses API input. - - The bridge used to pass Google content parts through raw ({"type": "text"}), - which the Responses API rejects with a 400, and it dropped the role encoded - in step types and in the legacy "model" turn role. - """ - - def test_step_input_maps_roles_and_content_types(self): - transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input( - [ - {"type": "user_input", "content": [{"type": "text", "text": "I like apples."}]}, - {"type": "model_output", "content": [{"type": "text", "text": "I like oranges."}]}, - {"type": "user_input", "content": [{"type": "text", "text": "What did you say?"}]}, - ] - ) - assert transformed == [ - {"role": "user", "content": [{"type": "input_text", "text": "I like apples."}]}, - {"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]}, - {"role": "user", "content": [{"type": "input_text", "text": "What did you say?"}]}, - ] - - def test_legacy_turn_input_maps_model_role_to_assistant(self): - transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input( - [ - {"role": "user", "content": [{"type": "text", "text": "I like apples."}]}, - {"role": "model", "content": [{"type": "text", "text": "I like oranges."}]}, - ] - ) - assert transformed == [ - {"role": "user", "content": [{"type": "input_text", "text": "I like apples."}]}, - {"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]}, - ] - - def test_turn_pydantic_model_with_string_content(self): - transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input( - [Turn(role="model", content="I like oranges.")] - ) - assert transformed == [ - {"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]} - ] - - def test_string_input_passes_through(self): - transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input("Hello") - assert transformed == "Hello" - - def test_content_list_input_becomes_single_user_message(self): - transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input( - [{"type": "text", "text": "Hello"}, "world"] - ) - assert transformed == [ - { - "role": "user", - "content": [ - {"type": "input_text", "text": "Hello"}, - {"type": "input_text", "text": "world"}, - ], - } - ] - - def test_non_text_content_passes_through_unchanged(self): - image_part = {"type": "image", "data": "base64data", "mime_type": "image/png"} - transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input( - [{"type": "user_input", "content": [image_part]}] - ) - assert transformed == [{"role": "user", "content": [image_part]}] diff --git a/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py b/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py index 7a69b676667..f692259db2e 100644 --- a/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py +++ b/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py @@ -9,171 +9,6 @@ import os import pytest -from litellm.llms.cometapi.chat.transformation import ( - CometAPIChatCompletionStreamingHandler, - CometAPIConfig, -) -from litellm.llms.cometapi.common_utils import CometAPIException - - -class TestCometAPIChatCompletionStreamingHandler: - def test_chunk_parser_successful(self): - handler = CometAPIChatCompletionStreamingHandler( - streaming_response=None, sync_stream=True - ) - - # Test input chunk - chunk = { - "id": "test_id", - "created": 1234567890, - "model": "gpt-3.5-turbo", - "usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}, - "choices": [ - {"delta": {"content": "test content", "reasoning": "test reasoning"}} - ], - } - - # Parse chunk - result = handler.chunk_parser(chunk) - - # Verify response - assert result.id == "test_id" - assert result.object == "chat.completion.chunk" - assert result.created == 1234567890 - assert result.model == "gpt-3.5-turbo" - assert result.usage.prompt_tokens == chunk["usage"]["prompt_tokens"] - assert result.usage.completion_tokens == chunk["usage"]["completion_tokens"] - assert result.usage.total_tokens == chunk["usage"]["total_tokens"] - assert len(result.choices) == 1 - assert result.choices[0]["delta"]["reasoning_content"] == "test reasoning" - - def test_chunk_parser_error_response(self): - handler = CometAPIChatCompletionStreamingHandler( - streaming_response=None, sync_stream=True - ) - - # Test error chunk - error_chunk = { - "error": { - "message": "test error", - "code": 400, - } - } - - # Verify error handling - with pytest.raises(CometAPIException) as exc_info: - handler.chunk_parser(error_chunk) - - assert "CometAPI Error: test error" in str(exc_info.value) - assert exc_info.value.status_code == 400 - - def test_chunk_parser_key_error(self): - handler = CometAPIChatCompletionStreamingHandler( - streaming_response=None, sync_stream=True - ) - - # Test invalid chunk missing required fields - invalid_chunk = {"incomplete": "data"} - - # Verify KeyError handling - with pytest.raises(CometAPIException) as exc_info: - handler.chunk_parser(invalid_chunk) - - assert "KeyError" in str(exc_info.value) - assert exc_info.value.status_code == 400 - - -class TestCometAPIConfig: - def test_transform_request_basic(self): - """Test basic request transformation""" - config = CometAPIConfig() - - transformed_request = config.transform_request( - model="cometapi/gpt-3.5-turbo", - messages=[{"role": "user", "content": "Hello, world!"}], - optional_params={}, - litellm_params={}, - headers={}, - ) - - assert transformed_request["model"] == "cometapi/gpt-3.5-turbo" - assert transformed_request["messages"] == [ - {"role": "user", "content": "Hello, world!"} - ] - - def test_transform_request_with_extra_body(self): - """Test request transformation with extra_body parameters""" - config = CometAPIConfig() - - transformed_request = config.transform_request( - model="cometapi/gpt-4", - messages=[{"role": "user", "content": "Hello, world!"}], - optional_params={"extra_body": {"custom_param": "custom_value"}}, - litellm_params={}, - headers={}, - ) - - # Validate that extra_body parameters are merged into the request - assert transformed_request["custom_param"] == "custom_value" - assert transformed_request["messages"] == [ - {"role": "user", "content": "Hello, world!"} - ] - - def test_cache_control_flag_removal(self): - """Test cache control flag removal from messages""" - config = CometAPIConfig() - - transformed_request = config.transform_request( - model="cometapi/gpt-3.5-turbo", - messages=[ - { - "role": "user", - "content": "Hello, world!", - "cache_control": {"type": "ephemeral"}, - } - ], - optional_params={}, - litellm_params={}, - headers={}, - ) - - # CometAPI should remove cache_control flags by default - assert transformed_request["messages"][0].get("cache_control") is None - - def test_map_openai_params(self): - """Test OpenAI parameter mapping""" - config = CometAPIConfig() - - non_default_params = { - "temperature": 0.7, - "max_tokens": 100, - "top_p": 0.9, - } - - mapped_params = config.map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model="cometapi/gpt-3.5-turbo", - drop_params=False, - ) - - assert mapped_params["temperature"] == 0.7 - assert mapped_params["max_tokens"] == 100 - assert mapped_params["top_p"] == 0.9 - - def test_get_error_class(self): - """Test error class creation""" - config = CometAPIConfig() - - error = config.get_error_class( - error_message="Test error", - status_code=400, - headers={"Content-Type": "application/json"}, - ) - - assert isinstance(error, CometAPIException) - assert error.message == "Test error" - assert error.status_code == 400 # Integration test example (requires real API key) diff --git a/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py b/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py deleted file mode 100644 index a3391a2c585..00000000000 --- a/tests/test_litellm/llms/databricks/chat/test_databricks_chat_transformation.py +++ /dev/null @@ -1,79 +0,0 @@ -import json -from typing import Final - -import httpx -import respx - -import litellm - - -def test_completion_merges_leading_system_and_developer_messages_for_chat_template_models( - respx_mock: respx.MockRouter, -): - upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock( - return_value=httpx.Response( - status_code=200, - json={ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "my-custom-model", - "choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}], - "usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10}, - }, - ) - ) - - response: Final = litellm.completion( - model="databricks/my-custom-model", - messages=[ - {"role": "system", "content": "You are terse."}, - {"role": "developer", "content": "Skills: none."}, - {"role": "user", "content": "Hello"}, - ], - api_base="https://example.databricks.test/serving-endpoints", - api_key="fake-databricks-api-key", - num_retries=0, - ) - - assert upstream.call_count == 1 - request_body: Final = json.loads(upstream.calls[0].request.read()) - assert request_body["messages"] == [ - {"role": "system", "content": "You are terse.\n\nSkills: none."}, - {"role": "user", "content": "Hello"}, - ] - assert response.choices[0].message.content == "Answer" - - -def test_completion_merges_system_messages_when_one_has_empty_content(respx_mock: respx.MockRouter): - upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock( - return_value=httpx.Response( - status_code=200, - json={ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "my-custom-model", - "choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}], - "usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10}, - }, - ) - ) - - litellm.completion( - model="databricks/my-custom-model", - messages=[ - {"role": "system", "content": "You are terse."}, - {"role": "system", "content": ""}, - {"role": "user", "content": "Hello"}, - ], - api_base="https://example.databricks.test/serving-endpoints", - api_key="fake-databricks-api-key", - num_retries=0, - ) - - request_body: Final = json.loads(upstream.calls[0].request.read()) - assert request_body["messages"] == [ - {"role": "system", "content": "You are terse."}, - {"role": "user", "content": "Hello"}, - ] diff --git a/tests/test_litellm/llms/deepinfra/test_deepinfra_rerank_integration.py b/tests/test_litellm/llms/deepinfra/test_deepinfra_rerank_integration.py deleted file mode 100644 index 5b013681864..00000000000 --- a/tests/test_litellm/llms/deepinfra/test_deepinfra_rerank_integration.py +++ /dev/null @@ -1,433 +0,0 @@ -""" -Integration tests for DeepInfra rerank functionality. -Tests the full rerank flow following the repository patterns. -""" - -import asyncio -import json -from unittest.mock import AsyncMock, MagicMock, patch - -import pytest - -import litellm - - -def assert_response_shape(response, custom_llm_provider): - """Helper function to validate response structure specific to DeepInfra.""" - assert hasattr(response, "id") - assert hasattr(response, "results") - assert hasattr(response, "meta") - assert isinstance(response.results, list) - - for result in response.results: - assert "index" in result - assert "relevance_score" in result - assert isinstance(result["index"], int) - assert isinstance(result["relevance_score"], (int, float)) - - # Check meta structure - assert "tokens" in response.meta - assert "billed_units" in response.meta - assert "input_tokens" in response.meta["tokens"] - assert "total_tokens" in response.meta["billed_units"] - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post") -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") -def test_basic_rerank_deepinfra(mock_sync_post, mock_async_post, sync_mode): - """Test basic DeepInfra rerank functionality.""" - # Mock response data that matches DeepInfra API format - mock_response_data = { - "scores": [0.9, 0.1], - "input_tokens": 25, - "request_id": "deepinfra-request-123", - "inference_status": { - "status": "success", - "runtime_ms": 150, - "cost": 0.0001, - "tokens_generated": 0, - "tokens_input": 25, - }, - } - - def return_val(): - return mock_response_data - - api_key = "test_deepinfra_api_key" - api_base = "https://api.deepinfra.com" - - if sync_mode: - # Create mock response object for sync - mock_response = MagicMock() - mock_response.json = return_val - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.text = json.dumps(mock_response_data) - mock_sync_post.return_value = mock_response - - response = litellm.rerank( - model="deepinfra/Qwen/Qwen3-Reranker-0.6B", - query="hello", - documents=["hello", "world"], - top_n=2, - custom_llm_provider="deepinfra", - api_key=api_key, - api_base=api_base, - ) - mock_sync_post.assert_called_once() - else: - # Create mock response object for async - mock_response = AsyncMock() - - def return_val(): - return mock_response_data - - mock_response.json = return_val - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.text = json.dumps(mock_response_data) - mock_async_post.return_value = mock_response - - response = asyncio.run( - litellm.arerank( - model="deepinfra/Qwen/Qwen3-Reranker-0.6B", - query="hello", - documents=["hello", "world"], - top_n=2, - custom_llm_provider="deepinfra", - api_key=api_key, - api_base=api_base, - ) - ) - mock_async_post.assert_called_once() - - # Verify response structure - assert response.id == "deepinfra-request-123" - assert response.results is not None - assert len(response.results) == 2 - assert response.results[0]["index"] == 0 - assert response.results[0]["relevance_score"] == 0.9 - assert response.results[1]["index"] == 1 - assert response.results[1]["relevance_score"] == 0.1 - - # Verify metadata - assert response.meta["tokens"]["input_tokens"] == 25 - assert response.meta["billed_units"]["total_tokens"] == 25 - - # Verify hidden params specific to DeepInfra - assert response._hidden_params["status"] == "success" - assert response._hidden_params["runtime_ms"] == 150 - assert response._hidden_params["cost"] == 0.0001 - # Note: The model name is processed and the 'deepinfra/' prefix is removed - assert response._hidden_params["model"] == "Qwen/Qwen3-Reranker-0.6B" - - assert_response_shape(response, custom_llm_provider="deepinfra") - - -@pytest.mark.parametrize("sync_mode", [True, False]) -@patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post") -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") -def test_deepinfra_rerank_with_queries_param( - mock_sync_post, mock_async_post, sync_mode -): - """Test DeepInfra rerank with multiple queries parameter.""" - mock_response_data = { - "scores": [0.8, 0.6, 0.2], - "input_tokens": 35, - "request_id": "deepinfra-multi-query-123", - "inference_status": {"status": "success", "runtime_ms": 200}, - } - - def return_val(): - return mock_response_data - - if sync_mode: - mock_response = MagicMock() - mock_response.json = return_val - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.text = json.dumps(mock_response_data) - mock_sync_post.return_value = mock_response - - response = litellm.rerank( - model="deepinfra/Qwen/Qwen3-Reranker-4B", - query="hello", - documents=["hello", "world", "test"], - queries=["hello", "hi there"], # DeepInfra specific param - custom_llm_provider="deepinfra", - api_key="test_key", - api_base="https://api.deepinfra.com", - ) - - mock_sync_post.assert_called_once() - # Verify that queries parameter was passed in request - call_data = json.loads(mock_sync_post.call_args.kwargs["data"]) - assert "queries" in call_data - assert call_data["queries"] == ["hello", "hi there"] - else: - mock_response = AsyncMock() - mock_response.json = return_val - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.text = json.dumps(mock_response_data) - mock_async_post.return_value = mock_response - - response = asyncio.run( - litellm.arerank( - model="deepinfra/Qwen/Qwen3-Reranker-4B", - query="hello", - documents=["hello", "world", "test"], - queries=["hello", "hi there"], - custom_llm_provider="deepinfra", - api_key="test_key", - api_base="https://api.deepinfra.com", - ) - ) - - mock_async_post.assert_called_once() - call_data = json.loads(mock_async_post.call_args.kwargs["data"]) - assert "queries" in call_data - assert call_data["queries"] == ["hello", "hi there"] - - assert response.results is not None - assert len(response.results) == 3 - - -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") -def test_deepinfra_rerank_with_service_tier(mock_post): - """Test DeepInfra rerank with service_tier parameter.""" - mock_response_data = { - "scores": [0.95, 0.75], - "input_tokens": 30, - "request_id": "deepinfra-premium-123", - } - - def return_val(): - return mock_response_data - - mock_response = MagicMock() - mock_response.json = return_val - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.text = json.dumps(mock_response_data) - mock_post.return_value = mock_response - - response = litellm.rerank( - model="deepinfra/Qwen/Qwen3-Reranker-8B", - query="premium search", - documents=["doc1", "doc2"], - service_tier="premium", # DeepInfra specific param - custom_llm_provider="deepinfra", - api_key="test_key", - api_base="https://api.deepinfra.com", - ) - - mock_post.assert_called_once() - - # Verify URL - call_url = mock_post.call_args.kwargs["url"] - assert "api.deepinfra.com/inference/Qwen/Qwen3-Reranker-8B" in call_url - - # Verify request contains service_tier - call_data = json.loads(mock_post.call_args.kwargs["data"]) - assert call_data["service_tier"] == "premium" - - assert response.results is not None - - -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") -def test_deepinfra_rerank_with_env_vars(mock_post, monkeypatch): - """Test DeepInfra rerank with environment variable configuration.""" - monkeypatch.setenv("DEEPINFRA_API_KEY", "env_test_key") - monkeypatch.setenv("DEEPINFRA_API_BASE", "https://custom-deepinfra.com") - - mock_response_data = { - "scores": [0.88, 0.22], - "input_tokens": 28, - "request_id": "env-test-123", - } - - def return_val(): - return mock_response_data - - mock_response = MagicMock() - mock_response.json = return_val - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.text = json.dumps(mock_response_data) - mock_post.return_value = mock_response - - response = litellm.rerank( - model="deepinfra/Qwen/Qwen3-Reranker-0.6B", - query="hello", - documents=["hello", "world"], - custom_llm_provider="deepinfra", - ) - - mock_post.assert_called_once() - - # Verify headers contain env API key - headers = mock_post.call_args.kwargs.get("headers", {}) - assert "Bearer env_test_key" in headers.get("Authorization", "") - - assert response.results is not None - - -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") -def test_deepinfra_rerank_error_handling(mock_post): - """Test DeepInfra rerank error handling.""" - error_response = {"detail": {"error": "Invalid API key"}} - - def return_val(): - return error_response - - mock_response = MagicMock() - mock_response.status_code = 401 - mock_response.json = return_val - mock_response.text = json.dumps(error_response) - mock_response.headers = {"content-type": "application/json"} - mock_post.return_value = mock_response - - # The current implementation handles errors gracefully, so we expect a successful response - # with the error information in the hidden params - response = litellm.rerank( - model="deepinfra/Qwen/Qwen3-Reranker-0.6B", - query="hello", - documents=["hello", "world"], - custom_llm_provider="deepinfra", - api_key="invalid_key", - api_base="https://api.deepinfra.com", - ) - - # Verify that the response contains error information - assert ( - response._hidden_params["status"] == "unknown" - ) # Default status when error occurs - - -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") -def test_deepinfra_rerank_defaults_api_base_when_missing(mock_post, monkeypatch): - """With no api_base anywhere, the call still goes out against DeepInfra's own base.""" - monkeypatch.delenv("DEEPINFRA_API_BASE", raising=False) - - mock_response = MagicMock() - mock_response.json = lambda: {"scores": [0.9, 0.1], "input_tokens": 20} - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_post.return_value = mock_response - - response = litellm.rerank( - model="deepinfra/Qwen/Qwen3-Reranker-0.6B", - query="hello", - documents=["hello", "world"], - custom_llm_provider="deepinfra", - api_key="test_key", - # api_base is intentionally missing - ) - - assert "api.deepinfra.com" in mock_post.call_args.kwargs["url"] - assert [result["relevance_score"] for result in response.results] == [0.9, 0.1] - - -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") -def test_deepinfra_rerank_request_format(mock_post): - """Test that the request is properly formatted for DeepInfra API.""" - mock_response_data = {"scores": [0.9, 0.1], "input_tokens": 20} - - def return_val(): - return mock_response_data - - mock_response = MagicMock() - mock_response.json = return_val - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.text = json.dumps(mock_response_data) - mock_post.return_value = mock_response - - response = litellm.rerank( - model="deepinfra/Qwen/Qwen3-Reranker-0.6B", - query="test query", - documents=["doc1", "doc2"], - custom_llm_provider="deepinfra", - api_key="test_key", - api_base="https://api.deepinfra.com", - instruction="custom instruction", - webhook="https://webhook.example.com", - ) - - mock_post.assert_called_once() - - # Verify URL format - call_url = mock_post.call_args.kwargs["url"] - assert call_url == "https://api.deepinfra.com/inference/Qwen/Qwen3-Reranker-0.6B" - - # Verify headers - headers = mock_post.call_args.kwargs["headers"] - assert headers["Authorization"] == "Bearer test_key" - assert headers["accept"] == "application/json" - assert headers["content-type"] == "application/json" - - # Verify request body format - request_data = json.loads(mock_post.call_args.kwargs["data"]) - assert request_data["queries"] == [ - "test query", - "test query", - ] # DeepInfra requires queries to match documents length - assert request_data["documents"] == ["doc1", "doc2"] - assert request_data["instruction"] == "custom instruction" - assert request_data["webhook"] == "https://webhook.example.com" - - assert response.results is not None - - -def test_deepinfra_rerank_models(): - """Test that DeepInfra Qwen rerank models are recognized.""" - # These should not raise errors during model validation - models = [ - "deepinfra/Qwen/Qwen3-Reranker-0.6B", - "deepinfra/Qwen/Qwen3-Reranker-4B", - "deepinfra/Qwen/Qwen3-Reranker-8B", - ] - - for model in models: - resolved_model, provider, _, api_base = litellm.get_llm_provider(model=model) - assert provider == "deepinfra" - assert resolved_model == model.removeprefix("deepinfra/") - assert api_base == "https://api.deepinfra.com/v1/openai" - - -@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") -def test_deepinfra_rerank_minimal_response(mock_post): - """Test handling of minimal DeepInfra response.""" - # Minimal response with just scores - mock_response_data = {"scores": [0.7, 0.3]} - - def return_val(): - return mock_response_data - - mock_response = MagicMock() - mock_response.json = return_val - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.text = json.dumps(mock_response_data) - mock_post.return_value = mock_response - - response = litellm.rerank( - model="deepinfra/Qwen/Qwen3-Reranker-0.6B", - query="hello", - documents=["hello", "world"], - custom_llm_provider="deepinfra", - api_key="test_key", - api_base="https://api.deepinfra.com", - ) - - # Should handle minimal response gracefully - assert response.results is not None - assert len(response.results) == 2 - assert response.results[0]["relevance_score"] == 0.7 - assert response.results[1]["relevance_score"] == 0.3 - - # Should have default values for missing fields - assert response.meta["tokens"]["input_tokens"] == 0 # Default when missing - assert response._hidden_params["status"] == "unknown" # Default when missing diff --git a/tests/test_litellm/llms/gemini/files/__init__.py b/tests/test_litellm/llms/gemini/files/__init__.py deleted file mode 100644 index f48fe7dbe2b..00000000000 --- a/tests/test_litellm/llms/gemini/files/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Tests for Gemini files functionality""" diff --git a/tests/test_litellm/llms/gemini/videos/__init__.py b/tests/test_litellm/llms/gemini/videos/__init__.py deleted file mode 100644 index e0780c08321..00000000000 --- a/tests/test_litellm/llms/gemini/videos/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# Gemini Video Generation Tests diff --git a/tests/test_litellm/llms/manus/__init__.py b/tests/test_litellm/llms/manus/__init__.py deleted file mode 100644 index c9121a7b2a4..00000000000 --- a/tests/test_litellm/llms/manus/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# Manus provider tests diff --git a/tests/test_litellm/llms/manus/responses/__init__.py b/tests/test_litellm/llms/manus/responses/__init__.py deleted file mode 100644 index ea7ebb64d55..00000000000 --- a/tests/test_litellm/llms/manus/responses/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# Manus Responses API tests diff --git a/tests/test_litellm/llms/minimax/__init__.py b/tests/test_litellm/llms/minimax/__init__.py deleted file mode 100644 index 451f542f4ad..00000000000 --- a/tests/test_litellm/llms/minimax/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# MiniMax tests diff --git a/tests/test_litellm/llms/minimax/chat/__init__.py b/tests/test_litellm/llms/minimax/chat/__init__.py deleted file mode 100644 index 4a7916ae6cf..00000000000 --- a/tests/test_litellm/llms/minimax/chat/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# MiniMax chat tests diff --git a/tests/test_litellm/llms/minimax/messages/__init__.py b/tests/test_litellm/llms/minimax/messages/__init__.py deleted file mode 100644 index de5a80602ea..00000000000 --- a/tests/test_litellm/llms/minimax/messages/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# MiniMax messages tests diff --git a/tests/test_litellm/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py b/tests/test_litellm/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py index d1eb6241ceb..db77eabba23 100644 --- a/tests/test_litellm/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py +++ b/tests/test_litellm/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py @@ -1,19 +1,9 @@ import os from typing import Dict -from unittest.mock import MagicMock -import httpx import litellm import pytest -from litellm.llms.base_llm.audio_transcription.transformation import ( - BaseAudioTranscriptionConfig, -) -from litellm.llms.mistral.audio_transcription.transformation import ( - MistralAudioTranscriptionConfig, -) -from litellm.types.utils import TranscriptionResponse -from litellm.utils import ProviderConfigManager from tests.llm_translation.base_audio_transcription_unit_tests import ( BaseLLMAudioTranscriptionTest, ) @@ -37,184 +27,3 @@ class TestMistralAudioTranscription(BaseLLMAudioTranscriptionTest): "Async audio transcription test for Mistral is skipped in this suite; " "async test plugins (e.g. pytest-asyncio/anyio) are not configured here." ) - - -def test_mistral_audio_transcription_config_installed(): - """Ensure Mistral audio transcription config is registered with ProviderConfigManager.""" - config = ProviderConfigManager.get_provider_audio_transcription_config( - model="mistral/voxtral-mini-latest", - provider=litellm.LlmProviders.MISTRAL, - ) - assert config is not None - assert isinstance(config, BaseAudioTranscriptionConfig) - assert isinstance(config, MistralAudioTranscriptionConfig) - - -def test_mistral_audio_transcription_get_complete_url(): - config = MistralAudioTranscriptionConfig() - url = config.get_complete_url( - api_base=None, - api_key="fake-key", - model="voxtral-mini-latest", - optional_params={}, - litellm_params={}, - ) - assert url == "https://api.mistral.ai/v1/audio/transcriptions" - - -def test_mistral_audio_transcription_get_complete_url_custom_base(): - config = MistralAudioTranscriptionConfig() - url = config.get_complete_url( - api_base="https://custom.api.example.com/v1/", - api_key="fake-key", - model="voxtral-mini-latest", - optional_params={}, - litellm_params={}, - ) - assert url == "https://custom.api.example.com/v1/audio/transcriptions" - - -def test_mistral_audio_transcription_validate_environment(): - config = MistralAudioTranscriptionConfig() - headers = config.validate_environment( - headers={}, - model="voxtral-mini-latest", - messages=[], - optional_params={}, - litellm_params={}, - api_key="test-key-123", - ) - assert headers["Authorization"] == "Bearer test-key-123" - assert headers["accept"] == "application/json" - - -def test_mistral_audio_transcription_supported_params(): - config = MistralAudioTranscriptionConfig() - params = config.get_supported_openai_params("voxtral-mini-latest") - assert "language" in params - assert "temperature" in params - assert "response_format" in params - assert "timestamp_granularities" in params - - -def test_mistral_audio_transcription_request_transform(): - config = MistralAudioTranscriptionConfig() - - wav_path = os.path.join( - os.path.dirname(__file__), - "../../../../..", - "tests", - "llm_translation", - "gettysburg.wav", - ) - audio_file = open(wav_path, "rb") - - result = config.transform_audio_transcription_request( - model="voxtral-mini-latest", - audio_file=audio_file, - optional_params={"language": "en", "temperature": 0.0}, - litellm_params={}, - ) - - audio_file.close() - - assert isinstance(result.data, dict) - assert result.data["model"] == "voxtral-mini-latest" - assert result.data["language"] == "en" - assert result.data["temperature"] == 0.0 - assert result.files is not None - assert "file" in result.files - - -def test_mistral_audio_transcription_request_with_diarize(): - """Test that Mistral-specific params like diarize are passed through.""" - config = MistralAudioTranscriptionConfig() - - wav_path = os.path.join( - os.path.dirname(__file__), - "../../../../..", - "tests", - "llm_translation", - "gettysburg.wav", - ) - audio_file = open(wav_path, "rb") - - result = config.transform_audio_transcription_request( - model="voxtral-mini-latest", - audio_file=audio_file, - optional_params={"diarize": True}, - litellm_params={}, - ) - - audio_file.close() - - assert isinstance(result.data, dict) - assert result.data["diarize"] == "true" - - -def test_mistral_audio_transcription_response_transform(): - config = MistralAudioTranscriptionConfig() - - mock_response = MagicMock(spec=httpx.Response) - mock_response.json.return_value = {"text": "Four score and seven years ago..."} - - response = config.transform_audio_transcription_response(mock_response) - - assert isinstance(response, TranscriptionResponse) - assert response.text == "Four score and seven years ago..." - - -def test_mistral_audio_transcription_response_transform_diarized(): - """Test that diarized responses preserve segments and language.""" - config = MistralAudioTranscriptionConfig() - - mock_response = MagicMock(spec=httpx.Response) - mock_response.json.return_value = { - "model": "voxtral-mini-latest", - "text": "Hello, how are you? I am fine.", - "language": None, - "segments": [ - { - "text": "Hello, how are you?", - "start": 0.3, - "end": 2.1, - "speaker_id": "speaker_1", - "type": "transcription_segment", - }, - { - "text": "I am fine.", - "start": 2.5, - "end": 3.8, - "speaker_id": "speaker_2", - "type": "transcription_segment", - }, - ], - "usage": { - "prompt_audio_seconds": 4, - "prompt_tokens": 5, - "total_tokens": 50, - "completion_tokens": 20, - }, - } - - response = config.transform_audio_transcription_response(mock_response) - - assert isinstance(response, TranscriptionResponse) - assert response.text == "Hello, how are you? I am fine." - assert response["segments"] is not None - assert len(response["segments"]) == 2 - assert response["segments"][0]["speaker_id"] == "speaker_1" - assert response["segments"][1]["speaker_id"] == "speaker_2" - assert response["language"] is None - - -def test_mistral_audio_transcription_response_transform_empty(): - config = MistralAudioTranscriptionConfig() - - mock_response = MagicMock(spec=httpx.Response) - mock_response.json.return_value = {} - - response = config.transform_audio_transcription_response(mock_response) - - assert isinstance(response, TranscriptionResponse) - assert response.text == "" diff --git a/tests/test_litellm/llms/openai_like/test_json_providers.py b/tests/test_litellm/llms/openai_like/test_json_providers.py index d84cc8d3237..55703063fae 100644 --- a/tests/test_litellm/llms/openai_like/test_json_providers.py +++ b/tests/test_litellm/llms/openai_like/test_json_providers.py @@ -3,321 +3,12 @@ Tests for JSON-based provider configuration system. """ import os -import sys -from unittest.mock import patch -try: - import pytest -except ImportError: - # pytest not available, will run as standalone script - pytest = None - -# Add workspace to path -workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../..")) -sys.path.insert(0, workspace_path) +import pytest import litellm -class TestJSONProviderLoader: - """Test JSON provider loading and configuration""" - - def test_load_json_providers(self): - """Test that JSON providers load correctly""" - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - # Verify publicai is loaded - assert JSONProviderRegistry.exists("publicai") - - # Get publicai config - publicai = JSONProviderRegistry.get("publicai") - assert publicai is not None - assert publicai.base_url == "https://api.publicai.co/v1" - assert publicai.api_key_env == "PUBLICAI_API_KEY" - assert publicai.api_base_env == "PUBLICAI_API_BASE" - assert publicai.param_mappings.get("max_completion_tokens") == "max_tokens" - - def test_dynamic_config_generation(self): - """Test dynamic config class creation""" - from litellm.llms.openai_like.dynamic_config import create_config_class - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - provider = JSONProviderRegistry.get("publicai") - config_class = create_config_class(provider) - config = config_class() - - # Test API info resolution - api_base, api_key = config._get_openai_compatible_provider_info(None, None) - assert api_base == "https://api.publicai.co/v1" - - # Test with custom base - api_base, api_key = config._get_openai_compatible_provider_info( - "https://custom.api.com", "test-key" - ) - assert api_base == "https://custom.api.com" - assert api_key == "test-key" - - def test_parameter_mapping(self): - """Test parameter mapping works""" - from litellm.llms.openai_like.dynamic_config import create_config_class - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - provider = JSONProviderRegistry.get("publicai") - config_class = create_config_class(provider) - config = config_class() - - # Test parameter mapping - optional_params = {} - non_default_params = {"max_completion_tokens": 100, "temperature": 0.7} - result = config.map_openai_params( - non_default_params, optional_params, "gpt-4", False - ) - - # max_completion_tokens should be mapped to max_tokens - assert "max_tokens" in result - assert result["max_tokens"] == 100 - assert "max_completion_tokens" not in result - - # temperature should be passed through - assert result["temperature"] == 0.7 - - def test_supported_params(self): - """Test that config returns supported params""" - from litellm.llms.openai_like.dynamic_config import create_config_class - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - provider = JSONProviderRegistry.get("publicai") - config_class = create_config_class(provider) - config = config_class() - - # Get supported params - supported = config.get_supported_openai_params("gpt-4") - - # Should have standard OpenAI params - assert isinstance(supported, list) - assert len(supported) > 0 - - def test_tool_params_excluded_when_function_calling_not_supported(self): - """Test that tool-related params are excluded for models that don't support - function calling. Regression test for https://github.com/BerriAI/litellm/issues/21125 - """ - from litellm.llms.openai_like.dynamic_config import create_config_class - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - provider = JSONProviderRegistry.get("publicai") - config_class = create_config_class(provider) - config = config_class() - - # Mock supports_function_calling to return False - with patch("litellm.utils.supports_function_calling", return_value=False): - supported = config.get_supported_openai_params("some-model-without-fc") - - tool_params = [ - "tools", - "tool_choice", - "function_call", - "functions", - "parallel_tool_calls", - ] - for param in tool_params: - assert ( - param not in supported - ), f"'{param}' should not be in supported params when function calling is not supported" - - # Non-tool params should still be present - assert "temperature" in supported - assert "max_tokens" in supported - assert "stop" in supported - - def test_tool_params_included_when_function_calling_supported(self): - """Test that tool-related params are included for models that support function calling.""" - from litellm.llms.openai_like.dynamic_config import create_config_class - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - provider = JSONProviderRegistry.get("publicai") - config_class = create_config_class(provider) - config = config_class() - - # Mock supports_function_calling to return True - with patch("litellm.utils.supports_function_calling", return_value=True): - supported = config.get_supported_openai_params("some-model-with-fc") - - assert "tools" in supported - assert "tool_choice" in supported - - def test_provider_resolution(self): - """Test that provider resolution finds JSON providers""" - from litellm.litellm_core_utils.get_llm_provider_logic import ( - get_llm_provider, - ) - - model, provider, api_key, api_base = get_llm_provider( - model="publicai/gpt-4", - custom_llm_provider=None, - api_base=None, - api_key=None, - ) - - assert model == "gpt-4" - assert provider == "publicai" - assert api_base == "https://api.publicai.co/v1" - - def test_provider_config_manager(self): - """Test that ProviderConfigManager returns JSON-based configs""" - from litellm import LlmProviders - from litellm.utils import ProviderConfigManager - - config = ProviderConfigManager.get_provider_chat_config( - model="gpt-4", provider=LlmProviders.PUBLICAI - ) - - assert config is not None - assert config.custom_llm_provider == "publicai" - - -class TestPinstripes: - """Tests for Pinstripes JSON-configured provider""" - - def test_pinstripes_json_config_exists(self): - """Test that pinstripes is configured in providers.json""" - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - assert JSONProviderRegistry.exists("pinstripes") - - pinstripes = JSONProviderRegistry.get("pinstripes") - assert pinstripes is not None - assert pinstripes.base_url == "https://pinstripes.io/v1" - assert pinstripes.api_key_env == "PINSTRIPES_API_KEY" - assert pinstripes.param_mappings.get("max_completion_tokens") == "max_tokens" - - def test_pinstripes_provider_resolution(self): - """Test that provider resolution finds pinstripes and returns the default base URL""" - from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider - - model, provider, api_key, api_base = get_llm_provider( - model="pinstripes/ps/glm-4.5-air", - custom_llm_provider=None, - api_base=None, - api_key=None, - ) - - assert model == "ps/glm-4.5-air" - assert provider == "pinstripes" - assert api_base == "https://pinstripes.io/v1" - - def test_pinstripes_dynamic_config(self): - """Test dynamic config class creation for pinstripes""" - from litellm.llms.openai_like.dynamic_config import create_config_class - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - provider = JSONProviderRegistry.get("pinstripes") - config_class = create_config_class(provider) - config = config_class() - - api_base, api_key = config._get_openai_compatible_provider_info(None, None) - assert api_base == "https://pinstripes.io/v1" - - api_base, api_key = config._get_openai_compatible_provider_info( - "https://custom.pinstripes.io/v1", "test-key" - ) - assert api_base == "https://custom.pinstripes.io/v1" - assert api_key == "test-key" - - def test_pinstripes_parameter_mapping(self): - """Test that max_completion_tokens is mapped to max_tokens for pinstripes""" - from litellm.llms.openai_like.dynamic_config import create_config_class - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - provider = JSONProviderRegistry.get("pinstripes") - config_class = create_config_class(provider) - config = config_class() - - optional_params = {} - non_default_params = {"max_completion_tokens": 100, "temperature": 0.7} - result = config.map_openai_params( - non_default_params, optional_params, "ps/glm-4.5-air", False - ) - - assert "max_tokens" in result - assert result["max_tokens"] == 100 - assert "max_completion_tokens" not in result - assert result["temperature"] == 0.7 - - -class TestDarkbloom: - def test_darkbloom_json_config_exists(self): - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - darkbloom = JSONProviderRegistry.get("darkbloom") - assert darkbloom is not None - assert darkbloom.base_url == "https://api.darkbloom.dev/v1" - assert darkbloom.api_key_env == "DARKBLOOM_API_KEY" - assert darkbloom.api_base_env == "DARKBLOOM_API_BASE" - assert darkbloom.param_mappings.get("max_completion_tokens") == "max_tokens" - - def test_darkbloom_provider_resolution(self): - from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider - - model, provider, api_key, api_base = get_llm_provider( - model="darkbloom/gemma-4-26b", - custom_llm_provider=None, - api_base=None, - api_key=None, - ) - - assert model == "gemma-4-26b" - assert provider == "darkbloom" - assert api_key is None - assert api_base == "https://api.darkbloom.dev/v1" - - def test_darkbloom_dynamic_config(self): - from litellm.llms.openai_like.dynamic_config import create_config_class - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - provider = JSONProviderRegistry.get("darkbloom") - config_class = create_config_class(provider) - config = config_class() - - api_base, api_key = config._get_openai_compatible_provider_info(None, None) - assert api_base == "https://api.darkbloom.dev/v1" - - api_base, api_key = config._get_openai_compatible_provider_info( - "https://custom.darkbloom.dev/v1", "test-key" - ) - assert api_base == "https://custom.darkbloom.dev/v1" - assert api_key == "test-key" - - def test_darkbloom_complete_url_appends_endpoint(self): - from litellm.llms.openai_like.dynamic_config import create_config_class - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - provider = JSONProviderRegistry.get("darkbloom") - config_class = create_config_class(provider) - config = config_class() - - url = config.get_complete_url( - api_base="https://api.darkbloom.dev/v1", - api_key="test-key", - model="darkbloom/gemma-4-26b", - optional_params={}, - litellm_params={}, - stream=True, - ) - - assert url == "https://api.darkbloom.dev/v1/chat/completions" - - def test_darkbloom_provider_config_manager(self): - from litellm import LlmProviders - from litellm.utils import ProviderConfigManager - - config = ProviderConfigManager.get_provider_chat_config( - model="gemma-4-26b", provider=LlmProviders.DARKBLOOM - ) - - assert config is not None - assert config.custom_llm_provider == "darkbloom" - - class TestPublicAIIntegration: """Integration tests for PublicAI provider""" @@ -457,55 +148,3 @@ class TestPublicAIIntegration: pytest.fail(f"Content list conversion test failed: {str(e)}") else: raise - - -if __name__ == "__main__": - # Run basic tests - print("Testing JSON Provider System...") - - test_loader = TestJSONProviderLoader() - print("\n1. Testing JSON provider loading...") - test_loader.test_load_json_providers() - print(" ✓ JSON providers loaded") - - print("\n2. Testing dynamic config generation...") - test_loader.test_dynamic_config_generation() - print(" ✓ Dynamic config works") - - print("\n3. Testing parameter mapping...") - test_loader.test_parameter_mapping() - print(" ✓ Parameter mapping works") - - print("\n4. Testing excluded params...") - test_loader.test_excluded_params() - print(" ✓ Excluded params work") - - print("\n5. Testing provider resolution...") - test_loader.test_provider_resolution() - print(" ✓ Provider resolution works") - - print("\n6. Testing provider config manager...") - test_loader.test_provider_config_manager() - print(" ✓ Config manager works") - - print("\n" + "=" * 50) - print("PublicAI Integration Tests...") - print("=" * 50) - - test_integration = TestPublicAIIntegration() - - print("\n7. Testing basic completion...") - test_integration.test_publicai_completion_basic() - - print("\n8. Testing streaming...") - test_integration.test_publicai_completion_with_streaming() - - print("\n9. Testing parameter mapping...") - test_integration.test_publicai_parameter_mapping() - - print("\n10. Testing content list conversion...") - test_integration.test_publicai_content_list_conversion() - - print("\n" + "=" * 50) - print("✓ All tests passed!") - print("=" * 50) diff --git a/tests/test_litellm/llms/openai_like/test_xiaomi_mimo.py b/tests/test_litellm/llms/openai_like/test_xiaomi_mimo.py index 8104fb12943..580994f60b8 100644 --- a/tests/test_litellm/llms/openai_like/test_xiaomi_mimo.py +++ b/tests/test_litellm/llms/openai_like/test_xiaomi_mimo.py @@ -4,86 +4,12 @@ Related to issue #18794 """ import os -import sys -from unittest.mock import MagicMock, patch -try: - import pytest -except ImportError: - pytest = None - -# Add workspace to path -workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../..")) -sys.path.insert(0, workspace_path) +import pytest import litellm -class TestXiaomiMiMoProviderConfig: - """Test Xiaomi MiMo provider configuration""" - - def test_xiaomi_mimo_in_provider_list(self): - """Test that xiaomi_mimo is in the provider list (fixes #18794)""" - from litellm import LlmProviders - - # Verify xiaomi_mimo is in the enum - assert hasattr(LlmProviders, "XIAOMI_MIMO") - assert LlmProviders.XIAOMI_MIMO.value == "xiaomi_mimo" - - # Verify it's in the provider list - assert "xiaomi_mimo" in litellm.provider_list - - def test_xiaomi_mimo_json_config_exists(self): - """Test that xiaomi_mimo is configured in providers.json""" - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - # Verify xiaomi_mimo is loaded - assert JSONProviderRegistry.exists("xiaomi_mimo") - - # Get xiaomi_mimo config - xiaomi_mimo = JSONProviderRegistry.get("xiaomi_mimo") - assert xiaomi_mimo is not None - assert xiaomi_mimo.base_url == "https://api.xiaomimimo.com/v1" - assert xiaomi_mimo.api_key_env == "XIAOMI_MIMO_API_KEY" - assert xiaomi_mimo.param_mappings.get("max_completion_tokens") == "max_tokens" - - def test_xiaomi_mimo_provider_resolution(self): - """Test that provider resolution finds xiaomi_mimo""" - from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider - - model, provider, api_key, api_base = get_llm_provider( - model="xiaomi_mimo/mimo-v2-flash", - custom_llm_provider=None, - api_base=None, - api_key=None, - ) - - assert model == "mimo-v2-flash" - assert provider == "xiaomi_mimo" - assert api_base == "https://api.xiaomimimo.com/v1" - - def test_xiaomi_mimo_router_config(self): - """Test that xiaomi_mimo can be used in Router configuration (fixes #18794)""" - from litellm import Router - - # This should not raise "Unsupported provider - xiaomi_mimo" - router = Router( - model_list=[ - { - "model_name": "mimo-v2-flash", - "litellm_params": { - "model": "xiaomi_mimo/mimo-v2-flash", - "api_key": "test-key", - }, - } - ] - ) - - # Verify the deployment was created successfully - assert len(router.model_list) == 1 - assert router.model_list[0]["model_name"] == "mimo-v2-flash" - - class TestXiaomiMiMoIntegration: """Integration tests for Xiaomi MiMo provider""" @@ -128,30 +54,3 @@ class TestXiaomiMiMoIntegration: pytest.fail(f"Xiaomi MiMo completion failed: {str(e)}") else: raise - - -if __name__ == "__main__": - # Run basic tests - print("Testing Xiaomi MiMo Provider...") - - test_config = TestXiaomiMiMoProviderConfig() - - print("\n1. Testing provider in list...") - test_config.test_xiaomi_mimo_in_provider_list() - print(" ✓ xiaomi_mimo in provider list") - - print("\n2. Testing JSON config...") - test_config.test_xiaomi_mimo_json_config_exists() - print(" ✓ xiaomi_mimo JSON config loaded") - - print("\n3. Testing provider resolution...") - test_config.test_xiaomi_mimo_provider_resolution() - print(" ✓ Provider resolution works") - - print("\n4. Testing router configuration...") - test_config.test_xiaomi_mimo_router_config() - print(" ✓ Router configuration works (issue #18794 fixed)") - - print("\n" + "=" * 50) - print("✓ All configuration tests passed!") - print("=" * 50) diff --git a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py index c8751fb2d95..8cc46dc98d0 100644 --- a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py +++ b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py @@ -54,61 +54,3 @@ def test_ovhcloud_audio_transcription_config_installed(): assert config is not None assert isinstance(config, BaseAudioTranscriptionConfig) - - - -class TestOVHCloudDurationFieldMigration: - """Tests for OVHCloud duration -> seconds field migration.""" - - def test_seconds_field_mapped_to_duration(self): - """New `seconds` field should be normalized to `duration`.""" - from litellm.llms.ovhcloud.audio_transcription.transformation import ( - OVHCloudAudioTranscriptionConfig, - ) - from unittest.mock import MagicMock - - config = OVHCloudAudioTranscriptionConfig() - mock_response = MagicMock() - mock_response.json.return_value = { - "text": "Hello world", - "seconds": 3.14, - } - - result = config.transform_audio_transcription_response(mock_response) - - assert result.text == "Hello world" - assert result._hidden_params["duration"] == 3.14 - - def test_legacy_duration_field_still_works(self): - """Legacy `duration` field should still be accepted.""" - from litellm.llms.ovhcloud.audio_transcription.transformation import ( - OVHCloudAudioTranscriptionConfig, - ) - from unittest.mock import MagicMock - - config = OVHCloudAudioTranscriptionConfig() - mock_response = MagicMock() - mock_response.json.return_value = { - "text": "Hello world", - "duration": 2.71, - } - - result = config.transform_audio_transcription_response(mock_response) - - assert result.text == "Hello world" - assert result._hidden_params["duration"] == 2.71 - - - - def test_seconds_zero_mapped_to_duration(self): - """seconds=0.0 must not be treated as falsy and lost.""" - from litellm.llms.ovhcloud.audio_transcription.transformation import ( - OVHCloudAudioTranscriptionConfig, - ) - from unittest.mock import MagicMock - - config = OVHCloudAudioTranscriptionConfig() - mock_response = MagicMock() - 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 \ No newline at end of file diff --git a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py index 057ab9ede9a..34954587ed0 100644 --- a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py +++ b/tests/test_litellm/llms/ovhcloud/test_ovhcloud_chat_transformation.py @@ -6,174 +6,12 @@ import os import pytest -from litellm.llms.ovhcloud.utils import OVHCloudException -from litellm.utils import get_optional_params -from litellm.llms.ovhcloud.chat.transformation import ( - OVHCloudChatCompletionStreamingHandler, - OVHCloudChatConfig, -) -config = OVHCloudChatConfig() model = "ovhcloud/Mistral-7B-Instruct-v0.3" -class TestOvhCloudChatCompletionStreamingHandler: - def test_chunk_parser_successful(self): - handler = OVHCloudChatCompletionStreamingHandler( - streaming_response=None, sync_stream=True - ) - - chunk = { - "id": "test_id", - "created": 1234567890, - "model": "gpt-oss-20b", - "usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}, - "choices": [ - {"delta": {"content": "test content", "reasoning": "test reasoning"}} - ], - } - - result = handler.chunk_parser(chunk) - - assert result.id == "test_id" - assert result.object == "chat.completion.chunk" - assert result.created == 1234567890 - assert result.model == "gpt-oss-20b" - assert result.usage.prompt_tokens == chunk["usage"]["prompt_tokens"] - assert result.usage.completion_tokens == chunk["usage"]["completion_tokens"] - assert result.usage.total_tokens == chunk["usage"]["total_tokens"] - assert len(result.choices) == 1 - assert result.choices[0]["delta"]["reasoning_content"] == "test reasoning" - - def test_chunk_parser_error_response(self): - handler = OVHCloudChatCompletionStreamingHandler( - streaming_response=None, sync_stream=True - ) - - error_chunk = { - "error": { - "message": "test error", - "code": 400, - } - } - - with pytest.raises(OVHCloudException) as exc_info: - handler.chunk_parser(error_chunk) - - assert "OVHCloud Error: test error" in str(exc_info.value) - assert exc_info.value.status_code == 400 - - def test_chunk_parser_key_error(self): - handler = OVHCloudChatCompletionStreamingHandler( - streaming_response=None, sync_stream=True - ) - - invalid_chunk = {"incomplete": "data"} - - with pytest.raises(OVHCloudException) as exc_info: - handler.chunk_parser(invalid_chunk) - - assert "KeyError" in str(exc_info.value) - assert exc_info.value.status_code == 400 - - -class TestOVHCloudConfig: - def test_transform_request_basic(self): - """Test basic request transformation""" - transformed_request = config.transform_request( - model, - messages=[{"role": "user", "content": "Hello, world!"}], - optional_params={}, - litellm_params={}, - headers={}, - ) - - assert transformed_request["model"] == model - assert transformed_request["messages"] == [ - {"role": "user", "content": "Hello, world!"} - ] - - def test_transform_request_with_extra_body(self): - """Test request transformation with extra_body parameters""" - transformed_request = config.transform_request( - model, - messages=[{"role": "user", "content": "Hello, world!"}], - optional_params={"extra_body": {"custom_param": "custom_value"}}, - litellm_params={}, - headers={}, - ) - - assert transformed_request["custom_param"] == "custom_value" - assert transformed_request["messages"] == [ - {"role": "user", "content": "Hello, world!"} - ] - - def test_map_openai_params(self): - """Test OpenAI parameter mapping""" - non_default_params = { - "temperature": 0.7, - "max_tokens": 100, - "top_p": 0.9, - } - - mapped_params = config.map_openai_params( - non_default_params=non_default_params, - optional_params={}, - model=model, - drop_params=False, - ) - - assert mapped_params["temperature"] == 0.7 - assert mapped_params["max_tokens"] == 100 - assert mapped_params["top_p"] == 0.9 - - def test_get_error_class(self): - """Test error class creation""" - error = config.get_error_class( - error_message="Test error", - status_code=400, - headers={"Content-Type": "application/json"}, - ) - - assert isinstance(error, OVHCloudException) - assert error.message == "Test error" - assert error.status_code == 400 - - @pytest.mark.parametrize( - "model", - [ - "Meta-Llama-3_3-70B-Instruct", - "Meta-Llama-3_1-70B-Instruct", - "Mixtral-8x7B-Instruct-v0.1", - "gpt-oss-120b", - "some-model-not-in-the-cost-map", - ], - ) - def test_tools_not_filtered_by_static_model_map(self, model): - """ - OVHCloud AI Endpoints are OpenAI-compatible; tools/tool_choice must pass - through for any model. The server is responsible for rejecting unsupported - tool calls — LiteLLM must not strip them based on a stale static catalog. - """ - - params = get_optional_params( - model=model, - custom_llm_provider="ovhcloud", - tools=[ - { - "type": "function", - "function": {"name": "x", "parameters": {}}, - } - ], - tool_choice="auto", - ) - - assert "tools" in params - assert "tool_choice" in params - - def test_ovhcloud_integration(): from litellm import completion @@ -285,78 +123,3 @@ def test_ovhcloud_with_custom_base_url(): if __name__ == "__main__": pytest.main([__file__, "-v"]) - - -class TestOVHCloudReasoningFieldMigration: - """Tests for OVHCloud reasoning_content -> reasoning field migration.""" - - def test_streaming_new_reasoning_field(self): - """New `reasoning` field should be mapped to `reasoning_content`.""" - handler = OVHCloudChatCompletionStreamingHandler( - streaming_response=iter([]), - sync_stream=True, - ) - chunk = { - "id": "test-id", - "created": 1234567890, - "model": "test-model", - "choices": [ - { - "delta": { - "role": "assistant", - "reasoning": "Let me think...", - }, - "index": 0, - } - ], - } - result = handler.chunk_parser(chunk) - assert result.choices[0]["delta"]["reasoning_content"] == "Let me think..." - - def test_streaming_legacy_reasoning_content_unchanged(self): - """Legacy `reasoning_content` field should pass through untouched.""" - handler = OVHCloudChatCompletionStreamingHandler( - streaming_response=iter([]), - sync_stream=True, - ) - chunk = { - "id": "test-id", - "created": 1234567890, - "model": "test-model", - "choices": [ - { - "delta": { - "role": "assistant", - "reasoning_content": "Already correct field.", - }, - "index": 0, - } - ], - } - result = handler.chunk_parser(chunk) - assert result.choices[0]["delta"]["reasoning_content"] == "Already correct field." - - def test_streaming_both_fields_legacy_wins(self): - """When both fields present, existing `reasoning_content` is not overwritten.""" - handler = OVHCloudChatCompletionStreamingHandler( - streaming_response=iter([]), - sync_stream=True, - ) - chunk = { - "id": "test-id", - "created": 1234567890, - "model": "test-model", - "choices": [ - { - "delta": { - "reasoning": "new field", - "reasoning_content": "legacy field", - }, - "index": 0, - } - ], - } - result = handler.chunk_parser(chunk) - assert result.choices[0]["delta"]["reasoning_content"] == "legacy field" - - diff --git a/tests/test_litellm/llms/reducto/__init__.py b/tests/test_litellm/llms/reducto/__init__.py deleted file mode 100644 index 8b137891791..00000000000 --- a/tests/test_litellm/llms/reducto/__init__.py +++ /dev/null @@ -1 +0,0 @@ - diff --git a/tests/test_litellm/llms/s3_vectors/__init__.py b/tests/test_litellm/llms/s3_vectors/__init__.py deleted file mode 100644 index d4b0c4d8550..00000000000 --- a/tests/test_litellm/llms/s3_vectors/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# S3 Vectors tests diff --git a/tests/test_litellm/llms/s3_vectors/vector_stores/__init__.py b/tests/test_litellm/llms/s3_vectors/vector_stores/__init__.py deleted file mode 100644 index 231735c1de7..00000000000 --- a/tests/test_litellm/llms/s3_vectors/vector_stores/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# S3 Vectors vector store tests diff --git a/tests/test_litellm/llms/soniox/__init__.py b/tests/test_litellm/llms/soniox/__init__.py deleted file mode 100644 index b2cd496d66a..00000000000 --- a/tests/test_litellm/llms/soniox/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Soniox provider tests.""" diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index 4679b978f78..d3a7ba7a1bd 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -1,1681 +1,13 @@ -import base64 - import pytest from litellm.litellm_core_utils.prompt_templates.factory import ( convert_to_gemini_tool_call_result, ) -from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - _transform_request_body, - check_if_part_exists_in_parts, - _get_highest_media_resolution, - _extract_max_media_resolution_from_messages, -) from litellm.types.llms.vertex_ai import BlobType -from litellm.types.utils import Message - - -def test_check_if_part_exists_in_parts(): - parts = [ - {"text": "Hello", "thought": True}, - {"text": "World", "thought": False}, - ] - part = {"text": "Hello", "thought": True} - new_part = {"text": "Hello World", "thought": True} - assert check_if_part_exists_in_parts(parts, part) - assert not check_if_part_exists_in_parts(parts, new_part, ["thought"]) - assert check_if_part_exists_in_parts(parts, new_part, ["text"]) - - -def test_check_if_part_exists_in_parts_camel_case_snake_case(): - """Test that function handles both camelCase and snake_case key variations""" - # Test snake_case to camelCase matching - parts_with_snake_case = [ - { - "function_call": { - "name": "get_current_weather", - "args": {"location": "San Francisco, CA"}, - } - }, - {"text": "Some other content"}, - ] - - part_with_camel_case = { - "functionCall": { - "name": "get_current_weather", - "args": {"location": "San Francisco, CA"}, - } - } - - # Should find match between function_call and functionCall - assert check_if_part_exists_in_parts(parts_with_snake_case, part_with_camel_case) - - # Test camelCase to snake_case matching - parts_with_camel_case = [ - {"functionCall": {"name": "calculate_sum", "args": {"a": 1, "b": 2}}} - ] - - part_with_snake_case = { - "function_call": {"name": "calculate_sum", "args": {"a": 1, "b": 2}} - } - - # Should find match between functionCall and function_call - assert check_if_part_exists_in_parts(parts_with_camel_case, part_with_snake_case) - - # Test no match when values differ - part_with_different_values = { - "function_call": {"name": "different_function", "args": {"x": 5}} - } - - assert not check_if_part_exists_in_parts( - parts_with_snake_case, part_with_different_values - ) - - # Test multiple keys with mixed casing - parts_mixed = [ - { - "function_call": {"name": "test"}, - "thoughtSignature": "reasoning", - "text": "content", - } - ] - - part_mixed_casing = { - "functionCall": {"name": "test"}, - "thought_signature": "reasoning", - "text": "content", - } - - assert check_if_part_exists_in_parts(parts_mixed, part_mixed_casing) - - -def test_cached_content_respects_modify_params_for_cache_incompatible_fields(): - """Regression: cachedContent drops system/tools/toolConfig only when modify_params=True.""" - import litellm - - cache_name = "projects/p/locations/us-central1/cachedContents/abc123" - messages = [ - {"role": "system", "content": "You are helpful"}, - {"role": "user", "content": "hi"}, - ] - optional_params = { - "tools": [ - { - "functionDeclarations": [ - {"name": "get_weather", "description": "Get weather"}, - ] - } - ], - "tool_choice": {"functionCallingConfig": {"mode": "AUTO"}}, - } - - original_modify_params = litellm.modify_params - try: - # With modify_params=False (default), keep fields even with cachedContent. - litellm.modify_params = False - result = _transform_request_body( - messages=list(messages), - model="gemini-2.5-pro", - optional_params=dict(optional_params), - custom_llm_provider="vertex_ai", - litellm_params={}, - cached_content=cache_name, - ) - assert result.get("cachedContent") == cache_name - assert "system_instruction" in result - assert "tools" in result - assert "toolConfig" in result - assert "contents" in result - - # With modify_params=True, drop cache-incompatible fields. - litellm.modify_params = True - result_modify_true = _transform_request_body( - messages=list(messages), - model="gemini-2.5-pro", - optional_params=dict(optional_params), - custom_llm_provider="vertex_ai", - litellm_params={}, - cached_content=cache_name, - ) - assert result_modify_true.get("cachedContent") == cache_name - assert "system_instruction" not in result_modify_true - assert "tools" not in result_modify_true - assert "toolConfig" not in result_modify_true - assert "contents" in result_modify_true - - # Without cache, fields are always included. - result_no_cache = _transform_request_body( - messages=list(messages), - model="gemini-2.5-pro", - optional_params=dict(optional_params), - custom_llm_provider="vertex_ai", - litellm_params={}, - cached_content=None, - ) - assert "system_instruction" in result_no_cache - assert "tools" in result_no_cache - assert "toolConfig" in result_no_cache - finally: - litellm.modify_params = original_modify_params - - -# Tests for issue #14556: Labels field provider-aware filtering -def test_google_genai_excludes_labels(): - """Test that Google GenAI/AI Studio endpoints exclude labels when custom_llm_provider='gemini'""" - messages = [{"role": "user", "content": "test"}] - optional_params = {"labels": {"project": "test", "team": "ai"}} - litellm_params = {} - - result = _transform_request_body( - messages=messages, - model="gemini-2.5-pro", - optional_params=optional_params, - custom_llm_provider="gemini", - litellm_params=litellm_params, - cached_content=None, - ) - - # Google GenAI/AI Studio should NOT include labels - assert "labels" not in result - assert "contents" in result - - -def test_vertex_ai_includes_labels(): - """Test that Vertex AI endpoints include labels when custom_llm_provider='vertex_ai'""" - messages = [{"role": "user", "content": "test"}] - optional_params = {"labels": {"project": "test", "team": "ai"}} - litellm_params = {} - - result = _transform_request_body( - messages=messages, - model="gemini-2.5-pro", - optional_params=optional_params, - custom_llm_provider="vertex_ai", - litellm_params=litellm_params, - cached_content=None, - ) - - # Vertex AI SHOULD include labels - assert "labels" in result - assert result["labels"] == {"project": "test", "team": "ai"} - - -def test_service_tier_forwarded_to_vertex_ai(): - """Test that service_tier in optional_params is mapped to serviceTier in request body.""" - messages = [{"role": "user", "content": "test"}] - optional_params = {"service_tier": "flex"} - litellm_params = {} - - result = _transform_request_body( - messages=messages, - model="gemini-2.5-pro", - optional_params=optional_params, - custom_llm_provider="vertex_ai", - litellm_params=litellm_params, - cached_content=None, - ) - - assert "serviceTier" in result - assert result["serviceTier"] == "flex" - - -def test_extra_body_cache_not_forwarded_to_vertex_ai(): - """ - 'cache' inside extra_body is a LiteLLM-internal proxy caching control. - It must NOT be forwarded to the Vertex AI request body. - - Regression test for: "Invalid JSON payload received. Unknown name \"cache\": Cannot find field." - Vertex AI enforces a strict JSON schema and rejects any unknown field. - """ - messages = [{"role": "user", "content": "test"}] - optional_params = { - "extra_body": { - "cache": {"use-cache": True, "ttl": 86400}, # LiteLLM-internal - "some_vertex_param": "value", # legitimate provider extra - }, - } - litellm_params = {} - - result = _transform_request_body( - messages=messages, - model="gemini-2.5-pro", - optional_params=optional_params, - custom_llm_provider="vertex_ai", - litellm_params=litellm_params, - cached_content=None, - ) - - # 'cache' must be stripped — Vertex AI has no such field - assert "cache" not in result, ( - "extra_body.cache must not be forwarded to Vertex AI. " - 'Vertex AI rejects it with 400: Unknown name "cache": Cannot find field.' - ) - - # Other legitimate extra_body keys should still pass through - assert "some_vertex_param" in result - assert result["some_vertex_param"] == "value" - - # Core request fields must be present - assert "contents" in result - - -def test_extra_body_tags_not_forwarded_to_vertex_ai(): - """ - 'tags' inside extra_body is a LiteLLM-internal param for logging/tracking. - It must NOT be forwarded to the Vertex AI request body. - Documented in litellm_proxy.md: "Send tags by including them in the extra_body parameter" - """ - messages = [{"role": "user", "content": "test"}] - optional_params = { - "extra_body": { - "tags": ["user:alice", "env:prod"], - "custom_param": "allowed", - }, - } - litellm_params = {} - - result = _transform_request_body( - messages=messages, - model="gemini-2.5-pro", - optional_params=optional_params, - custom_llm_provider="vertex_ai", - litellm_params=litellm_params, - cached_content=None, - ) - - assert "tags" not in result - assert "custom_param" in result - assert result["custom_param"] == "allowed" - - -def test_extra_body_google_maps_rewrites_json_response_format(): - messages = [{"role": "user", "content": "test"}] - optional_params = { - "response_mime_type": "application/json", - "response_schema": { - "type": "object", - "properties": {"answer": {"type": "string"}}, - }, - "extra_body": { - "tools": [{"googleMaps": {}}], - }, - } - - result = _transform_request_body( - messages=messages, - model="gemini-2.5-pro", - optional_params=optional_params, - custom_llm_provider="vertex_ai", - litellm_params={}, - cached_content=None, - ) - - generation_config = result["generationConfig"] - assert "response_mime_type" not in generation_config - assert generation_config["responseFormat"] == { - "text": { - "mimeType": "APPLICATION_JSON", - "schema": { - "type": "object", - "properties": {"answer": {"type": "string"}}, - }, - } - } - - -def test_extra_body_generation_config_cannot_restore_google_maps_json_mime_type(): - messages = [{"role": "user", "content": "test"}] - optional_params = { - "tools": [{"googleMaps": {}}], - "response_mime_type": "application/json", - "extra_body": { - "generationConfig": { - "response_mime_type": "application/json", - "response_json_schema": { - "type": "object", - "properties": {"answer": {"type": "string"}}, - }, - }, - }, - } - - result = _transform_request_body( - messages=messages, - model="gemini-2.5-pro", - optional_params=optional_params, - custom_llm_provider="vertex_ai", - litellm_params={}, - cached_content=None, - ) - - generation_config = result["generationConfig"] - assert "response_mime_type" not in generation_config - assert "response_json_schema" not in generation_config - assert generation_config["responseFormat"] == { - "text": { - "mimeType": "APPLICATION_JSON", - "schema": { - "type": "object", - "properties": {"answer": {"type": "string"}}, - }, - } - } - - -def test_metadata_to_labels_vertex_only(): - """Test that metadata->labels conversion only happens for Vertex AI""" - messages = [{"role": "user", "content": "test"}] - optional_params = {} - litellm_params = { - "metadata": { - "requester_metadata": {"user": "john_doe", "project": "test-project"} - } - } - - # Google GenAI/AI Studio should not include labels from metadata - result = _transform_request_body( - messages=messages, - model="gemini-2.5-pro", - optional_params=optional_params.copy(), - custom_llm_provider="gemini", - litellm_params=litellm_params.copy(), - cached_content=None, - ) - assert "labels" not in result - - # Vertex AI should include labels from metadata - result = _transform_request_body( - messages=messages, - model="gemini-2.5-pro", - optional_params=optional_params.copy(), - custom_llm_provider="vertex_ai", - litellm_params=litellm_params.copy(), - cached_content=None, - ) - assert "labels" in result - assert result["labels"] == {"user": "john_doe", "project": "test-project"} - - -def test_empty_content_handling(): - """Test that empty content strings are properly handled in Gemini message transformation""" - # Test with empty content in user message - messages = [{"content": "", "role": "user"}] - - contents = _gemini_convert_messages_with_history(messages=messages) - - # Verify that the content was properly transformed - assert len(contents) == 1 - assert contents[0]["role"] == "user" - assert len(contents[0]["parts"]) == 1 - assert "text" in contents[0]["parts"][0] - assert contents[0]["parts"][0]["text"] == "" - - -def test_thought_signature_extraction_from_response(): - """Test that thought signatures are extracted from Gemini response parts and stored in provider_specific_fields""" - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - from litellm.types.llms.vertex_ai import HttpxPartType - - # Test case: Single function call with thought signature - test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" - - parts_with_signature = [ - HttpxPartType( - functionCall={ - "name": "get_current_temperature", - "args": {"location": "Paris"}, - }, - thoughtSignature=test_signature, - ) - ] - - function, tools, _ = VertexGeminiConfig._transform_parts( - parts=parts_with_signature, - cumulative_tool_call_idx=0, - is_function_call=False, - ) - - # Verify thought signature is stored in provider_specific_fields - assert tools is not None - assert len(tools) == 1 - assert "provider_specific_fields" in tools[0] - assert tools[0]["provider_specific_fields"]["thought_signature"] == test_signature - - -def test_thought_signature_parallel_function_calls(): - """Test that only the first function call in parallel calls has thought signature""" - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - from litellm.types.llms.vertex_ai import HttpxPartType - - test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" - - # Parallel function calls - only first has signature - parts_parallel = [ - HttpxPartType( - functionCall={ - "name": "get_current_temperature", - "args": {"location": "Paris"}, - }, - thoughtSignature=test_signature, # First FC has signature - ), - HttpxPartType( - functionCall={ - "name": "get_current_temperature", - "args": {"location": "London"}, - }, - # Second FC has no signature (parallel call) - ), - ] - - function, tools, _ = VertexGeminiConfig._transform_parts( - parts=parts_parallel, - cumulative_tool_call_idx=0, - is_function_call=False, - ) - - # Verify only first tool call has thought signature - assert tools is not None - assert len(tools) == 2 - assert "provider_specific_fields" in tools[0] - assert tools[0]["provider_specific_fields"]["thought_signature"] == test_signature - # Second tool call should not have thought signature - assert "provider_specific_fields" not in tools[ - 1 - ] or "thought_signature" not in tools[1].get("provider_specific_fields", {}) - - -def test_thought_signature_preservation_in_conversion(): - """Test that thought signatures are preserved when converting assistant messages back to Gemini format""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" - - # Assistant message with tool calls containing thought signatures - assistant_message = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_abc123", - "type": "function", - "function": { - "name": "get_current_temperature", - "arguments": '{"location": "Paris"}', - }, - "index": 0, - "provider_specific_fields": { - "thought_signature": test_signature, - }, - }, - { - "id": "call_def456", - "type": "function", - "function": { - "name": "get_current_temperature", - "arguments": '{"location": "London"}', - }, - "index": 1, - # No thought signature for parallel call - }, - ], - } - - gemini_parts = convert_to_gemini_tool_call_invoke(assistant_message) - - # Verify thought signature is preserved in first function call part - assert len(gemini_parts) == 2 - assert "function_call" in gemini_parts[0] - assert "thoughtSignature" in gemini_parts[0] - assert gemini_parts[0]["thoughtSignature"] == test_signature - - # Verify second function call part does not have thought signature - assert "function_call" in gemini_parts[1] - assert "thoughtSignature" not in gemini_parts[1] - - -def test_thought_signature_sequential_function_calls(): - """Test that each sequential function call preserves its own thought signature""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - signature_1 = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" - signature_2 = "DifferentSignatureForSecondCall1234567890ABCDEFGHIJKLMNOPQRSTUVWXYZ" - - # Sequential function calls - each has its own signature - # This simulates a multi-step conversation where each step has a signature - assistant_message_step1 = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_step1", - "type": "function", - "function": { - "name": "check_flight", - "arguments": '{"flight": "AA100"}', - }, - "index": 0, - "provider_specific_fields": { - "thought_signature": signature_1, - }, - }, - ], - } - - assistant_message_step2 = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_step2", - "type": "function", - "function": { - "name": "book_taxi", - "arguments": '{"destination": "airport"}', - }, - "index": 0, - "provider_specific_fields": { - "thought_signature": signature_2, - }, - }, - ], - } - - gemini_parts_step1 = convert_to_gemini_tool_call_invoke(assistant_message_step1) - gemini_parts_step2 = convert_to_gemini_tool_call_invoke(assistant_message_step2) - - # Verify each step preserves its own signature - assert len(gemini_parts_step1) == 1 - assert gemini_parts_step1[0]["thoughtSignature"] == signature_1 - - assert len(gemini_parts_step2) == 1 - assert gemini_parts_step2[0]["thoughtSignature"] == signature_2 - - -def test_thought_signature_with_function_call_mode(): - """Test thought signature extraction in function_call mode (is_function_call=True)""" - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - from litellm.types.llms.vertex_ai import HttpxPartType - - test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" - - parts_with_signature = [ - HttpxPartType( - functionCall={ - "name": "get_current_weather", - "args": {"location": "Tokyo"}, - }, - thoughtSignature=test_signature, - ) - ] - - function, tools, _ = VertexGeminiConfig._transform_parts( - parts=parts_with_signature, - cumulative_tool_call_idx=0, - is_function_call=True, - ) - - # Verify thought signature is stored in function's provider_specific_fields - assert function is not None - # Function should be dict-like (TypedDict or dict) - assert hasattr(function, "__getitem__") or isinstance(function, dict) - assert "provider_specific_fields" in function - assert function["provider_specific_fields"]["thought_signature"] == test_signature - assert tools is None - - -def test_dummy_signature_added_for_gemini_3_conversation_history(): - """Test that dummy signatures are added when transferring conversation history from older models (like gemini-2.5-flash) to gemini-3.""" - import base64 - - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - # Simulate conversation history from gemini-2.5-flash (no thought signature) - assistant_message_from_older_model = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_abc123", - "type": "function", - "function": { - "name": "get_current_temperature", - "arguments": '{"location": "Paris"}', - }, - "index": 0, - # No provider_specific_fields - older model doesn't provide signatures - }, - ], - } - - # Convert to Gemini format for gemini-3-pro-preview (should add dummy signature) - gemini_parts = convert_to_gemini_tool_call_invoke( - assistant_message_from_older_model, model="gemini-3-pro-preview" - ) - - # Verify dummy signature is added - assert len(gemini_parts) == 1 - assert "function_call" in gemini_parts[0] - assert "thoughtSignature" in gemini_parts[0] - - # Verify it's the expected dummy signature (base64 encoded "skip_thought_signature_validator") - expected_dummy = base64.b64encode(b"skip_thought_signature_validator").decode( - "utf-8" - ) - assert gemini_parts[0]["thoughtSignature"] == expected_dummy - - -def test_dummy_signature_not_added_for_gemini_2_5(): - """Test that dummy signatures are NOT added when target model is not gemini-3.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - # Simulate conversation history from gemini-2.5-flash (no thought signature) - assistant_message = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_abc123", - "type": "function", - "function": { - "name": "get_current_temperature", - "arguments": '{"location": "Paris"}', - }, - "index": 0, - # No provider_specific_fields - }, - ], - } - - # Convert to Gemini format for gemini-2.5-flash (should NOT add dummy signature) - gemini_parts = convert_to_gemini_tool_call_invoke( - assistant_message, model="gemini-2.5-flash" - ) - - # Verify no dummy signature is added for non-gemini-3 models - assert len(gemini_parts) == 1 - assert "function_call" in gemini_parts[0] - assert "thoughtSignature" not in gemini_parts[0] - - -def test_dummy_signature_not_added_when_signature_exists(): - """Test that dummy signatures are NOT added when a real signature already exists.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - real_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" - - # Assistant message with existing thought signature - assistant_message_with_signature = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_abc123", - "type": "function", - "function": { - "name": "get_current_temperature", - "arguments": '{"location": "Paris"}', - "provider_specific_fields": { - "thought_signature": real_signature, - }, - }, - "index": 0, - }, - ], - } - - # Convert to Gemini format for gemini-3-pro-preview - gemini_parts = convert_to_gemini_tool_call_invoke( - assistant_message_with_signature, model="gemini-3-pro-preview" - ) - - # Verify real signature is preserved, not replaced with dummy - assert len(gemini_parts) == 1 - assert "function_call" in gemini_parts[0] - assert "thoughtSignature" in gemini_parts[0] - assert gemini_parts[0]["thoughtSignature"] == real_signature - - -def test_dummy_signature_with_function_call_mode(): - """Test that dummy signatures are added for function_call mode when converting to gemini-3.""" - import base64 - - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - # Assistant message with function_call (not tool_calls) and no signature - assistant_message_function_call = { - "role": "assistant", - "content": None, - "function_call": { - "name": "get_current_temperature", - "arguments": '{"location": "Paris"}', - # No provider_specific_fields - }, - } - - # Convert to Gemini format for gemini-3-pro-preview - gemini_parts = convert_to_gemini_tool_call_invoke( - assistant_message_function_call, model="gemini-3-pro-preview" - ) - - # Verify dummy signature is added - assert len(gemini_parts) == 1 - assert "function_call" in gemini_parts[0] - assert "thoughtSignature" in gemini_parts[0] - - # Verify it's the expected dummy signature - expected_dummy = base64.b64encode(b"skip_thought_signature_validator").decode( - "utf-8" - ) - assert gemini_parts[0]["thoughtSignature"] == expected_dummy - - -def _parallel_tool_calls(*signatures): - return [ - { - "id": f"call_{idx}", - "type": "function", - "function": { - "name": f"tool_{idx}", - "arguments": '{"location": "Paris"}', - **( - {"provider_specific_fields": {"thought_signature": signature}} - if signature is not None - else {} - ), - }, - "index": idx, - } - for idx, signature in enumerate(signatures) - ] - - -def _parallel_tool_calls_signed_via_id(*signatures): - """Parallel tool calls in the shape LiteLLM actually hands back to clients. - - The signature rides in the tool call id behind __thought__, which is what an - OpenAI-format client echoes back on the next turn. - """ - from litellm.litellm_core_utils.prompt_templates.factory import ( - _encode_tool_call_id_with_signature, - ) - - return [ - { - "id": _encode_tool_call_id_with_signature(f"call_{idx}", signature), - "type": "function", - "function": {"name": f"tool_{idx}", "arguments": '{"location": "Paris"}'}, - "index": idx, - } - for idx, signature in enumerate(signatures) - ] - - -REAL_THOUGHT_SIGNATURE = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n" -PLACEHOLDER_SIGNATURE = base64.b64encode(b"skip_thought_signature_validator").decode( - "utf-8" -) - - -def test_dummy_signature_only_on_first_parallel_tool_call(): - """Google documents the placeholder as a last resort that degrades quality, so an unsigned - parallel turn replayed to gemini-3 gets a budget of exactly one.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - gemini_parts = convert_to_gemini_tool_call_invoke( - { - "role": "assistant", - "content": None, - "tool_calls": _parallel_tool_calls(None, None, None), - }, - model="gemini-3-pro-preview", - ) - - assert len(gemini_parts) == 3 - assert gemini_parts[0]["thoughtSignature"] == PLACEHOLDER_SIGNATURE - assert "thoughtSignature" not in gemini_parts[1] - assert "thoughtSignature" not in gemini_parts[2] - - -def test_real_signature_on_first_parallel_tool_call_leaves_siblings_empty(): - """Gemini signs only the first of N parallel function calls, so a faithful replay has - nothing to attach to the siblings.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - gemini_parts = convert_to_gemini_tool_call_invoke( - { - "role": "assistant", - "content": None, - "tool_calls": _parallel_tool_calls(REAL_THOUGHT_SIGNATURE, None, None), - }, - model="gemini-3-pro-preview", - ) - - assert len(gemini_parts) == 3 - assert gemini_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE - assert "thoughtSignature" not in gemini_parts[1] - assert "thoughtSignature" not in gemini_parts[2] - - -def test_real_signature_on_later_parallel_tool_call_is_preserved(): - """Clients may reorder or drop calls, so a signature that lands on a non-first call is - still the model's own and must survive the round trip.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - gemini_parts = convert_to_gemini_tool_call_invoke( - { - "role": "assistant", - "content": None, - "tool_calls": _parallel_tool_calls(None, REAL_THOUGHT_SIGNATURE), - }, - model="gemini-3-pro-preview", - ) - - assert len(gemini_parts) == 2 - assert gemini_parts[0]["thoughtSignature"] == PLACEHOLDER_SIGNATURE - assert gemini_parts[1]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE - - -def test_no_signatures_on_parallel_tool_calls_for_gemini_2_5(): - """Non-gemini-3 models never get a placeholder signature, on any call.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - gemini_parts = convert_to_gemini_tool_call_invoke( - { - "role": "assistant", - "content": None, - "tool_calls": _parallel_tool_calls(None, None), - }, - model="gemini-2.5-flash", - ) - - assert len(gemini_parts) == 2 - assert all("thoughtSignature" not in part for part in gemini_parts) - - -def test_signature_embedded_in_tool_call_id_only_on_first_parallel_call(): - """The production shape: the signature arrives inside the first call's id, siblings have bare ids.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - gemini_parts = convert_to_gemini_tool_call_invoke( - { - "role": "assistant", - "content": None, - "tool_calls": _parallel_tool_calls_signed_via_id( - REAL_THOUGHT_SIGNATURE, None, None - ), - }, - model="gemini-3-pro-preview", - ) - - assert len(gemini_parts) == 3 - assert gemini_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE - assert "thoughtSignature" not in gemini_parts[1] - assert "thoughtSignature" not in gemini_parts[2] - - -def test_tool_level_provider_specific_fields_signature_leaves_siblings_empty(): - """A signature on the tool call itself, rather than on its function, behaves the same way.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - tool_calls = _parallel_tool_calls(None, None) - tool_calls[0]["provider_specific_fields"] = { - "thought_signature": REAL_THOUGHT_SIGNATURE - } - - gemini_parts = convert_to_gemini_tool_call_invoke( - {"role": "assistant", "content": None, "tool_calls": tool_calls}, - model="gemini-3-pro-preview", - ) - - assert len(gemini_parts) == 2 - assert gemini_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE - assert "thoughtSignature" not in gemini_parts[1] - - -def test_placeholder_lands_on_first_emitted_part_not_first_tool_call_entry(): - """A non-function entry (e.g. an OpenAI custom tool call) emits no part, so it must not - consume the one placeholder slot and leave the real first function call bare.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - tool_calls = [ - {"id": "call_custom", "type": "custom", "custom": {"name": "noop", "input": ""}} - ] + _parallel_tool_calls(None, None) - - gemini_parts = convert_to_gemini_tool_call_invoke( - {"role": "assistant", "content": None, "tool_calls": tool_calls}, - model="gemini-3-pro-preview", - ) - - assert len(gemini_parts) == 2 - assert gemini_parts[0]["thoughtSignature"] == PLACEHOLDER_SIGNATURE - assert "thoughtSignature" not in gemini_parts[1] - - -def test_no_placeholder_when_model_is_unknown(): - """Without a model there is nothing to prove the target needs a placeholder, so none is added.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - gemini_parts = convert_to_gemini_tool_call_invoke( - { - "role": "assistant", - "content": None, - "tool_calls": _parallel_tool_calls(None, None), - }, - ) - - assert len(gemini_parts) == 2 - assert all("thoughtSignature" not in part for part in gemini_parts) - - -def test_real_signature_forwarded_to_gemini_2_5_without_placeholder_siblings(): - """Older models still receive a real signature that a client replays, and still get no placeholder.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - gemini_parts = convert_to_gemini_tool_call_invoke( - { - "role": "assistant", - "content": None, - "tool_calls": _parallel_tool_calls(REAL_THOUGHT_SIGNATURE, None), - }, - model="gemini-2.5-flash", - ) - - assert len(gemini_parts) == 2 - assert gemini_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE - assert "thoughtSignature" not in gemini_parts[1] - - -def test_parallel_tool_call_history_replayed_through_full_message_conversion(): - """End to end through the message-history converter, the path a real /chat/completions replay takes.""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - messages = [ - {"role": "user", "content": "Weather in Paris, London and Tokyo?"}, - { - "role": "assistant", - "content": None, - "tool_calls": _parallel_tool_calls_signed_via_id( - REAL_THOUGHT_SIGNATURE, None, None - ), - }, - ] - - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-3-pro-preview" - ) - - model_parts = contents[1]["parts"] - assert len(model_parts) == 3 - assert model_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE - assert "thoughtSignature" not in model_parts[1] - assert "thoughtSignature" not in model_parts[2] - - -@pytest.mark.parametrize( - "model", - ["gemini-3.5-flash", "vertex_ai/gemini-3.5-flash", "gemini/gemini-3.5-flash"], -) -def test_natively_signed_parallel_turn_never_carries_a_placeholder(model): - """A native gemini-3.5 parallel turn replays with zero skip_thought_signature_validator parts. - - Fabricating the placeholder alongside a real signature is what produced empty text responses - on gemini-3.5 parallel function calling, so the whole payload has to stay placeholder-free. - """ - import json - - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - messages = [ - {"role": "user", "content": "Weather in Paris, London and Tokyo?"}, - { - "role": "assistant", - "content": None, - "tool_calls": _parallel_tool_calls_signed_via_id( - REAL_THOUGHT_SIGNATURE, None, None - ), - }, - ] - - contents = _gemini_convert_messages_with_history(messages=messages, model=model) - - model_parts = contents[1]["parts"] - assert len(model_parts) == 3 - assert model_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE - assert "thoughtSignature" not in model_parts[1] - assert "thoughtSignature" not in model_parts[2] - assert PLACEHOLDER_SIGNATURE not in json.dumps(contents) - - -@pytest.mark.parametrize( - "model", - [ - "gemini-3-pro-preview", - "gemini-3-flash-preview", - "gemini-3.1-pro-preview", - "gemini-3.5-flash", - "gemini-3.6-flash", - "gemini-3.7-flash", - "gemini-3.8-flash", - "vertex_ai/gemini-3.5-flash", - "vertex_ai/gemini-3.7-flash", - "vertex_ai/gemini-3.8-flash", - "gemini/gemini-3.5-flash", - "gemini/gemini-3.7-flash", - "gemini/gemini-3.8-flash", - ], -) -def test_placeholder_scoped_to_first_call_across_gemini_3_variants(model): - """The gemini-3 gate is a substring match, so every family member and prefix form has to - land on the same one-placeholder budget rather than only the versions we happened to try.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_invoke, - ) - - gemini_parts = convert_to_gemini_tool_call_invoke( - { - "role": "assistant", - "content": None, - "tool_calls": _parallel_tool_calls(None, None, None), - }, - model=model, - ) - - assert len(gemini_parts) == 3 - assert gemini_parts[0]["thoughtSignature"] == PLACEHOLDER_SIGNATURE - assert "thoughtSignature" not in gemini_parts[1] - assert "thoughtSignature" not in gemini_parts[2] - - -def test_signed_text_part_survives_alongside_unsigned_parallel_tool_calls(): - """Text-part and function-call signatures are collected by separate code paths, so scoping the - placeholder must not disturb a real signature that arrived on the text part.""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - msg = { - "role": "assistant", - "content": "Checking all three cities.", - "provider_specific_fields": {"thought_signatures": ["real_25_signature"]}, - "tool_calls": _parallel_tool_calls(None, None, None), - } - - parts = _gemini_convert_messages_with_history( - messages=[msg], model="gemini-3-pro-preview" - )[0]["parts"] - - assert parts[0]["text"] == "Checking all three cities." - assert parts[0]["thoughtSignature"] == "real_25_signature" - assert parts[1]["thoughtSignature"] == PLACEHOLDER_SIGNATURE - assert "thoughtSignature" not in parts[2] - assert "thoughtSignature" not in parts[3] - - -# Tests for media_resolution (detail parameter) handling - Issue #17084 -class TestMediaResolution: - """Tests for media_resolution handling in Gemini 2.x models""" - - def test_get_highest_media_resolution_high_wins(self): - """Test that 'high' resolution takes precedence over 'low'""" - assert _get_highest_media_resolution("low", "high") == "high" - assert _get_highest_media_resolution("high", "low") == "high" - assert _get_highest_media_resolution(None, "high") == "high" - assert _get_highest_media_resolution("high", None) == "high" - - def test_get_highest_media_resolution_low_over_none(self): - """Test that 'low' resolution takes precedence over None""" - assert _get_highest_media_resolution(None, "low") == "low" - assert _get_highest_media_resolution("low", None) == "low" - - def test_get_highest_media_resolution_same_values(self): - """Test handling of same resolution values""" - assert _get_highest_media_resolution("high", "high") == "high" - assert _get_highest_media_resolution("low", "low") == "low" - assert _get_highest_media_resolution(None, None) is None - - def test_get_highest_media_resolution_medium(self): - """Test that 'medium' resolution is correctly ranked between 'low' and 'high'""" - assert _get_highest_media_resolution("low", "medium") == "medium" - assert _get_highest_media_resolution("medium", "low") == "medium" - assert _get_highest_media_resolution("medium", "high") == "high" - assert _get_highest_media_resolution("high", "medium") == "high" - assert _get_highest_media_resolution(None, "medium") == "medium" - assert _get_highest_media_resolution("medium", None) == "medium" - - def test_get_highest_media_resolution_ultra_high(self): - """Test that 'ultra_high' resolution takes precedence over all others""" - assert _get_highest_media_resolution("high", "ultra_high") == "ultra_high" - assert _get_highest_media_resolution("ultra_high", "high") == "ultra_high" - assert _get_highest_media_resolution("medium", "ultra_high") == "ultra_high" - assert _get_highest_media_resolution("low", "ultra_high") == "ultra_high" - assert _get_highest_media_resolution(None, "ultra_high") == "ultra_high" - assert _get_highest_media_resolution("ultra_high", None) == "ultra_high" - - def test_extract_max_media_resolution_single_image_high(self): - """Test extraction of media resolution from single image with detail=high""" - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - { - "type": "image_url", - "image_url": { - "url": "data:image/png;base64,abc123", - "detail": "high", - }, - }, - ], - } - ] - assert _extract_max_media_resolution_from_messages(messages) == "high" - - def test_extract_max_media_resolution_single_image_low(self): - """Test extraction of media resolution from single image with detail=low""" - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - { - "type": "image_url", - "image_url": { - "url": "data:image/png;base64,abc123", - "detail": "low", - }, - }, - ], - } - ] - assert _extract_max_media_resolution_from_messages(messages) == "low" - - def test_extract_max_media_resolution_no_detail(self): - """Test extraction when no detail parameter is provided""" - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - { - "type": "image_url", - "image_url": {"url": "data:image/png;base64,abc123"}, - }, - ], - } - ] - assert _extract_max_media_resolution_from_messages(messages) is None - - def test_extract_max_media_resolution_multiple_images_mixed(self): - """Test that highest resolution is returned when multiple images have different details""" - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "Compare these images"}, - { - "type": "image_url", - "image_url": { - "url": "data:image/png;base64,abc123", - "detail": "low", - }, - }, - { - "type": "image_url", - "image_url": { - "url": "data:image/png;base64,def456", - "detail": "high", - }, - }, - ], - } - ] - assert _extract_max_media_resolution_from_messages(messages) == "high" - - def test_extract_max_media_resolution_text_only(self): - """Test extraction from messages with no images""" - messages = [ - {"role": "user", "content": "Hello, how are you?"}, - {"role": "assistant", "content": "I'm doing well!"}, - ] - assert _extract_max_media_resolution_from_messages(messages) is None - - def test_transform_request_body_gemini_2x_adds_media_resolution(self): - """Test that media_resolution is added to generationConfig for Gemini 2.x models""" - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - { - "type": "image_url", - "image_url": { - "url": "data:image/png;base64,iVBORw0KGgo=", - "detail": "high", - }, - }, - ], - } - ] - - result = _transform_request_body( - messages=messages, - model="gemini-2.5-flash", - optional_params={}, - custom_llm_provider="gemini", - litellm_params={}, - cached_content=None, - ) - - assert "generationConfig" in result - assert "mediaResolution" in result["generationConfig"] - assert result["generationConfig"]["mediaResolution"] == "MEDIA_RESOLUTION_HIGH" - - def test_transform_request_body_gemini_2x_low_resolution(self): - """Test that low media_resolution is correctly added for Gemini 2.x""" - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - { - "type": "image_url", - "image_url": { - "url": "data:image/png;base64,iVBORw0KGgo=", - "detail": "low", - }, - }, - ], - } - ] - - result = _transform_request_body( - messages=messages, - model="gemini-2.5-flash", - optional_params={}, - custom_llm_provider="gemini", - litellm_params={}, - cached_content=None, - ) - - assert "generationConfig" in result - assert "mediaResolution" in result["generationConfig"] - assert result["generationConfig"]["mediaResolution"] == "MEDIA_RESOLUTION_LOW" - - def test_transform_request_body_gemini_3_no_global_media_resolution(self): - """Test that Gemini 3 models don't add media_resolution to generationConfig (they use per-part)""" - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - { - "type": "image_url", - "image_url": { - "url": "data:image/png;base64,iVBORw0KGgo=", - "detail": "high", - }, - }, - ], - } - ] - - result = _transform_request_body( - messages=messages, - model="gemini-3-pro-preview", - optional_params={}, - custom_llm_provider="gemini", - litellm_params={}, - cached_content=None, - ) - - # Gemini 3 should NOT have mediaResolution in generationConfig - # (it's handled per-part in the content transformation) - if "generationConfig" in result: - assert "mediaResolution" not in result["generationConfig"] - - def test_transform_request_body_no_detail_no_media_resolution(self): - """Test that no mediaResolution is added when detail is not specified""" - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - { - "type": "image_url", - "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, - }, - ], - } - ] - - result = _transform_request_body( - messages=messages, - model="gemini-2.5-flash", - optional_params={}, - custom_llm_provider="gemini", - litellm_params={}, - cached_content=None, - ) - - # When no detail is specified, mediaResolution should not be in generationConfig - if "generationConfig" in result: - assert "mediaResolution" not in result["generationConfig"] - - def test_extract_max_media_resolution_file_type_with_detail(self): - """Test that detail is extracted from file content type, not just image_url""" - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is in this file?"}, - { - "type": "file", - "file": { - "url": "data:image/png;base64,abc123", - "detail": "high", - }, - }, - ], - } - ] - assert _extract_max_media_resolution_from_messages(messages) == "high" - - def test_extract_max_media_resolution_mixed_image_and_file(self): - """Test that highest detail is returned across both image_url and file types""" - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "Compare these"}, - { - "type": "image_url", - "image_url": { - "url": "data:image/png;base64,abc123", - "detail": "low", - }, - }, - { - "type": "file", - "file": { - "url": "data:image/png;base64,def456", - "detail": "high", - }, - }, - ], - } - ] - assert _extract_max_media_resolution_from_messages(messages) == "high" - - def test_transform_request_body_gemini_1x_no_media_resolution(self): - """Test that Gemini 1.x models don't get mediaResolution in generationConfig""" - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is this?"}, - { - "type": "image_url", - "image_url": { - "url": "data:image/png;base64,iVBORw0KGgo=", - "detail": "high", - }, - }, - ], - } - ] - - result = _transform_request_body( - messages=messages, - model="gemini-1.5-pro", - optional_params={}, - custom_llm_provider="gemini", - litellm_params={}, - cached_content=None, - ) - - # Gemini 1.x should NOT have mediaResolution (not supported) - if "generationConfig" in result: - assert "mediaResolution" not in result["generationConfig"] - - -# Tests for VideoMetadata support across all Gemini models (Issue #25474) -class TestVideoMetadataAllGeminiModels: - """Tests that video_metadata (fps, start_offset, end_offset) works for all Gemini models""" - - def _make_video_messages(self, video_metadata: dict) -> list: - return [ - { - "role": "user", - "content": [ - {"type": "text", "text": "Analyze this video"}, - { - "type": "file", - "file": { - "file_id": "gs://bucket/video.mp4", - "format": "video/mp4", - "video_metadata": video_metadata, - }, - }, - ], - } - ] - - def _get_file_part(self, contents: list) -> dict: - for part in contents[0]["parts"]: - if "file_data" in part: - return part - raise AssertionError("No file part found in contents") - - def test_video_metadata_fps_gemini_2_5_flash(self): - """Gemini 2.5 Flash: fps in video_metadata should be forwarded (Issue #25474)""" - messages = self._make_video_messages({"fps": 5}) - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-2.5-flash" - ) - file_part = self._get_file_part(contents) - assert "video_metadata" in file_part - assert file_part["video_metadata"]["fps"] == 5 - - def test_video_metadata_fps_gemini_2_5_pro(self): - """Gemini 2.5 Pro: fps in video_metadata should be forwarded (Issue #25474)""" - messages = self._make_video_messages({"fps": 10}) - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-2.5-pro" - ) - file_part = self._get_file_part(contents) - assert "video_metadata" in file_part - assert file_part["video_metadata"]["fps"] == 10 - - def test_video_metadata_offsets_gemini_2_5_flash(self): - """Gemini 2.5 Flash: start_offset/end_offset converted to camelCase (Issue #25474)""" - messages = self._make_video_messages( - {"start_offset": "5s", "end_offset": "30s"} - ) - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-2.5-flash" - ) - file_part = self._get_file_part(contents) - assert "video_metadata" in file_part - vm = file_part["video_metadata"] - assert vm["startOffset"] == "5s" - assert vm["endOffset"] == "30s" - - def test_video_metadata_all_fields_gemini_2_5_flash(self): - """Gemini 2.5 Flash: all video_metadata fields forwarded correctly (Issue #25474)""" - messages = self._make_video_messages( - {"fps": 5, "start_offset": "10s", "end_offset": "60s"} - ) - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-2.5-flash" - ) - file_part = self._get_file_part(contents) - assert "video_metadata" in file_part - vm = file_part["video_metadata"] - assert vm["fps"] == 5 - assert vm["startOffset"] == "10s" - assert vm["endOffset"] == "60s" - - def test_video_metadata_gemini_1_5_pro(self): - """Gemini 1.5 Pro: video_metadata should also be forwarded (Issue #25474)""" - messages = self._make_video_messages({"fps": 2}) - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-1.5-pro" - ) - file_part = self._get_file_part(contents) - assert "video_metadata" in file_part - assert file_part["video_metadata"]["fps"] == 2 - - -def test_convert_tool_response_with_base64_image(): - """Test tool response with base64 data URI image.""" - # Create a small test image (1x1 red pixel PNG) - test_image_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" - image_data_uri = f"data:image/png;base64,{test_image_base64}" - - # Create tool message with image - tool_message = { - "role": "tool", - "tool_call_id": "call_test123", - "content": [ - { - "type": "text", - "text": '{"url": "https://example.com", "status": "success"}', - }, - {"type": "input_image", "image_url": image_data_uri}, - ], - } - - # Mock last message with tool calls - last_message_with_tool_calls = { - "tool_calls": [ - { - "id": "call_test123", - "function": {"name": "click_at", "arguments": '{"x": 100, "y": 200}'}, - } - ] - } - - # Convert tool response with nested multimodal functionResponse.parts. - result = convert_to_gemini_tool_call_result( - tool_message, last_message_with_tool_calls - ) - - assert isinstance(result, list), "Should return a parts list when media is present" - assert len(result) == 1, "Should return one function_response part" - result_part = result[0] - assert "function_response" in result_part - assert "inline_data" not in result_part - function_response = result_part["function_response"] - assert function_response["name"] == "click_at" - assert "response" in function_response - # Verify JSON response is parsed correctly - assert "url" in function_response["response"] - assert function_response["response"]["url"] == "https://example.com" - - # Check inline_data is nested under functionResponse.parts. - assert "parts" in function_response - assert len(function_response["parts"]) == 1 - inline_data: BlobType = function_response["parts"][0]["inline_data"] - assert "data" in inline_data - assert "mime_type" in inline_data - assert inline_data["mime_type"] == "image/png" - assert inline_data["data"] == test_image_base64 - - -def test_gemini_history_nests_multimodal_tool_response_parts(): - """Full history conversion should not emit sibling inline_data tool result parts.""" - test_image_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" - messages = [ - {"role": "user", "content": "Get me an image"}, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_get_image", - "type": "function", - "function": {"name": "get_image", "arguments": "{}"}, - } - ], - }, - { - "role": "tool", - "tool_call_id": "call_get_image", - "content": [ - {"type": "text", "text": '{"image_ref": "inline"}'}, - { - "type": "image", - "source": { - "type": "base64", - "media_type": "image/png", - "data": test_image_base64, - }, - }, - ], - }, - ] - - contents = _gemini_convert_messages_with_history(messages=messages) - - tool_response_parts = contents[-1]["parts"] - assert len(tool_response_parts) == 1 - assert "inline_data" not in tool_response_parts[0] - function_response = tool_response_parts[0]["function_response"] - assert function_response["parts"] == [ - { - "inline_data": { - "data": test_image_base64, - "mime_type": "image/png", - } - } - ] def test_convert_tool_response_with_url_image(): """Test tool response with HTTP URL image (will download and convert).""" - import pytest - # Use a publicly accessible test image URL test_image_url = "https://via.placeholder.com/1x1.png" @@ -1701,13 +33,9 @@ def test_convert_tool_response_with_url_image(): } try: - result = convert_to_gemini_tool_call_result( - tool_message, last_message_with_tool_calls - ) + result = convert_to_gemini_tool_call_result(tool_message, last_message_with_tool_calls) - assert isinstance( - result, list - ), "Should return a parts list when media is present" + assert isinstance(result, list), "Should return a parts list when media is present" assert len(result) == 1, "Should return one function_response part" result_part = result[0] assert "function_response" in result_part @@ -1724,1060 +52,3 @@ def test_convert_tool_response_with_url_image(): except Exception as e: # Skip test if URL download fails (no internet connection, etc.) pytest.skip(f"Failed to download image from URL: {e}") - - -def test_convert_tool_response_text_only(): - """Test tool response with only text (no image).""" - tool_message = { - "role": "tool", - "tool_call_id": "call_test789", - "content": [ - {"type": "text", "text": '{"status": "completed", "result": "success"}'} - ], - } - - last_message_with_tool_calls = { - "tool_calls": [ - { - "id": "call_test789", - "function": {"name": "wait_5_seconds", "arguments": "{}"}, - } - ] - } - - result = convert_to_gemini_tool_call_result( - tool_message, last_message_with_tool_calls - ) - - # Should be a single part (no list) when no image - assert not isinstance(result, list), "Should return single part when no image" - - # Check function_response exists - assert "function_response" in result - function_response = result["function_response"] - assert function_response["name"] == "wait_5_seconds" - # Verify JSON response is parsed correctly - assert "status" in function_response["response"] - assert function_response["response"]["status"] == "completed" - - # Check inline_data does NOT exist (no image provided) - assert "inline_data" not in result - - -def test_file_data_field_order(): - """ - Test that file_data fields are in the correct order (mime_type before file_uri). - - The Gemini API is sensitive to field order in the file_data object. - This test verifies that mime_type comes before file_uri in both: - 1. Dictionary key order - 2. JSON serialization - - Related issue: Gemini API returns 400 INVALID_ARGUMENT when fields are in wrong order. - """ - import json - - from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_media - - # Test with HTTPS URL and explicit format (audio file) - file_url = "https://generativelanguage.googleapis.com/v1beta/files/test123" - format = "audio/mpeg" - - result = _process_gemini_media(image_url=file_url, format=format) - - # Verify the result has file_data - assert "file_data" in result - file_data = result["file_data"] - - # Verify both fields are present - assert "mime_type" in file_data - assert "file_uri" in file_data - assert file_data["mime_type"] == "audio/mpeg" - assert file_data["file_uri"] == file_url - - # Verify field order by checking dictionary keys - # In Python 3.7+, dict maintains insertion order - file_data_keys = list(file_data.keys()) - assert file_data_keys.index("mime_type") < file_data_keys.index( - "file_uri" - ), "mime_type must come before file_uri in the file_data dict" - - # Also verify by serializing to JSON string - json_str = json.dumps(file_data) - mime_type_pos = json_str.find('"mime_type"') - file_uri_pos = json_str.find('"file_uri"') - assert ( - mime_type_pos < file_uri_pos - ), "mime_type must appear before file_uri in JSON serialization" - - -def test_file_data_field_order_gcs_urls(): - """Test that GCS URLs also maintain correct field order.""" - import json - - from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_media - - # Test with GCS URL - gcs_url = "gs://bucket/audio.mp3" - - result = _process_gemini_media(image_url=gcs_url) - - # Verify the result has file_data - assert "file_data" in result - file_data = result["file_data"] - - # Verify both fields are present - assert "mime_type" in file_data - assert "file_uri" in file_data - - # Verify field order - file_data_keys = list(file_data.keys()) - assert file_data_keys.index("mime_type") < file_data_keys.index( - "file_uri" - ), "mime_type must come before file_uri in the file_data dict" - - -def test_gemini_files_api_uri_without_format(): - """ - Test that Gemini Files API URIs work WITHOUT an explicit format/mime_type. - - When a user uploads a file via the Gemini Files API and then references it - by URI (https://generativelanguage.googleapis.com/v1beta/files/...), - the file is already on Google's servers. These URLs return 403 when - fetched directly, so _process_gemini_media must NOT try to resolve the - MIME type via HTTP. Instead it should pass the URI through as file_data - and let the Gemini API resolve the type from its stored metadata. - - Related issue: https://github.com/BerriAI/litellm/issues/24907 - """ - from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_media - - file_url = "https://generativelanguage.googleapis.com/v1beta/files/37eh7rsw1vfe" - - # Should NOT raise — previously this hit the generic https:// handler - # which called _get_image_mime_type_from_url() and got a 403. - result = _process_gemini_media(image_url=file_url) - - assert "file_data" in result - file_data = result["file_data"] - assert file_data["file_uri"] == file_url - # When no format is provided, mime_type should be absent so the - # Gemini API infers it from the stored file metadata. - assert "mime_type" not in file_data - - -def test_gemini_files_api_uri_with_format(): - """ - Test that Gemini Files API URIs correctly forward an explicit format. - - Related issue: https://github.com/BerriAI/litellm/issues/24907 - """ - from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_media - - file_url = "https://generativelanguage.googleapis.com/v1beta/files/n1vhxa28lyaw" - - result = _process_gemini_media(image_url=file_url, format="text/plain") - - assert "file_data" in result - file_data = result["file_data"] - assert file_data["file_uri"] == file_url - assert file_data["mime_type"] == "text/plain" - - -def test_extract_file_data_with_path_object(): - """ - Test that filename is correctly extracted from Path objects for MIME type detection. - - When uploading files using Path objects (e.g., Path("speech.mp3")), the filename - must be extracted to enable proper MIME type detection. Without this, files get - uploaded with 'application/octet-stream' instead of the correct MIME type. - - Related issue: Files uploaded with wrong MIME type cause Gemini API to reject - requests where the specified format doesn't match the uploaded file's MIME type. - """ - import os - import tempfile - from pathlib import Path - - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - extract_file_data, - ) - - # Create a temporary MP3 file - with tempfile.NamedTemporaryFile(suffix=".mp3", delete=False) as tmp: - tmp.write(b"fake mp3 content") - tmp_path = tmp.name - - try: - # Test with Path object - path_obj = Path(tmp_path) - extracted = extract_file_data(path_obj) - - # Verify filename was extracted - assert extracted["filename"] is not None - assert extracted["filename"].endswith(".mp3") - - # Verify MIME type was correctly detected - assert ( - extracted["content_type"] == "audio/mpeg" - ), f"Expected 'audio/mpeg' but got '{extracted['content_type']}'" - - # Verify content was read - assert extracted["content"] == b"fake mp3 content" - - finally: - # Clean up temporary file - os.unlink(tmp_path) - - -def test_extract_file_data_with_pathlib_path(): - """Test that filename is correctly extracted from pathlib.Path inputs. - Bare str paths are rejected — when this runs in a proxy request handler - the value is attacker-controlled and opening it as a path is an LFI.""" - import os - import tempfile - from pathlib import Path - - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - extract_file_data, - ) - - with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp: - tmp.write(b"fake wav content") - tmp_path = Path(tmp.name) - - try: - extracted = extract_file_data(tmp_path) - - assert extracted["filename"] is not None - assert extracted["filename"].endswith(".wav") - assert extracted["content_type"] in [ - "audio/wav", - "audio/x-wav", - ], f"Expected 'audio/wav' or 'audio/x-wav' but got '{extracted['content_type']}'" - assert extracted["content"] == b"fake wav content" - finally: - os.unlink(str(tmp_path)) - - -def test_extract_file_data_with_tuple_format(): - """Test that tuple format (with explicit content_type) still works correctly.""" - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - extract_file_data, - ) - - # Test with tuple format: (filename, content, content_type) - filename = "test_audio.mp3" - content = b"test audio content" - content_type = "audio/mpeg" - - extracted = extract_file_data((filename, content, content_type)) - - # Verify all fields are correct - assert extracted["filename"] == filename - assert extracted["content"] == content - assert extracted["content_type"] == content_type - - -def test_extract_file_data_fallback_to_octet_stream(): - """Unknown file types fall back to application/octet-stream.""" - import os - import tempfile - from pathlib import Path - - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - extract_file_data, - ) - - with tempfile.NamedTemporaryFile(suffix=".xyz123", delete=False) as tmp: - tmp.write(b"unknown content") - tmp_path = Path(tmp.name) - - try: - extracted = extract_file_data(tmp_path) - - assert extracted["filename"] is not None - assert extracted["filename"].endswith(".xyz123") - assert ( - extracted["content_type"] == "application/octet-stream" - ), f"Expected 'application/octet-stream' for unknown type, got '{extracted['content_type']}'" - finally: - os.unlink(str(tmp_path)) - - -def test_convert_tool_response_with_pdf_file(): - """Test tool response with PDF file content using file_data field.""" - # Create a minimal test PDF (base64 encoded) - test_pdf_base64 = "JVBERi0xLjQKJeLjz9MKMSAwIG9iago8PC9UeXBlL0NhdGFsb2cvUGFnZXMgMiAwIFI+PgplbmRvYmoKdHJhaWxlcgo8PC9TaXplIDQvUm9vdCAxIDAgUj4+CnN0YXJ0eHJlZgoyMTYKJSVFT0Y=" - file_data_uri = f"data:application/pdf;base64,{test_pdf_base64}" - - # Create tool message with file - tool_message = { - "role": "tool", - "tool_call_id": "call_pdf_test", - "content": [ - {"type": "text", "text": '{"status": "success", "pages": 1}'}, - {"type": "file", "file_data": file_data_uri}, - ], - } - - # Mock last message with tool calls - last_message_with_tool_calls = { - "tool_calls": [ - { - "id": "call_pdf_test", - "function": { - "name": "analyze_document", - "arguments": '{"path": "/tmp/doc.pdf"}', - }, - } - ] - } - - # Convert tool response with nested multimodal functionResponse.parts. - result = convert_to_gemini_tool_call_result( - tool_message, last_message_with_tool_calls - ) - - assert isinstance(result, list), "Should return a parts list when media is present" - assert len(result) == 1, "Should return one function_response part" - result_part = result[0] - assert "function_response" in result_part - assert "inline_data" not in result_part - function_response = result_part["function_response"] - assert function_response["name"] == "analyze_document" - assert "response" in function_response - # Verify JSON response is parsed correctly - assert "status" in function_response["response"] - assert function_response["response"]["status"] == "success" - - # Check inline_data is nested under functionResponse.parts. - assert "parts" in function_response - assert len(function_response["parts"]) == 1 - inline_data: BlobType = function_response["parts"][0]["inline_data"] - assert "data" in inline_data - assert "mime_type" in inline_data - assert inline_data["mime_type"] == "application/pdf" - assert inline_data["data"] == test_pdf_base64 - - -def test_convert_tool_response_with_input_file_type(): - """Test tool response with input_file content type (Responses API format).""" - # Create a minimal test PDF (base64 encoded) - test_pdf_base64 = "JVBERi0xLjQKJeLjz9MKMSAwIG9iago8PC9UeXBlL0NhdGFsb2cvUGFnZXMgMiAwIFI+PgplbmRvYmoKdHJhaWxlcgo8PC9TaXplIDQvUm9vdCAxIDAgUj4+CnN0YXJ0eHJlZgoyMTYKJSVFT0Y=" - file_data_uri = f"data:application/pdf;base64,{test_pdf_base64}" - - # Create tool message with input_file type - tool_message = { - "role": "tool", - "tool_call_id": "call_input_file_test", - "content": [{"type": "input_file", "file_data": file_data_uri}], - } - - # Mock last message with tool calls - last_message_with_tool_calls = { - "tool_calls": [ - { - "id": "call_input_file_test", - "function": {"name": "read_file", "arguments": "{}"}, - } - ] - } - - # Convert tool response - result = convert_to_gemini_tool_call_result( - tool_message, last_message_with_tool_calls - ) - - # Check inline_data is nested under functionResponse.parts. - assert isinstance(result, list), "Should return a parts list when media is present" - assert len(result) == 1, "Should return one function_response part" - function_response = result[0]["function_response"] - assert ( - function_response["parts"][0]["inline_data"]["mime_type"] == "application/pdf" - ) - - -def test_convert_tool_response_with_nested_file_object(): - """Test tool response with file content using nested file object format.""" - # Create a minimal test PDF (base64 encoded) - test_pdf_base64 = "JVBERi0xLjQKJeLjz9MKMSAwIG9iago8PC9UeXBlL0NhdGFsb2cvUGFnZXMgMiAwIFI+PgplbmRvYmoKdHJhaWxlcgo8PC9TaXplIDQvUm9vdCAxIDAgUj4+CnN0YXJ0eHJlZgoyMTYKJSVFT0Y=" - file_data_uri = f"data:application/pdf;base64,{test_pdf_base64}" - - # Create tool message with nested file object (OpenAI Agents SDK format) - tool_message = { - "role": "tool", - "tool_call_id": "call_nested_test", - "content": [{"type": "file", "file": {"file_data": file_data_uri}}], - } - - # Mock last message with tool calls - last_message_with_tool_calls = { - "tool_calls": [ - { - "id": "call_nested_test", - "function": {"name": "process_document", "arguments": "{}"}, - } - ] - } - - # Convert tool response - result = convert_to_gemini_tool_call_result( - tool_message, last_message_with_tool_calls - ) - - # Check inline_data is nested under functionResponse.parts. - assert isinstance(result, list), "Should return a parts list when media is present" - assert len(result) == 1, "Should return one function_response part" - function_response = result[0]["function_response"] - inline_data: BlobType = function_response["parts"][0]["inline_data"] - assert "data" in inline_data - assert "mime_type" in inline_data - assert inline_data["mime_type"] == "application/pdf" - assert inline_data["data"] == test_pdf_base64 - - -def test_assistant_message_with_images_field(): - """ - Test that assistant messages with images field are properly converted to Gemini format. - - This handles the case where an assistant message contains generated images in the - `images` field (e.g., from image generation models like gemini-2.5-flash-image). - The images should be converted to inline_data parts in the Gemini format. - """ - # Create a small test image (1x1 red pixel PNG) - test_image_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" - image_data_uri = f"data:image/png;base64,{test_image_base64}" - - # Create messages with assistant message containing images field - messages = [ - { - "role": "user", - "content": "Generate an image of a banana wearing a costume that says LiteLLM", - }, - { - "role": "assistant", - "content": "Here's your banana in a LiteLLM costume!", - "images": [ - { - "image_url": {"url": image_data_uri, "detail": "auto"}, - "index": 0, - "type": "image_url", - } - ], - }, - ] - - # Convert messages to Gemini format - contents = _gemini_convert_messages_with_history(messages=messages) - - # Verify structure - assert len(contents) == 2, f"Expected 2 content blocks, got {len(contents)}" - - # Verify user message - assert contents[0]["role"] == "user" - assert len(contents[0]["parts"]) == 1 - assert ( - contents[0]["parts"][0]["text"] - == "Generate an image of a banana wearing a costume that says LiteLLM" - ) - - # Verify assistant message - assert contents[1]["role"] == "model" - assert ( - len(contents[1]["parts"]) == 2 - ), f"Expected 2 parts (text + image), got {len(contents[1]['parts'])}" - - # Find text part and inline_data part - text_part = None - inline_data_part = None - for part in contents[1]["parts"]: - if "text" in part: - text_part = part - elif "inline_data" in part: - inline_data_part = part - - # Verify text part - assert text_part is not None, "Missing text part in assistant message" - assert text_part["text"] == "Here's your banana in a LiteLLM costume!" - - # Verify inline_data part (image) - assert inline_data_part is not None, "Missing inline_data part in assistant message" - inline_data: BlobType = inline_data_part["inline_data"] - assert "data" in inline_data - assert "mime_type" in inline_data - assert inline_data["mime_type"] == "image/png" - assert inline_data["data"] == test_image_base64 - - -def test_assistant_message_with_multiple_images(): - """Test that assistant messages with multiple images are properly converted.""" - # Create two test images - test_image1_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" - test_image2_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8DwHwAFBQIAX8jx0gAAAABJRU5ErkJggg==" - image1_data_uri = f"data:image/png;base64,{test_image1_base64}" - image2_data_uri = f"data:image/jpeg;base64,{test_image2_base64}" - - messages = [ - {"role": "user", "content": "Generate two images"}, - { - "role": "assistant", - "content": "Here are your images:", - "images": [ - { - "image_url": {"url": image1_data_uri, "detail": "auto"}, - "index": 0, - "type": "image_url", - }, - { - "image_url": {"url": image2_data_uri, "detail": "high"}, - "index": 1, - "type": "image_url", - }, - ], - }, - ] - - # Convert messages to Gemini format - contents = _gemini_convert_messages_with_history(messages=messages) - - # Verify assistant message has 3 parts (1 text + 2 images) - assert contents[1]["role"] == "model" - assert ( - len(contents[1]["parts"]) == 3 - ), f"Expected 3 parts (text + 2 images), got {len(contents[1]['parts'])}" - - # Count inline_data parts - inline_data_parts = [part for part in contents[1]["parts"] if "inline_data" in part] - assert ( - len(inline_data_parts) == 2 - ), f"Expected 2 inline_data parts, got {len(inline_data_parts)}" - - # Verify first image - assert inline_data_parts[0]["inline_data"]["mime_type"] == "image/png" - assert inline_data_parts[0]["inline_data"]["data"] == test_image1_base64 - - # Verify second image - assert inline_data_parts[1]["inline_data"]["mime_type"] == "image/jpeg" - assert inline_data_parts[1]["inline_data"]["data"] == test_image2_base64 - - -def test_assistant_message_with_images_using_message_object(): - """Test that Message objects with images field are properly converted.""" - # Create a small test image - test_image_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" - image_data_uri = f"data:image/png;base64,{test_image_base64}" - - # Create messages using Message object (as returned by LiteLLM) - user_message = {"role": "user", "content": "Generate an image"} - - assistant_message = Message( - content="Here's your image!", - role="assistant", - tool_calls=None, - function_call=None, - images=[ - { - "image_url": {"url": image_data_uri, "detail": "auto"}, - "index": 0, - "type": "image_url", - } - ], - ) - - messages = [user_message, assistant_message] - - # Convert messages to Gemini format - contents = _gemini_convert_messages_with_history(messages=messages) - - # Verify assistant message has both text and image - assert contents[1]["role"] == "model" - assert len(contents[1]["parts"]) == 2 - - # Verify image was converted - inline_data_parts = [part for part in contents[1]["parts"] if "inline_data" in part] - assert len(inline_data_parts) == 1 - assert inline_data_parts[0]["inline_data"]["mime_type"] == "image/png" - assert inline_data_parts[0]["inline_data"]["data"] == test_image_base64 - - -def test_assistant_message_with_images_in_conversation_history(): - """ - Test multi-turn conversation where assistant message with images is in history. - - This simulates the real use case where: - 1. User asks for image generation - 2. Assistant generates image (with images field) - 3. User asks follow-up question about the image - """ - test_image_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" - image_data_uri = f"data:image/png;base64,{test_image_base64}" - - messages = [ - {"role": "user", "content": "Generate an image of a cat"}, - { - "role": "assistant", - "content": "Here's a cat image:", - "images": [ - { - "image_url": {"url": image_data_uri, "detail": "auto"}, - "index": 0, - "type": "image_url", - } - ], - }, - {"role": "user", "content": "Can you make it more colorful?"}, - ] - - # Convert messages to Gemini format - contents = _gemini_convert_messages_with_history(messages=messages) - - # Verify structure: user -> model (with image) -> user - assert len(contents) == 3 - assert contents[0]["role"] == "user" - assert contents[1]["role"] == "model" - assert contents[2]["role"] == "user" - - # Verify assistant message has image in history - inline_data_parts = [part for part in contents[1]["parts"] if "inline_data" in part] - assert len(inline_data_parts) == 1 - assert inline_data_parts[0]["inline_data"]["mime_type"] == "image/png" - - -def test_function_response_has_user_role(): - """ - Test that function response ContentType blocks include role="user". - - Gemini API only accepts two roles: "user" and "model". Function responses - must be sent with role="user". Previously, LiteLLM omitted the role field - entirely, causing 400 errors from the Gemini API. - - Fixes: https://github.com/BerriAI/litellm/issues/22003 - Fixes: https://github.com/BerriAI/litellm/issues/20690 - """ - messages = [ - {"role": "user", "content": "What is the weather in Berlin?"}, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_abc123", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"city": "Berlin"}', - }, - } - ], - }, - { - "role": "tool", - "tool_call_id": "call_abc123", - "content": '{"temperature": "15°C", "condition": "Cloudy"}', - }, - ] - - contents = _gemini_convert_messages_with_history(messages=messages) - - # Expect: user -> model (functionCall) -> user (functionResponse) - assert len(contents) == 3 - - assert contents[0]["role"] == "user" - assert contents[1]["role"] == "model" - assert "function_call" in contents[1]["parts"][0] - - # The critical assertion: function response must have role="user" - assert contents[2]["role"] == "user" - assert "function_response" in contents[2]["parts"][0] - - -def test_multi_turn_function_calling_roles(): - """ - Test a full multi-turn function calling conversation produces correct roles. - - Simulates: user asks → model calls tool → tool responds → model answers → user asks again. - Every content block must have an explicit role of "user" or "model". - - Fixes: https://github.com/BerriAI/litellm/issues/22003 - """ - messages = [ - {"role": "user", "content": "What is the weather in Berlin?"}, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_001", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"city": "Berlin"}', - }, - } - ], - }, - { - "role": "tool", - "tool_call_id": "call_001", - "content": '{"temperature": "15°C"}', - }, - { - "role": "assistant", - "content": "The weather in Berlin is 15°C.", - }, - {"role": "user", "content": "And in Paris?"}, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_002", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"city": "Paris"}', - }, - } - ], - }, - { - "role": "tool", - "tool_call_id": "call_002", - "content": '{"temperature": "18°C"}', - }, - ] - - contents = _gemini_convert_messages_with_history(messages=messages) - - # Every content block must have a valid role - for i, content in enumerate(contents): - assert "role" in content, f"Content block {i} missing 'role' field" - assert content["role"] in ( - "user", - "model", - ), f"Content block {i} has invalid role: {content.get('role')}" - - # Verify the function response blocks specifically have role="user" - for i, content in enumerate(contents): - for part in content["parts"]: - if "function_response" in part: - assert ( - content["role"] == "user" - ), f"Content block {i} with function_response has role='{content['role']}', expected 'user'" - - -def test_gemini_thought_signature_preservation_real_response(): - """Test that thought signatures are preserved on the text part if originally there, without dropping or duplicating (real response case).""" - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - real_candidate = { - "content": { - "parts": [ - { - "text": "I will explain and then list files.", - "thoughtSignature": "mock_signature_from_text_part", - }, - { - "functionCall": { - "name": "list_files", - "args": {}, - } - }, - ] - } - } - - parts = real_candidate["content"]["parts"] - - content, reasoning_content = ( - VertexGeminiConfig().get_assistant_content_message(parts=parts) - ) - thought_signatures = ( - VertexGeminiConfig()._extract_thought_signatures_from_parts( - parts=parts - ) - ) - functions, tools, _ = VertexGeminiConfig._transform_parts( - parts=parts, - cumulative_tool_call_idx=0, - is_function_call=False, - ) - - msg: dict = {"role": "assistant"} - if content is not None: - msg["content"] = content - if tools: - msg["tool_calls"] = tools - if functions is not None: - msg["function_call"] = functions - if thought_signatures is not None: - msg["provider_specific_fields"] = { - "thought_signatures": thought_signatures - } - - converted_real = _gemini_convert_messages_with_history( - messages=[msg], - model="gemini-2.5-pro", - ) - - assert len(converted_real) == 1 - assert "parts" in converted_real[0] - parts_out = converted_real[0]["parts"] - assert len(parts_out) == 2 - assert "text" in parts_out[0] - assert ( - parts_out[0]["thoughtSignature"] == "mock_signature_from_text_part" - ) - assert "function_call" in parts_out[1] - assert "thoughtSignature" not in parts_out[1] - - -def test_gemini_thought_signature_deduplication_assumed_response(): - """Test that thought signatures are deduplicated and not attached to the text part if already present in the tool call (assumed response case).""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - pr_assumed_msg = { - "role": "assistant", - "content": "I will list the directory.", - "provider_specific_fields": { - "thought_signatures": ["mock_signature_63k"] - }, - "tool_calls": [ - { - "id": "call_1", - "type": "function", - "function": {"name": "list_files", "arguments": "{}"}, - "provider_specific_fields": { - "thought_signature": "mock_signature_63k" - }, - } - ], - } - - converted_pr = _gemini_convert_messages_with_history( - messages=[pr_assumed_msg], - model="gemini-2.5-pro", - ) - - assert len(converted_pr) == 1 - assert "parts" in converted_pr[0] - parts_out = converted_pr[0]["parts"] - assert len(parts_out) == 2 - assert "text" in parts_out[0] - assert "thoughtSignature" not in parts_out[0] - assert "function_call" in parts_out[1] - assert parts_out[1]["thoughtSignature"] == "mock_signature_63k" - - -def test_gemini_thought_signature_pure_text(): - """Test that thought signatures are preserved on the text part for responses with no tool calls.""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - msg = { - "role": "assistant", - "content": "Hello, I am a model.", - "provider_specific_fields": { - "thought_signatures": ["pure_text_signature"] - }, - } - - converted = _gemini_convert_messages_with_history( - messages=[msg], - model="gemini-2.5-pro", - ) - - assert len(converted) == 1 - assert "parts" in converted[0] - parts_out = converted[0]["parts"] - assert len(parts_out) == 1 - assert "text" in parts_out[0] - assert parts_out[0]["thoughtSignature"] == "pure_text_signature" - - -def test_gemini_thought_signature_pure_tool_call(): - """Test that thought signatures are preserved on the tool call for responses with no intermediate text.""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - msg = { - "role": "assistant", - "content": None, - "provider_specific_fields": { - "thought_signatures": ["pure_tool_signature"] - }, - "tool_calls": [ - { - "id": "call_1", - "type": "function", - "function": {"name": "list_files", "arguments": "{}"}, - "provider_specific_fields": { - "thought_signature": "pure_tool_signature" - }, - } - ], - } - - converted = _gemini_convert_messages_with_history( - messages=[msg], - model="gemini-2.5-pro", - ) - - assert len(converted) == 1 - assert "parts" in converted[0] - parts_out = converted[0]["parts"] - assert len(parts_out) == 1 - assert "function_call" in parts_out[0] - assert parts_out[0]["thoughtSignature"] == "pure_tool_signature" - - -def test_gemini_distinct_text_and_tool_signatures_are_both_preserved(): - """A text-part signature that differs from the tool-call signature must stay on the text part.""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - msg = { - "role": "assistant", - "content": "Some analysis.", - "provider_specific_fields": { - "thought_signatures": ["text_signature", "tool_signature"] - }, - "tool_calls": [ - { - "id": "call_1", - "type": "function", - "function": {"name": "list_files", "arguments": "{}"}, - "provider_specific_fields": {"thought_signature": "tool_signature"}, - } - ], - } - - parts = _gemini_convert_messages_with_history( - messages=[msg], model="gemini-2.5-pro" - )[0]["parts"] - - assert parts[0]["text"] == "Some analysis." - assert parts[0]["thoughtSignature"] == "text_signature" - assert "function_call" in parts[1] - assert parts[1]["thoughtSignature"] == "tool_signature" - - -def test_gemini_25_text_signature_survives_replay_to_gemini_3(): - """gemini-2.5 history (signed text, unsigned tool call) replayed to gemini-3 keeps the real - text signature; the dummy signature synthesized for the unsigned tool call must not suppress it.""" - from litellm.litellm_core_utils.prompt_templates.factory import ( - _get_dummy_thought_signature, - ) - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - msg = { - "role": "assistant", - "content": "I will list the directory.", - "provider_specific_fields": {"thought_signatures": ["real_25_signature"]}, - "tool_calls": [ - { - "id": "call_1", - "type": "function", - "function": {"name": "list_files", "arguments": "{}"}, - } - ], - } - - parts = _gemini_convert_messages_with_history(messages=[msg], model="gemini-3-pro")[ - 0 - ]["parts"] - - assert parts[0]["text"] == "I will list the directory." - assert parts[0]["thoughtSignature"] == "real_25_signature" - assert "function_call" in parts[1] - assert parts[1]["thoughtSignature"] == _get_dummy_thought_signature() - - -def test_gemini_function_call_signature_round_trip_no_duplicate(): - """End to end: a gemini-3-style response (unsigned text + signed functionCall) parsed and - re-serialized sends the signature exactly once, on the function-call part.""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - - response_parts = [ - {"text": "I will calculate the result for you."}, - { - "functionCall": {"name": "add_numbers", "args": {"a": 17, "b": 25}}, - "thoughtSignature": "signature_from_function_call", - }, - ] - - config = VertexGeminiConfig() - content, _ = config.get_assistant_content_message(parts=response_parts) - thought_signatures = config._extract_thought_signatures_from_parts( - parts=response_parts - ) - _, tools, _ = VertexGeminiConfig._transform_parts( - parts=response_parts, cumulative_tool_call_idx=0, is_function_call=False - ) - - msg = { - "role": "assistant", - "content": content, - "tool_calls": tools, - "provider_specific_fields": {"thought_signatures": thought_signatures}, - } - - parts = _gemini_convert_messages_with_history(messages=[msg], model="gemini-3-pro")[ - 0 - ]["parts"] - - signatures = [p["thoughtSignature"] for p in parts if "thoughtSignature" in p] - assert signatures == ["signature_from_function_call"] - assert "thoughtSignature" not in parts[0] - assert "function_call" in parts[1] - - -def test_gemini_server_side_tool_signature_not_duplicated_on_text(): - """A signature already re-injected on a server-side toolCall part is not attached to the text part again.""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - msg = { - "role": "assistant", - "content": "The weather in Buenos Aires is sunny.", - "provider_specific_fields": { - "thought_signatures": ["server_side_signature"], - "server_side_tool_invocations": [ - { - "tool_type": "GOOGLE_SEARCH_WEB", - "id": "abc123", - "args": {"queries": ["weather Buenos Aires"]}, - "response": {"weather": "Sunny"}, - "thought_signature": "server_side_signature", - } - ], - }, - } - - parts = _gemini_convert_messages_with_history( - messages=[msg], model="gemini-2.5-pro" - )[0]["parts"] - - text_part = next(p for p in parts if "text" in p) - assert "thoughtSignature" not in text_part - tool_call_part = next(p for p in parts if "toolCall" in p) - assert tool_call_part["thoughtSignature"] == "server_side_signature" diff --git a/tests/test_litellm/llms/vertex_ai/image_edit/__init__.py b/tests/test_litellm/llms/vertex_ai/image_edit/__init__.py deleted file mode 100644 index 50135ba1f92..00000000000 --- a/tests/test_litellm/llms/vertex_ai/image_edit/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# Vertex AI Image Edit Tests diff --git a/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py b/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py index 54607cc5284..aeba9f0fa3c 100644 --- a/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py @@ -1,13 +1,9 @@ import os -from unittest.mock import MagicMock, patch +from unittest.mock import patch -import httpx import pytest -from litellm.llms.vertex_ai.image_generation import ( - get_vertex_ai_image_generation_config, -) from litellm.llms.vertex_ai.image_generation.vertex_gemini_transformation import ( VertexAIGeminiImageGenerationConfig, ) @@ -16,588 +12,6 @@ from litellm.llms.vertex_ai.image_generation.vertex_imagen_transformation import ) -class TestVertexAIGeminiImageGenerationConfig: - def setup_method(self): - """Set up test fixtures""" - self.config = VertexAIGeminiImageGenerationConfig() - - def test_get_supported_openai_params(self): - """Test get_supported_openai_params returns correct params""" - supported = self.config.get_supported_openai_params("gemini-2.5-flash-image") - assert "n" in supported - assert "size" in supported - - def test_map_openai_params_n(self): - """Test mapping n parameter to candidate_count""" - non_default_params = {"n": 3} - optional_params = {} - result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False) - assert result.get("candidate_count") == 3 - - def test_map_openai_params_size(self): - """Test mapping size parameter to aspectRatio""" - non_default_params = {"size": "1024x1024"} - optional_params = {} - result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False) - assert result.get("aspectRatio") == "1:1" - - def test_map_openai_params_size_16_9(self): - """Test mapping 16:9 size""" - non_default_params = {"size": "1792x1024"} - optional_params = {} - result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False) - assert result.get("aspectRatio") == "16:9" - - def test_map_size_to_aspect_ratio(self): - """Test size to aspect ratio mapping""" - assert self.config._map_size_to_aspect_ratio("1024x1024") == "1:1" - assert self.config._map_size_to_aspect_ratio("1792x1024") == "16:9" - assert self.config._map_size_to_aspect_ratio("1024x1792") == "9:16" - assert self.config._map_size_to_aspect_ratio("1280x896") == "4:3" - assert self.config._map_size_to_aspect_ratio("896x1280") == "3:4" - assert self.config._map_size_to_aspect_ratio("unknown") == "1:1" # default - - def test_get_supported_openai_params_includes_native_gemini_params(self): - """Test that native Gemini imageConfig params are supported""" - supported = self.config.get_supported_openai_params("gemini-3-pro-image-preview") - assert "aspectRatio" in supported - assert "aspect_ratio" in supported - assert "imageSize" in supported - assert "image_size" in supported - assert "imageConfig" in supported - - def test_map_openai_params_aspect_ratio_camel_case(self): - """Test mapping native aspectRatio parameter""" - result = self.config.map_openai_params({"aspectRatio": "9:16"}, {}, "gemini-3-pro-image-preview", False) - assert result["aspectRatio"] == "9:16" - - def test_map_openai_params_aspect_ratio_snake_case(self): - """Test mapping native aspect_ratio parameter""" - result = self.config.map_openai_params({"aspect_ratio": "16:9"}, {}, "gemini-3-pro-image-preview", False) - assert result["aspectRatio"] == "16:9" - - def test_map_openai_params_image_size_camel_case(self): - """Test mapping native imageSize parameter""" - result = self.config.map_openai_params({"imageSize": "4K"}, {}, "gemini-3-pro-image-preview", False) - assert result["imageSize"] == "4K" - - def test_map_openai_params_image_size_snake_case(self): - """Test mapping native image_size parameter""" - result = self.config.map_openai_params({"image_size": "2K"}, {}, "gemini-3-pro-image-preview", False) - assert result["imageSize"] == "2K" - - def test_map_openai_params_image_config_dict_stored_whole(self): - """imageConfig dict is stored as-is so all fields survive""" - result = self.config.map_openai_params( - {"imageConfig": {"aspectRatio": "16:9", "imageSize": "2K"}}, - {}, - "gemini-3.1-flash-image", - False, - ) - assert result["imageConfig"] == {"aspectRatio": "16:9", "imageSize": "2K"} - - def test_map_openai_params_image_config_all_fields(self): - """All ImageConfig fields (personGeneration, imageOutputOptions) pass through""" - payload = { - "imageConfig": { - "aspectRatio": "9:16", - "imageSize": "4K", - "personGeneration": "DONT_ALLOW", - "imageOutputOptions": { - "mimeType": "image/jpeg", - "compressionQuality": 80, - }, - } - } - result = self.config.map_openai_params(payload, {}, "gemini-3.1-flash-image", False) - assert result["imageConfig"] == payload["imageConfig"] - - def test_map_openai_params_image_config_non_dict_warns_and_drops(self): - """Non-dict imageConfig is dropped with a warning, not silently discarded""" - with patch("litellm.llms.vertex_ai.image_generation.vertex_gemini_transformation.verbose_logger") as mock_log: - result = self.config.map_openai_params( - {"imageConfig": "bad-string-value"}, {}, "gemini-3.1-flash-image", False - ) - assert "imageConfig" not in result - mock_log.warning.assert_called_once() - - def test_transform_image_generation_request_from_image_config(self): - """Full imageConfig dict is forwarded verbatim into generationConfig""" - full_config = { - "aspectRatio": "16:9", - "imageSize": "2K", - "personGeneration": "DONT_ALLOW", - "imageOutputOptions": {"mimeType": "image/jpeg", "compressionQuality": 85}, - } - mapped = self.config.map_openai_params( - {"imageConfig": full_config}, - {}, - "gemini-3.1-flash-image", - False, - ) - request = self.config.transform_image_generation_request( - model="gemini-3.1-flash-image", - prompt="A nano banana on a desk", - optional_params=mapped, - litellm_params={}, - headers={}, - ) - assert request["generationConfig"]["imageConfig"] == full_config - - def test_transform_image_generation_flat_params_override_image_config(self): - """Explicit flat params win over the same key inside imageConfig""" - request = self.config.transform_image_generation_request( - model="gemini-3.1-flash-image", - prompt="A nano banana", - optional_params={ - "imageConfig": {"aspectRatio": "1:1", "personGeneration": "DONT_ALLOW"}, - "aspectRatio": "16:9", # should win - }, - litellm_params={}, - headers={}, - ) - assert request["generationConfig"]["imageConfig"]["aspectRatio"] == "16:9" - assert request["generationConfig"]["imageConfig"]["personGeneration"] == "DONT_ALLOW" - - def test_transform_image_generation_request_basic(self): - """Test basic request transformation""" - request = self.config.transform_image_generation_request( - model="gemini-2.5-flash-image", - prompt="A nano banana", - optional_params={}, - litellm_params={}, - headers={}, - ) - assert "contents" in request - assert "generationConfig" in request - assert request["generationConfig"]["responseModalities"] == ["IMAGE"] - assert request["contents"][0]["parts"][0]["text"] == "A nano banana" - - def test_transform_image_generation_request_with_aspect_ratio(self): - """Test request transformation with aspectRatio""" - request = self.config.transform_image_generation_request( - model="gemini-2.5-flash-image", - prompt="A nano banana", - optional_params={"aspectRatio": "16:9"}, - litellm_params={}, - headers={}, - ) - assert request["generationConfig"]["imageConfig"]["aspectRatio"] == "16:9" - - def test_transform_image_generation_request_with_image_size(self): - """Test request transformation with imageSize (Gemini 3 Pro)""" - request = self.config.transform_image_generation_request( - model="gemini-3-pro-image-preview", - prompt="A nano banana", - optional_params={"imageSize": "4K"}, - litellm_params={}, - headers={}, - ) - assert request["generationConfig"]["imageConfig"]["imageSize"] == "4K" - - def test_map_openai_params_web_search_options(self): - """Test web_search_options maps to googleSearch tool""" - result = self.config.map_openai_params({"web_search_options": {}}, {}, "gemini-3.1-flash-image-preview", False) - assert result["tools"] == [{"googleSearch": {}}] - - def test_transform_image_generation_request_with_web_search_tools(self): - """Test request transformation includes googleSearch tools""" - request = self.config.transform_image_generation_request( - model="gemini-3.1-flash-image-preview", - prompt="Generate an image of the latest iPhone", - optional_params={"tools": [{"googleSearch": {}}]}, - litellm_params={}, - headers={}, - ) - assert request["tools"] == [{"googleSearch": {}}] - - def test_transform_image_generation_request_forwards_tool_config(self): - """Test request transformation forwards toolConfig side-effects from tool mapping""" - mapped = self.config.map_openai_params( - {"tools": [{"googleMaps": {"latitude": 37.7, "longitude": -122.4}}]}, - {}, - "gemini-3.1-flash-image-preview", - False, - ) - request = self.config.transform_image_generation_request( - model="gemini-3.1-flash-image-preview", - prompt="Generate an image of a coffee shop nearby", - optional_params=mapped, - litellm_params={}, - headers={}, - ) - assert request["tools"] == [{"googleMaps": {}}] - assert request["toolConfig"] == {"retrievalConfig": {"latLng": {"latitude": 37.7, "longitude": -122.4}}} - - def test_transform_image_generation_request_with_candidate_count(self): - """Test request transformation with candidate_count""" - request = self.config.transform_image_generation_request( - model="gemini-2.5-flash-image", - prompt="A nano banana", - optional_params={"candidate_count": 2}, - litellm_params={}, - headers={}, - ) - assert request["generationConfig"]["candidateCount"] == 2 - - def test_transform_image_generation_request_with_n(self): - """Test request transformation with n parameter""" - request = self.config.transform_image_generation_request( - model="gemini-2.5-flash-image", - prompt="A nano banana", - optional_params={"n": 2}, - litellm_params={}, - headers={}, - ) - assert request["generationConfig"]["candidateCount"] == 2 - - def test_transform_image_generation_response(self): - """Test response transformation""" - mock_response = MagicMock(spec=httpx.Response) - mock_response.status_code = 200 - mock_response.json.return_value = { - "candidates": [ - { - "content": { - "parts": [ - { - "inlineData": { - "mimeType": "image/png", - "data": "base64_encoded_image_data", - } - } - ] - } - } - ], - "usageMetadata": { - "promptTokenCount": 93, - "promptTokensDetails": [ - { - "modality": "TEXT", - "tokenCount": 54, - }, - { - "modality": "IMAGE", - "tokenCount": 39, - }, - ], - "candidatesTokenCount": 17, - "totalTokenCount": 110, - }, - } - mock_response.headers = {} - - from litellm.types.utils import ImageResponse - - model_response = ImageResponse() - result = self.config.transform_image_generation_response( - model="gemini-2.5-flash-image", - raw_response=mock_response, - model_response=model_response, - logging_obj=MagicMock(), - request_data={}, - optional_params={}, - litellm_params={}, - encoding=None, - ) - - assert len(result.data) == 1 - assert result.data[0].b64_json == "base64_encoded_image_data" - assert result.data[0].url is None - assert result.usage.input_tokens == 93 - assert result.usage.input_tokens_details.text_tokens == 54 - assert result.usage.input_tokens_details.image_tokens == 39 - assert result.usage.output_tokens == 17 - assert result.usage.total_tokens == 110 - - def test_transform_image_generation_response_multiple_images(self): - """Test response transformation with multiple images""" - mock_response = MagicMock(spec=httpx.Response) - mock_response.status_code = 200 - mock_response.json.return_value = { - "candidates": [ - { - "content": { - "parts": [ - { - "inlineData": { - "mimeType": "image/png", - "data": "image1", - } - }, - { - "inlineData": { - "mimeType": "image/png", - "data": "image2", - } - }, - ] - } - } - ] - } - mock_response.headers = {} - - from litellm.types.utils import ImageResponse - - model_response = ImageResponse() - result = self.config.transform_image_generation_response( - model="gemini-2.5-flash-image", - raw_response=mock_response, - model_response=model_response, - logging_obj=MagicMock(), - request_data={}, - optional_params={}, - litellm_params={}, - encoding=None, - ) - - assert len(result.data) == 2 - assert result.data[0].b64_json == "image1" - assert result.data[1].b64_json == "image2" - - def test_transform_image_generation_response_signature(self): - """Test response transformation includes thoughtSignature for Gemini 3 Pro""" - mock_response = MagicMock(spec=httpx.Response) - mock_response.status_code = 200 - mock_response.json.return_value = { - "candidates": [ - { - "content": { - "parts": [ - { - "inlineData": { - "mimeType": "image/png", - "data": "base64_encoded_image_data", - }, - "thoughtSignature": "test_signature_abc123", - } - ] - } - } - ] - } - mock_response.headers = {} - - from litellm.types.utils import ImageResponse - - model_response = ImageResponse() - result = self.config.transform_image_generation_response( - model="gemini-3-pro-image-preview", - raw_response=mock_response, - model_response=model_response, - logging_obj=MagicMock(), - request_data={}, - optional_params={}, - litellm_params={}, - encoding=None, - ) - - assert len(result.data) == 1 - assert result.data[0].b64_json == "base64_encoded_image_data" - assert result.data[0].provider_specific_fields["thought_signature"] == "test_signature_abc123" - - def test_transform_image_generation_response_tracks_web_search_requests(self): - """Grounding queries are carried onto usage so search spend can be billed""" - mock_response = MagicMock(spec=httpx.Response) - mock_response.status_code = 200 - mock_response.json.return_value = { - "candidates": [ - { - "content": { - "parts": [ - { - "inlineData": { - "mimeType": "image/png", - "data": "base64_encoded_image_data", - } - } - ] - }, - "groundingMetadata": {"webSearchQueries": ["eiffel tower", "paris skyline"]}, - } - ], - "usageMetadata": { - "promptTokenCount": 93, - "candidatesTokenCount": 17, - "totalTokenCount": 110, - }, - } - mock_response.headers = {} - - from litellm.types.utils import ImageResponse - - result = self.config.transform_image_generation_response( - model="gemini-2.5-flash-image", - raw_response=mock_response, - model_response=ImageResponse(), - logging_obj=MagicMock(), - request_data={}, - optional_params={}, - litellm_params={}, - encoding=None, - ) - - assert result.usage.web_search_requests == 2 - - -class TestVertexAIImagenImageGenerationConfig: - def setup_method(self): - """Set up test fixtures""" - self.config = VertexAIImagenImageGenerationConfig() - - def test_get_supported_openai_params(self): - """Test get_supported_openai_params returns correct params""" - supported = self.config.get_supported_openai_params("imagegeneration@006") - assert "n" in supported - assert "size" in supported - - def test_map_openai_params_n(self): - """Test mapping n parameter to sampleCount""" - non_default_params = {"n": 3} - optional_params = {} - result = self.config.map_openai_params(non_default_params, optional_params, "imagegeneration@006", False) - assert result.get("sampleCount") == 3 - - def test_map_openai_params_size(self): - """Test mapping size parameter to aspectRatio""" - non_default_params = {"size": "1024x1024"} - optional_params = {} - result = self.config.map_openai_params(non_default_params, optional_params, "imagegeneration@006", False) - assert result.get("aspectRatio") == "1:1" - - def test_map_size_to_aspect_ratio(self): - """Test size to aspect ratio mapping""" - assert self.config._map_size_to_aspect_ratio("1024x1024") == "1:1" - assert self.config._map_size_to_aspect_ratio("1792x1024") == "16:9" - assert self.config._map_size_to_aspect_ratio("unknown") == "1:1" # default - - def test_transform_image_generation_request_basic(self): - """Test basic request transformation""" - request = self.config.transform_image_generation_request( - model="imagegeneration@006", - prompt="A cat", - optional_params={}, - litellm_params={}, - headers={}, - ) - assert "instances" in request - assert "parameters" in request - assert request["instances"][0]["prompt"] == "A cat" - assert request["parameters"]["sampleCount"] == 1 - - def test_transform_image_generation_request_with_params(self): - """Test request transformation with parameters""" - request = self.config.transform_image_generation_request( - model="imagegeneration@006", - prompt="A cat", - optional_params={"sampleCount": 2, "aspectRatio": "16:9"}, - litellm_params={}, - headers={}, - ) - assert request["parameters"]["sampleCount"] == 2 - assert request["parameters"]["aspectRatio"] == "16:9" - - def test_transform_image_generation_request_labels_from_metadata(self): - """Billing labels from litellm_params.metadata.requester_metadata on predict body.""" - request = self.config.transform_image_generation_request( - model="imagegeneration@006", - prompt="A cat", - optional_params={}, - litellm_params={"metadata": {"requester_metadata": {"team": "platform", "env": "prod"}}}, - headers={}, - ) - assert request["labels"] == {"team": "platform", "env": "prod"} - assert "labels" not in request["parameters"] - - def test_transform_image_generation_response(self): - """Test response transformation""" - mock_response = MagicMock(spec=httpx.Response) - mock_response.status_code = 200 - mock_response.json.return_value = {"predictions": [{"bytesBase64Encoded": "base64_encoded_image_data"}]} - mock_response.headers = {} - - from litellm.types.utils import ImageResponse - - model_response = ImageResponse() - result = self.config.transform_image_generation_response( - model="imagegeneration@006", - raw_response=mock_response, - model_response=model_response, - logging_obj=MagicMock(), - request_data={}, - optional_params={}, - litellm_params={}, - encoding=None, - ) - - assert len(result.data) == 1 - assert result.data[0].b64_json == "base64_encoded_image_data" - assert result.data[0].url is None - - def test_transform_image_generation_response_multiple_images(self): - """Test response transformation with multiple images""" - mock_response = MagicMock(spec=httpx.Response) - mock_response.status_code = 200 - mock_response.json.return_value = { - "predictions": [ - {"bytesBase64Encoded": "image1"}, - {"bytesBase64Encoded": "image2"}, - ] - } - mock_response.headers = {} - - from litellm.types.utils import ImageResponse - - model_response = ImageResponse() - result = self.config.transform_image_generation_response( - model="imagegeneration@006", - raw_response=mock_response, - model_response=model_response, - logging_obj=MagicMock(), - request_data={}, - optional_params={}, - litellm_params={}, - encoding=None, - ) - - assert len(result.data) == 2 - assert result.data[0].b64_json == "image1" - assert result.data[1].b64_json == "image2" - - -class TestGetVertexAIImageGenerationConfig: - """Test the router function that selects the correct config""" - - def test_get_gemini_model_config(self): - """Test that Gemini models return Gemini config""" - config = get_vertex_ai_image_generation_config("gemini-2.5-flash-image") - assert isinstance(config, VertexAIGeminiImageGenerationConfig) - - config = get_vertex_ai_image_generation_config("gemini-3-pro-image-preview") - assert isinstance(config, VertexAIGeminiImageGenerationConfig) - - config = get_vertex_ai_image_generation_config("vertex_ai/gemini-2.5-flash-image") - assert isinstance(config, VertexAIGeminiImageGenerationConfig) - - def test_get_imagen_model_config(self): - """Test that Imagen models return Imagen config""" - config = get_vertex_ai_image_generation_config("imagegeneration@006") - assert isinstance(config, VertexAIImagenImageGenerationConfig) - - config = get_vertex_ai_image_generation_config("imagen-4.0-generate-001") - assert isinstance(config, VertexAIImagenImageGenerationConfig) - - config = get_vertex_ai_image_generation_config("vertex_ai/imagegeneration@006") - assert isinstance(config, VertexAIImagenImageGenerationConfig) - - def test_get_non_gemini_model_config(self): - """Test that non-Gemini models default to Imagen config""" - config = get_vertex_ai_image_generation_config("some-other-model") - assert isinstance(config, VertexAIImagenImageGenerationConfig) - - class TestVertexAIImageGenerationIntegration: """Integration tests for Vertex AI image generation""" @@ -642,39 +56,3 @@ class TestVertexAIImageGenerationIntegration: litellm_params={}, ) assert "Authorization" in headers - - def test_gemini_get_complete_url(self): - """Test Gemini config URL generation""" - config = VertexAIGeminiImageGenerationConfig() - url = config.get_complete_url( - api_base=None, - api_key=None, - model="gemini-2.5-flash-image", - optional_params={}, - litellm_params={ - "vertex_project": "test-project", - "vertex_location": "us-central1", - }, - ) - assert "test-project" in url - assert "us-central1" in url - assert "gemini-2.5-flash-image" in url - assert "generateContent" in url - - def test_imagen_get_complete_url(self): - """Test Imagen config URL generation""" - config = VertexAIImagenImageGenerationConfig() - url = config.get_complete_url( - api_base=None, - api_key=None, - model="imagegeneration@006", - optional_params={}, - litellm_params={ - "vertex_project": "test-project", - "vertex_location": "us-central1", - }, - ) - assert "test-project" in url - assert "us-central1" in url - assert "imagegeneration@006" in url - assert "predict" in url diff --git a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/__init__.py b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/__init__.py deleted file mode 100644 index 8b41c5ab3f8..00000000000 --- a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Tests for Vertex AI Gemma-AI models""" diff --git a/tests/test_litellm/llms/vertex_ai/videos/__init__.py b/tests/test_litellm/llms/vertex_ai/videos/__init__.py deleted file mode 100644 index f29c2a16fd5..00000000000 --- a/tests/test_litellm/llms/vertex_ai/videos/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -""" -Tests for Vertex AI video generation. -""" diff --git a/tests/test_litellm/messages/test_dispatch.py b/tests/test_litellm/messages/test_dispatch.py deleted file mode 100644 index 4da060f809a..00000000000 --- a/tests/test_litellm/messages/test_dispatch.py +++ /dev/null @@ -1,155 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping -from typing import Final - -import pytest -from pydantic import TypeAdapter - -import litellm -from litellm.messages import dispatch -from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.catalog import Route, RouteRule, Rules -from litellm.rust_bridge.configuration import Rollout -from litellm.rust_bridge.messages.entrypoints import ( - LiteLLMMessagesRequest, - NativeAmessages, - NativeMessages, -) -from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse - -MESSAGES: Final = [{"role": "user", "content": "hi"}] - - -@pytest.mark.asyncio -async def test_public_anthropic_messages_keeps_the_python_result() -> None: - response: Final = await litellm.anthropic_messages( - model="anthropic/claude-sonnet-4-5", messages=MESSAGES, max_tokens=10, mock_response="ok" - ) - - assert isinstance(response, dict) - content: Final = TypeAdapter(list[dict[str, object]]).validate_python(response.get("content", [])) - assert content[0]["text"] == "ok" - - -def test_sync_messages_request_projects_public_arguments() -> None: - rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),) - expected: Final = AnthropicMessagesResponse(model="claude-test") - - def native( - request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> AnthropicMessagesResponse: - assert request.model == "claude-test" - assert request.messages == MESSAGES - assert request.max_tokens == 10 - assert request.custom_llm_provider == "anthropic" - return expected - - binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None) - binding.override(native) - response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision - (), - { - "model": "claude-test", - "messages": MESSAGES, - "max_tokens": 10, - "custom_llm_provider": "anthropic", - }, - python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"), - binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), - rules=rules, - ) - - assert response is expected - - -def test_messages_binding_error_delegates_unchanged_to_python() -> None: - rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),) - expected: Final = AnthropicMessagesResponse(model="claude-test") - - def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse: - return expected - - def native( - request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> AnthropicMessagesResponse: - pytest.fail("a call without max_tokens cannot project a request and must stay on Python") - - binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None) - binding.override(native) - response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision - (), - {"model": "claude-test", "messages": MESSAGES, "custom_llm_provider": "anthropic"}, - python=python, - binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), - rules=rules, - ) - - assert response is expected - - -@pytest.mark.asyncio -async def test_async_messages_falls_back_after_native_declines() -> None: - from litellm.rust_bridge.bindings import native_exception_types - - native_types: Final = native_exception_types() - if native_types is None: - pytest.skip("native bridge is unavailable") - declined, _ = native_types - expected: Final = AnthropicMessagesResponse(model="claude-test") - rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_OPT_OUT),) - - async def native( - request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> AnthropicMessagesResponse: - raise declined("unsupported") - - async def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse: - return expected - - binding: Final[NativeBinding[NativeAmessages]] = NativeBinding("amessages", validate=lambda _: None) - binding.override(native) - response: Final = await dispatch._ADISPATCH.arun( # pyright: ignore[reportPrivateUsage] # test an explicit route decision - (), - {"model": "claude-test", "messages": MESSAGES, "max_tokens": 10}, - python=python, - binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), - rules=rules, - ) - - assert response is expected - - -def test_internal_is_async_marker_bypasses_native() -> None: - rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),) - expected: Final = AnthropicMessagesResponse(model="claude-test") - - def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse: - return expected - - def native( - request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> AnthropicMessagesResponse: - pytest.fail("anthropic_messages' inner handler call must stay on Python") - - binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None) - binding.override(native) - response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision - (), - { - "model": "claude-test", - "messages": MESSAGES, - "max_tokens": 10, - "custom_llm_provider": "anthropic", - "is_async": True, - }, - python=python, - binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), - rules=rules, - ) - - assert response is expected diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 1d3d7a452b6..45ad336368c 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -30,7 +30,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( BedrockTextContent, ) from litellm.types.utils import CallTypes, ModelResponse -from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe +from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py index bbc8fd539a3..5169d4c9ec6 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py @@ -23,7 +23,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( BedrockGuardrailResponse, ) from litellm.types.utils import Choices, Message, ModelResponse -from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe +from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe CONTENT_FILTER_CHECKS = {"contentFilter": {"categories": [{"category": "VIOLENCE"}]}} diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 1230c548281..21c0f565486 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -8029,6 +8029,37 @@ class TestPKCEStateCookieBinding: assert result is not None +@pytest.mark.asyncio +@pytest.mark.parametrize("enable_sso_debug_value", [None, "false", "0"]) +async def test_sso_debug_routes_return_404_unless_explicitly_enabled(enable_sso_debug_value): + """ + /sso/debug/login and /sso/debug/callback must 404 unless ENABLE_SSO_DEBUG is + explicitly set to a truthy value. + """ + from litellm.proxy.management_endpoints.ui_sso import debug_sso_callback, debug_sso_login + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://proxy.example.com/" + mock_request.cookies = {} + mock_request.query_params = {} + + env = {"GENERIC_CLIENT_ID": "test_client_id"} + if enable_sso_debug_value is not None: + env["ENABLE_SSO_DEBUG"] = enable_sso_debug_value + + with patch.dict(os.environ, env, clear=False): + if enable_sso_debug_value is None: + os.environ.pop("ENABLE_SSO_DEBUG", None) + + with pytest.raises(HTTPException) as login_exc: + await debug_sso_login(mock_request) + with pytest.raises(HTTPException) as callback_exc: + await debug_sso_callback(mock_request) + + assert login_exc.value.status_code == 404 + assert callback_exc.value.status_code == 404 + + @pytest.mark.asyncio async def test_debug_sso_callback_renders_full_jwt_claims(): """ @@ -8080,7 +8111,7 @@ async def test_debug_sso_callback_renders_full_jwt_claims(): with ( patch.dict( os.environ, - {"GENERIC_CLIENT_ID": "test_client_id"}, + {"GENERIC_CLIENT_ID": "test_client_id", "ENABLE_SSO_DEBUG": "true"}, clear=False, ), patch( @@ -8165,7 +8196,7 @@ async def test_debug_sso_callback_handles_missing_raw_response(): with ( patch.dict( os.environ, - {"MICROSOFT_CLIENT_ID": "test_microsoft_id"}, + {"MICROSOFT_CLIENT_ID": "test_microsoft_id", "ENABLE_SSO_DEBUG": "true"}, clear=False, ), patch.object( @@ -8213,7 +8244,7 @@ async def _render_debug_page(provider_env, id_jag_registered, force_inert=False) return parsed stack = [ - patch.dict(os.environ, provider_env, clear=False), + patch.dict(os.environ, {**provider_env, "ENABLE_SSO_DEBUG": "true"}, clear=False), patch( # test-quality-ok: endpoint test stubs the upstream generic IdP boundary "litellm.proxy.management_endpoints.ui_sso.get_generic_sso_response", side_effect=fake_generic ), diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 353ffadfa46..227921d6150 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -24,7 +24,7 @@ from starlette.datastructures import FormData import litellm from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe +from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( BaseOpenAIPassThroughHandler, diff --git a/tests/test_litellm/router_strategy/test_router_tag_routing.py b/tests/test_litellm/router_strategy/test_router_tag_routing.py index 16c641b8d29..e4b8860a7a6 100644 --- a/tests/test_litellm/router_strategy/test_router_tag_routing.py +++ b/tests/test_litellm/router_strategy/test_router_tag_routing.py @@ -2823,7 +2823,7 @@ def test_update_router_config_schema_includes_tag_routing_prefix(): # UpdateRouterConfig before calling update_settings; a field missing here # causes model_dump(exclude_none=True) to silently drop it before # update_settings is ever called -- the same bug shape LIT-3152 fixed for - # retry_policy (see tests/test_litellm/test_router_retry_policy_update.py). + # retry_policy (see tests/unit/test_router_retry_policy_update.py). from litellm.types.router import UpdateRouterConfig config = UpdateRouterConfig(tag_routing_prefix="route:") diff --git a/tests/test_litellm/test_compression.py b/tests/test_litellm/test_compression.py index 4fbcd4ed30d..997778d0a1b 100644 --- a/tests/test_litellm/test_compression.py +++ b/tests/test_litellm/test_compression.py @@ -3,20 +3,13 @@ Unit tests for litellm.compress(). """ import os -import importlib import pytest import litellm -from litellm.compression.scoring.bm25 import bm25_score_messages -from litellm.compression.scoring.embedding_scorer import embedding_score_messages -from litellm.compression.content_detection import detect_content_type -from litellm.compression.message_stubbing import extract_key, stub_message -from litellm.compression.retrieval_tool import build_retrieval_tool from litellm.types.utils import CallTypes CALL_TYPE = CallTypes.completion -ANTHROPIC_CALL_TYPE = CallTypes.anthropic_messages # --------------------------------------------------------------------------- @@ -24,420 +17,26 @@ ANTHROPIC_CALL_TYPE = CallTypes.anthropic_messages # --------------------------------------------------------------------------- -def test_bm25_relevance_ranking(): - query = "Fix the authentication bug in the login handler" - messages = [ - { - "role": "user", - "content": "def login_handler(): authentication check bug fix", - }, - {"role": "user", "content": "def render_template(name): css styling layout"}, - {"role": "user", "content": "def verify(): authentication token bug handler"}, - ] - scores = bm25_score_messages(query, messages) - # Messages sharing query terms should score higher than unrelated ones - assert scores[0] > scores[1] - assert scores[2] > scores[1] - - -def test_bm25_empty_query(): - scores = bm25_score_messages("", [{"role": "user", "content": "hello"}]) - assert scores == [0.0] - - -def test_bm25_empty_messages(): - scores = bm25_score_messages("query", []) - assert scores == [] - - -def test_bm25_empty_content(): - scores = bm25_score_messages("query", [{"role": "user", "content": ""}]) - assert scores == [0.0] - - # --------------------------------------------------------------------------- # Content detection # --------------------------------------------------------------------------- -def test_detect_code(): - code = """ -import os -from pathlib import Path - -def main(): - class Foo: - pass - return Foo() -""" - assert detect_content_type(code) == "code" - - -def test_detect_json(): - assert detect_content_type('{"key": "value", "num": 42}') == "json" - assert detect_content_type("[1, 2, 3]") == "json" - - -def test_detect_text(): - assert detect_content_type("This is a plain text paragraph about dogs.") == "text" - - -def test_detect_empty(): - assert detect_content_type("") == "text" - - # --------------------------------------------------------------------------- # Message stubbing # --------------------------------------------------------------------------- -def test_extract_key_with_filename(): - msg = {"role": "user", "content": "# auth.py\ndef authenticate():\n pass"} - used: set = set() - key = extract_key(msg, fallback_index=0, used_keys=used) - assert key == "auth.py" - - -def test_extract_key_fallback(): - msg = {"role": "user", "content": "Some random content without a filename"} - used: set = set() - key = extract_key(msg, fallback_index=5, used_keys=used) - assert key == "message_5" - - -def test_extract_key_duplicates(): - used: set = set() - msg = {"role": "user", "content": "# auth.py\ncode here"} - k1 = extract_key(msg, fallback_index=0, used_keys=used) - k2 = extract_key(msg, fallback_index=1, used_keys=used) - assert k1 == "auth.py" - assert k2 == "auth.py_2" - - -def test_stub_message(): - msg = {"role": "user", "content": "line1\nline2\nline3"} - stubbed = stub_message(msg, "test_key") - assert stubbed["role"] == "user" - assert "test_key" in stubbed["content"] - assert "litellm_content_retrieve" in stubbed["content"] - assert "3 lines" in stubbed["content"] - - # --------------------------------------------------------------------------- # Retrieval tool # --------------------------------------------------------------------------- -def test_retrieval_tool_schema(): - tool = build_retrieval_tool(["auth.py", "utils.py"]) - assert tool["type"] == "function" - assert tool["function"]["name"] == "litellm_content_retrieve" - assert "key" in tool["function"]["parameters"]["properties"] - assert tool["function"]["parameters"]["properties"]["key"]["enum"] == [ - "auth.py", - "utils.py", - ] - assert tool["function"]["parameters"]["required"] == ["key"] - - -def test_retrieval_tool_description_lists_keys(): - tool = build_retrieval_tool(["foo.py", "bar.js"]) - desc = tool["function"]["description"] - assert "foo.py" in desc - assert "bar.js" in desc - - # --------------------------------------------------------------------------- # compress() — end-to-end # --------------------------------------------------------------------------- -def test_compress_below_trigger_passthrough(): - messages = [{"role": "user", "content": "hello"}] - result = litellm.compress(messages, model="gpt-4o", call_type=CALL_TYPE) - assert result["messages"] == messages - assert result["cache"] == {} - assert result["tools"] == [] - assert result["compression_ratio"] == 0.0 - assert result["original_tokens"] == result["compressed_tokens"] - - -def test_compress_above_trigger(): - big_messages = [ - {"role": "system", "content": "You are a coding assistant."}, - { - "role": "user", - "content": "# auth.py\n" + "def authenticate():\n pass\n" * 2000, - }, - { - "role": "user", - "content": "# utils.py\n" + "def helper():\n pass\n" * 2000, - }, - { - "role": "user", - "content": "# readme.md\n" + "This is documentation. " * 2000, - }, - {"role": "user", "content": "Fix the bug in auth.py"}, - ] - - result = litellm.compress( - big_messages, - model="gpt-4o", - call_type=CALL_TYPE, - compression_trigger=1000, - compression_target=500, - ) - - assert result["compressed_tokens"] < result["original_tokens"] - assert result["compression_ratio"] > 0 - assert len(result["cache"]) > 0 - assert len(result["tools"]) == 1 - assert result["tools"][0]["function"]["name"] == "litellm_content_retrieve" - - -def test_compress_anthropic_list_content_is_boundary_stable(): - messages = [ - {"role": "system", "content": [{"type": "text", "text": "System prompt"}]}, - { - "role": "user", - "content": [ - {"type": "text", "text": "# a.py\n" + "alpha " * 2000}, - { - "type": "image_url", - "image_url": {"url": "https://example.com/a.png"}, - }, - ], - }, - { - "role": "user", - "content": [ - {"type": "text", "text": "# b.py\n" + "beta " * 2000}, - { - "type": "image_url", - "image_url": {"url": "https://example.com/b.png"}, - }, - ], - }, - { - "role": "user", - "content": [{"type": "text", "text": "Fix alpha bug in a.py"}], - }, - ] - - result = litellm.compress( - messages=messages, - model="claude-sonnet-4-20250514", - call_type=ANTHROPIC_CALL_TYPE, - compression_trigger=1000, - compression_target=500, - ) - - assert result["compressed_tokens"] < result["original_tokens"] - assert len(result["messages"]) == len(messages) - assert [m["role"] for m in result["messages"]] == [m["role"] for m in messages] - assert len(result["cache"]) > 0 - assert len(result["tools"]) == 1 - assert result["tools"][0]["type"] == "custom" - assert result["tools"][0]["name"] == "litellm_content_retrieve" - assert "input_schema" in result["tools"][0] - - -def test_compress_preserves_system_message(): - messages = [ - {"role": "system", "content": "System prompt. " * 500}, - {"role": "user", "content": "Large file content. " * 5000}, - {"role": "user", "content": "Fix the bug"}, - ] - result = litellm.compress( - messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000 - ) - assert result["messages"][0]["role"] == "system" - assert "System prompt" in result["messages"][0]["content"] - - -def test_compress_preserves_last_user_message(): - messages = [ - {"role": "user", "content": "Big context " * 5000}, - {"role": "user", "content": "Fix the bug in auth.py"}, - ] - result = litellm.compress( - messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000 - ) - last_user = [m for m in result["messages"] if m["role"] == "user"][-1] - assert "Fix the bug in auth.py" in last_user["content"] - - -def test_compress_preserves_last_assistant_message(): - messages = [ - {"role": "user", "content": "Big context " * 5000}, - {"role": "assistant", "content": "I'll help with that. " * 2000}, - {"role": "user", "content": "Now fix the bug"}, - ] - result = litellm.compress( - messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000 - ) - assistant_msgs = [m for m in result["messages"] if m["role"] == "assistant"] - assert len(assistant_msgs) >= 1 - # The last assistant message should be preserved (not stubbed) - last_assistant = assistant_msgs[-1] - assert "I'll help with that" in last_assistant["content"] - - -def test_cache_keys_match_stubs(): - messages = [ - {"role": "user", "content": "# auth.py\n" + "code " * 5000}, - {"role": "user", "content": "Fix it"}, - ] - result = litellm.compress( - messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000 - ) - if result["tools"]: - tool_desc = result["tools"][0]["function"]["description"] - for key in result["cache"]: - assert key in tool_desc - - -def test_compress_default_target(): - """compression_target defaults to compression_trigger // 2.""" - messages = [ - {"role": "user", "content": "content " * 5000}, - {"role": "user", "content": "query"}, - ] - result = litellm.compress( - messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=2000 - ) - # Should have compressed — target = 1000 - assert result["compressed_tokens"] <= result["original_tokens"] - - -def test_compress_nested_tool_result_extracts_text_only(): - messages = [ - {"role": "system", "content": [{"type": "text", "text": "System rules"}]}, - { - "role": "user", - "content": [ - {"type": "text", "text": "prefix"}, - { - "type": "tool_result", - "tool_use_id": "toolu_1", - "content": [ - {"type": "text", "text": "nested text fragment"}, - { - "type": "image_url", - "image_url": { - "url": "https://example.com/secret-tool.png", - }, - }, - ], - }, - { - "type": "image_url", - "image_url": {"url": "https://example.com/top.png"}, - }, - {"type": "text", "text": " " + ("irrelevant " * 3000)}, - ], - }, - { - "role": "user", - "content": [{"type": "text", "text": "final query that must remain"}], - }, - ] - - result = litellm.compress( - messages=messages, - model="claude-sonnet-4-20250514", - call_type=ANTHROPIC_CALL_TYPE, - compression_trigger=500, - compression_target=100, - ) - - cached_text = " ".join(result["cache"].values()) - assert "nested text fragment" in cached_text - assert "https://example.com/secret-tool.png" not in cached_text - assert "https://example.com/top.png" not in cached_text - - -def test_compress_default_call_type_is_completion(): - result = litellm.compress( - messages=[ - {"role": "user", "content": "Large context " * 4000}, - {"role": "user", "content": "query"}, - ], - model="gpt-4o", - compression_trigger=1000, - compression_target=500, - ) - - assert result["compressed_tokens"] <= result["original_tokens"] - assert isinstance(result["tools"], list) - - -def test_compress_forwards_embedding_model_params(monkeypatch): - captured = {} - - def fake_embedding_score_messages( - query, messages, model, cache=None, embedding_model_params=None - ): - captured["query"] = query - captured["model"] = model - captured["embedding_model_params"] = embedding_model_params - return [0.0] * len(messages) - - monkeypatch.setattr( - "litellm.compression.scoring.embedding_scorer.embedding_score_messages", - fake_embedding_score_messages, - ) - - result = litellm.compress( - messages=[ - {"role": "user", "content": "Authentication code " * 2000}, - {"role": "user", "content": "Fix auth"}, - ], - model="gpt-4o", - call_type=CALL_TYPE, - compression_trigger=1000, - embedding_model="text-embedding-3-small", - embedding_model_params={"api_base": "https://example-embeddings.test"}, - ) - - assert result["compressed_tokens"] <= result["original_tokens"] - assert captured["model"] == "text-embedding-3-small" - assert captured["embedding_model_params"] == { - "api_base": "https://example-embeddings.test" - } - - -def test_embedding_scorer_forwards_embedding_model_params(monkeypatch): - captured = {} - - class _MockResponse: - data = [ - {"embedding": [1.0, 0.0]}, - {"embedding": [1.0, 0.0]}, - {"embedding": [0.0, 1.0]}, - ] - - def fake_embedding(**kwargs): - captured.update(kwargs) - return _MockResponse() - - monkeypatch.setattr(litellm, "embedding", fake_embedding) - - scores = embedding_score_messages( - query="auth", - messages=[ - {"role": "user", "content": "auth code"}, - {"role": "user", "content": "cooking recipe"}, - ], - model="text-embedding-3-small", - embedding_model_params={"api_base": "https://example-embeddings.test"}, - ) - - assert len(scores) == 2 - assert captured["model"] == "text-embedding-3-small" - assert captured["api_base"] == "https://example-embeddings.test" - - # --------------------------------------------------------------------------- # Embedding scorer — integration test (skipped without API key) # --------------------------------------------------------------------------- @@ -458,210 +57,3 @@ def test_embedding_scorer(): ) assert result["compression_ratio"] > 0 assert len(result["cache"]) > 0 - - -@pytest.mark.parametrize( - "final_user_message, expected_content", - [ - ("How to cook?", "Unrelated cooking recipes "), - ("Fix auth", "Authentication code "), - ], -) -def test_simple_compression(final_user_message, expected_content): - messages = [ - {"role": "user", "content": "Authentication code " * 2000}, - {"role": "user", "content": "Unrelated cooking recipes " * 2000}, - {"role": "user", "content": final_user_message}, - ] - result = litellm.compress( - messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000 - ) - if expected_content == "Unrelated cooking recipes ": - assert "Unrelated cooking recipes " in result["messages"][1]["content"] - assert "Authentication code " not in result["messages"][0]["content"] - elif expected_content == "Authentication code ": - assert "Authentication code " in result["messages"][0]["content"] - assert "Unrelated cooking recipes " not in result["messages"][1]["content"] - else: - raise ValueError(f"Unexpected expected_content: {expected_content}") - - -def test_compress_anthropic_drops_irrelevant_tool_exchange_span(monkeypatch): - compress_module = importlib.import_module("litellm.compression.compress") - - def fake_bm25_score_messages(query, messages): - assert "final query" in query - assert len(messages) == 5 - # Prefer idx=0 and de-prioritize the tool exchange span (idx=1,2) - return [0.95, 0.01, 0.02, 0.8, 1.0] - - def fake_token_counter(model, messages=None, text=None): - if messages is not None: - return 1000 - if text is None: - return 0 - if "final query" in text: - return 50 - if "assistant_tail" in text: - return 20 - if "other_blob" in text: - return 220 - if "tool_payload_relevant" in text: - return 200 - if text == "": - return 1 - return 10 - - monkeypatch.setattr( - compress_module, "bm25_score_messages", fake_bm25_score_messages - ) - monkeypatch.setattr(compress_module, "token_counter", fake_token_counter) - - messages = [ - {"role": "user", "content": "other_blob " * 300}, - { - "role": "assistant", - "content": [ - { - "type": "tool_use", - "id": "toolu_drop", - "name": "litellm_content_retrieve", - "input": {"key": "message_1"}, - } - ], - }, - { - "role": "user", - "content": [ - { - "type": "tool_result", - "tool_use_id": "toolu_drop", - "content": [{"type": "text", "text": "tool_payload_relevant"}], - } - ], - }, - {"role": "assistant", "content": "assistant_tail"}, - {"role": "user", "content": "final query"}, - ] - - result = litellm.compress( - messages=messages, - model="claude-sonnet-4-20250514", - call_type=ANTHROPIC_CALL_TYPE, - compression_trigger=100, - compression_target=280, - ) - - # idx=1,2 should be dropped atomically (no orphan tool blocks left behind) - assert len(result["messages"]) == 3 - assert result["messages"][0]["role"] == "user" - assert "other_blob" in result["messages"][0]["content"] - assert result["messages"][1]["content"] == "assistant_tail" - assert result["messages"][2]["content"] == "final query" - assert result["cache"] == {} - - -def test_compress_anthropic_keeps_relevant_tool_exchange_span(monkeypatch): - compress_module = importlib.import_module("litellm.compression.compress") - - def fake_bm25_score_messages(query, messages): - assert "final query" in query - assert len(messages) == 5 - # Prefer the tool exchange span over idx=0 - return [0.05, 0.01, 0.92, 0.8, 1.0] - - def fake_token_counter(model, messages=None, text=None): - if messages is not None: - return 1000 - if text is None: - return 0 - if "final query" in text: - return 50 - if "assistant_tail" in text: - return 20 - if "other_blob" in text: - return 220 - if "tool_payload_relevant" in text: - return 200 - if text == "": - return 1 - return 10 - - monkeypatch.setattr( - compress_module, "bm25_score_messages", fake_bm25_score_messages - ) - monkeypatch.setattr(compress_module, "token_counter", fake_token_counter) - - messages = [ - {"role": "user", "content": "other_blob " * 300}, - { - "role": "assistant", - "content": [ - { - "type": "tool_use", - "id": "toolu_keep", - "name": "litellm_content_retrieve", - "input": {"key": "message_1"}, - } - ], - }, - { - "role": "user", - "content": [ - { - "type": "tool_result", - "tool_use_id": "toolu_keep", - "content": [{"type": "text", "text": "tool_payload_relevant"}], - } - ], - }, - {"role": "assistant", "content": "assistant_tail"}, - {"role": "user", "content": "final query"}, - ] - - result = litellm.compress( - messages=messages, - model="claude-sonnet-4-20250514", - call_type=ANTHROPIC_CALL_TYPE, - compression_trigger=100, - compression_target=280, - ) - - assert len(result["messages"]) == 5 - assert result["messages"][1]["role"] == "assistant" - assert result["messages"][2]["role"] == "user" - # idx=0 should be compressed instead - assert "litellm_content_retrieve" in result["messages"][0]["content"] - assert len(result["cache"]) == 1 - - -def test_compress_anthropic_malformed_tool_sequence_passes_through(): - messages = [ - {"role": "user", "content": "other_blob " * 300}, - { - "role": "assistant", - "content": [ - { - "type": "tool_use", - "id": "toolu_broken", - "name": "litellm_content_retrieve", - "input": {"key": "message_1"}, - } - ], - }, - {"role": "user", "content": [{"type": "text", "text": "missing tool_result"}]}, - {"role": "user", "content": "final query"}, - ] - - result = litellm.compress( - messages=messages, - model="claude-sonnet-4-20250514", - call_type=ANTHROPIC_CALL_TYPE, - compression_trigger=100, - compression_target=280, - ) - - assert result["messages"] == messages - assert result["cache"] == {} - assert result["tools"] == [] - assert result["compression_skipped_reason"] == "invalid_anthropic_tool_sequence" diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index 227fb48bb08..78728d6fd58 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -1,31 +1,12 @@ -import asyncio -import base64 -from datetime import datetime -import contextlib -import copy import json -import logging import os -from collections.abc import Mapping -from dataclasses import dataclass -from typing import Final -import httpx import pytest -import respx -from fastapi.testclient import TestClient -import urllib.parse -from importlib import import_module from unittest.mock import MagicMock, patch import litellm -from litellm import main as litellm_main -from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging -from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices, Usage async def _async_fake_bedrock_image_details(image_url): @@ -61,111 +42,6 @@ def add_api_keys_to_env(monkeypatch): monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False) -@pytest.fixture -def openai_api_response(): - mock_response_data = { - "id": "chatcmpl-B0W3vmiM78Xkgx7kI7dr7PC949DMS", - "choices": [ - { - "finish_reason": "stop", - "index": 0, - "logprobs": None, - "message": { - "content": "", - "refusal": None, - "role": "assistant", - "audio": None, - "function_call": None, - "tool_calls": None, - }, - } - ], - "created": 1739462947, - "model": "gpt-4o-mini-2024-07-18", - "object": "chat.completion", - "service_tier": "default", - "system_fingerprint": "fp_bd83329f63", - "usage": { - "completion_tokens": 1, - "prompt_tokens": 121, - "total_tokens": 122, - "completion_tokens_details": { - "accepted_prediction_tokens": 0, - "audio_tokens": 0, - "reasoning_tokens": 0, - "rejected_prediction_tokens": 0, - }, - "prompt_tokens_details": {"audio_tokens": 0, "cached_tokens": 0}, - }, - } - - return mock_response_data - - -def test_completion_missing_role(openai_api_response): - from openai import OpenAI - - from litellm.types.utils import ModelResponse - - client = OpenAI(api_key="test_api_key") - - mock_raw_response = MagicMock() - mock_raw_response.headers = { - "x-request-id": "123", - "openai-organization": "org-123", - "x-ratelimit-limit-requests": "100", - "x-ratelimit-remaining-requests": "99", - } - mock_raw_response.parse.return_value = ModelResponse(**openai_api_response) - - print(f"openai_api_response: {openai_api_response}") - - with patch.object( - client.chat.completions.with_raw_response, "create", MagicMock(return_value=mock_raw_response) - ) as mock_create: - litellm.completion( - model="gpt-4o-mini", - messages=[ - {"role": "user", "content": "Hey"}, - { - "content": "", - "tool_calls": [ - { - "id": "call_m0vFJjQmTH1McvaHBPR2YFwY", - "function": { - "arguments": '{"input": "dksjsdkjdhskdjshdskhjkhlk"}', - "name": "tool_name", - }, - "type": "function", - "index": 0, - }, - { - "id": "call_Vw6RaqV2n5aaANXEdp5pYxo2", - "function": { - "arguments": '{"input": "jkljlkjlkjlkjlk"}', - "name": "tool_name", - }, - "type": "function", - "index": 1, - }, - { - "id": "call_hBIKwldUEGlNh6NlSXil62K4", - "function": { - "arguments": '{"input": "jkjlkjlkjlkj;lj"}', - "name": "tool_name", - }, - "type": "function", - "index": 2, - }, - ], - }, - ], - client=client, - ) - - mock_create.assert_called_once() - - @pytest.mark.parametrize( "model", [ @@ -277,210 +153,6 @@ async def test_url_with_format_param(model, sync_mode, monkeypatch): assert "jpeg" not in json_str -@pytest.mark.parametrize("model", ["gpt-4o-mini"]) -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_url_with_format_param_openai(model, sync_mode): - from openai import AsyncOpenAI, OpenAI - - from litellm import acompletion, completion - - if sync_mode: - client = OpenAI() - else: - client = AsyncOpenAI() - - args = { - "model": model, - "messages": [ - { - "role": "user", - "content": [ - { - "type": "image_url", - "image_url": { - "url": "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png", - "format": "image/png", - }, - }, - {"type": "text", "text": "Describe this image"}, - ], - } - ], - } - with patch.object( - client.chat.completions.with_raw_response, "create" - ) as mock_client: - try: - if sync_mode: - response = completion(**args, client=client) - else: - response = await acompletion(**args, client=client) - print(response) - except Exception as e: - print(e) - - mock_client.assert_called() - - print(mock_client.call_args.kwargs) - - json_str = json.dumps(mock_client.call_args.kwargs) - - assert "format" not in json_str - - -def test_bedrock_latency_optimized_inference(): - from litellm.llms.custom_httpx.http_handler import HTTPHandler - - client = HTTPHandler() - with patch.object(client, "post") as mock_post: - try: - response = litellm.completion( - model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", - messages=[{"role": "user", "content": "Hello, how are you?"}], - performanceConfig={"latency": "optimized"}, - client=client, - ) - except Exception as e: - print(e) - - mock_post.assert_called_once() - json_data = json.loads(mock_post.call_args.kwargs["data"]) - assert json_data["performanceConfig"]["latency"] == "optimized" - - -@pytest.mark.parametrize( - ("custom_llm_provider", "model", "expected"), - [ - ("anthropic", "claude-sonnet-5", True), - ("bedrock", "us.anthropic.claude-sonnet-5-20260501-v1:0", True), - ("bedrock", "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123", True), - ("bedrock", "us.amazon.nova-2-lite-v1:0", False), - ("vertex_ai", "claude-sonnet-5", True), - ("vertex_ai", "gemini-3.8-flash", False), - ("azure_ai", "claude-sonnet-4-6", True), - ("azure_ai", "gpt-5.6", False), - ("openai", "gpt-5.6", False), - ("gemini", "gemini-3.8-flash", False), - ], -) -def test_is_claude_tool_target(custom_llm_provider: str, model: str, expected: bool): - assert litellm_main._is_claude_tool_target(custom_llm_provider=custom_llm_provider, model=model) is expected - - -@pytest.mark.parametrize("key", ["input_examples", "eager_input_streaming"]) -def test_drop_anthropic_only_tool_keys_strips_tool_and_function_levels(key: str): - tools = [ - {"type": "function", "name": "example_tool", key: True, "function": {"name": "example_tool", key: True}}, - "opaque_tool", - ] - - cleaned = litellm_main._drop_anthropic_only_tool_keys(tools=tools) - - assert cleaned == [ - {"type": "function", "name": "example_tool", "function": {"name": "example_tool"}}, - "opaque_tool", - ] - assert tools[0][key] is True - assert tools[0]["function"][key] is True - - -def test_completion_strips_eager_input_streaming_before_openai(respx_mock: respx.MockRouter, openai_api_response): - api_base: Final = "http://localhost:12346/v1" - mock_route: Final = respx_mock.post(url__regex=rf"{api_base}/chat/completions.*").mock( - return_value=httpx.Response(status_code=200, json=openai_api_response) - ) - - litellm.completion( - model="openai/gpt-5.6", - messages=[{"role": "user", "content": "Write the file"}], - tools=[ - { - "type": "function", - "function": {"name": "write_file", "parameters": {"type": "object", "properties": {}}}, - "eager_input_streaming": True, - } - ], - api_base=api_base, - api_key="fake_openai_api_key", - ) - - assert mock_route.called - sent_tool: Final = json.loads(respx_mock.calls[0].request.content)["tools"][0] - assert "eager_input_streaming" not in sent_tool - assert sent_tool["function"]["name"] == "write_file" - - -def test_custom_provider_with_extra_headers(): - from litellm.llms.custom_httpx.http_handler import HTTPHandler - - with patch.object( - litellm.llms.custom_httpx.http_handler.HTTPHandler, "post" - ) as mock_post: - response = litellm.completion( - model="custom/custom", - messages=[{"role": "user", "content": "Hello, how are you?"}], - headers={"X-Custom-Header": "custom-value"}, - api_base="https://example.com/api/v1", - ) - - mock_post.assert_called_once() - assert mock_post.call_args[1]["headers"]["X-Custom-Header"] == "custom-value" - - -def test_custom_provider_with_extra_body(): - from litellm.llms.custom_httpx.http_handler import HTTPHandler - - with patch.object( - litellm.llms.custom_httpx.http_handler.HTTPHandler, "post" - ) as mock_post: - response = litellm.completion( - model="custom/custom", - messages=[{"role": "user", "content": "Hello, how are you?"}], - extra_body={ - "X-Custom-BodyValue": "custom-value", - "X-Custom-BodyValue2": "custom-value2", - }, - api_base="https://example.com/api/v1", - ) - mock_post.assert_called_once() - - assert mock_post.call_args[1]["json"]["X-Custom-BodyValue"] == "custom-value" - assert mock_post.call_args[1]["json"] == { - "model": "custom", - "params": { - "prompt": ["Hello, how are you?"], - "max_tokens": None, - "temperature": None, - "top_p": None, - "top_k": None, - }, - "X-Custom-BodyValue": "custom-value", - "X-Custom-BodyValue2": "custom-value2", - } - - # test that extra_body is not passed if not provided - with patch.object( - litellm.llms.custom_httpx.http_handler.HTTPHandler, "post" - ) as mock_post: - response = litellm.completion( - model="custom/custom", - messages=[{"role": "user", "content": "Hello, how are you?"}], - api_base="https://example.com/api/v1", - ) - mock_post.assert_called_once() - assert mock_post.call_args[1]["json"] == { - "model": "custom", - "params": { - "prompt": ["Hello, how are you?"], - "max_tokens": None, - "temperature": None, - "top_p": None, - "top_k": None, - }, - } - - @pytest.fixture(autouse=True) def set_openrouter_api_key(): original_api_key = os.environ.get("OPENROUTER_API_KEY") @@ -490,3753 +162,3 @@ def set_openrouter_api_key(): os.environ["OPENROUTER_API_KEY"] = original_api_key else: del os.environ["OPENROUTER_API_KEY"] - - -@pytest.mark.asyncio -async def test_extra_body_with_fallback( - respx_mock: respx.MockRouter, set_openrouter_api_key, monkeypatch -): - """ - test regression for https://github.com/BerriAI/litellm/issues/8425. - - This was perhaps a wider issue with the acompletion function not passing kwargs such as extra_body correctly when fallbacks are specified. - """ - - # Save original state to restore after test - original_disable_aiohttp = litellm.disable_aiohttp_transport - - try: - # since this uses respx, we need to set use_aiohttp_transport to False - # Set both the global variable and environment variable to ensure it takes effect - litellm.disable_aiohttp_transport = True - monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") - # Flush cache to ensure no stale aiohttp clients are used - litellm.in_memory_llm_clients_cache.flush_cache() - - # Set up test parameters - model = "openrouter/deepseek/deepseek-chat" - messages = [{"role": "user", "content": "Hello, world!"}] - extra_body = { - "provider": { - "order": ["DeepSeek"], - "allow_fallbacks": False, - "require_parameters": True, - } - } - fallbacks = [{"model": "openrouter/google/gemini-flash-1.5-8b"}] - - # Set up mock to respond to any POST request to the OpenRouter endpoint - # This ensures it works for both primary and fallback models - mock_route = respx_mock.post("https://openrouter.ai/api/v1/chat/completions") - mock_route.return_value = httpx.Response( - 200, - json={ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": model, - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": "Hello from mocked response!", - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 9, - "completion_tokens": 12, - "total_tokens": 21, - }, - }, - ) - - response = await litellm.acompletion( - model=model, - messages=messages, - extra_body=extra_body, - fallbacks=fallbacks, - api_key="fake-openrouter-api-key", - ) - - # Verify the response - assert response is not None - assert ( - len(respx_mock.calls) > 0 - ), "Mock was not called - check if aiohttp transport is properly disabled" - - # Get the request from the mock - request: httpx.Request = respx_mock.calls[0].request - request_body = request.read() - request_body = json.loads(request_body) - - # Verify basic parameters - assert request_body["model"] == "deepseek/deepseek-chat" - assert request_body["messages"] == messages - - # Verify the extra_body parameters remain under the provider key - assert request_body["provider"]["order"] == ["DeepSeek"] - assert request_body["provider"]["allow_fallbacks"] is False - assert request_body["provider"]["require_parameters"] is True - finally: - # Restore original state to prevent test pollution - litellm.disable_aiohttp_transport = original_disable_aiohttp - litellm.in_memory_llm_clients_cache.flush_cache() - - -@pytest.mark.parametrize("env_base", ["OPENAI_BASE_URL", "OPENAI_API_BASE"]) -@pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=1) -async def test_openai_env_base( - respx_mock: respx.MockRouter, env_base, openai_api_response, monkeypatch -): - "This tests OpenAI env variables are honored, including legacy OPENAI_API_BASE" - # Ensure aiohttp transport is disabled to use httpx which respx can mock - litellm.disable_aiohttp_transport = True - - expected_base_url = "http://localhost:12345/v1" - - # Assign the environment variable based on env_base, and use a fake API key. - monkeypatch.setenv(env_base, expected_base_url) - monkeypatch.setenv("OPENAI_API_KEY", "fake_openai_api_key") - - model = "gpt-4o" - messages = [{"role": "user", "content": "Hello, how are you?"}] - - # Configure respx mock to intercept the request - mock_route = respx_mock.post( - url__regex=r"http://localhost:12345/v1/chat/completions.*" - ).mock( - return_value=httpx.Response( - status_code=200, - json={ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": model, - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": "Hello from mocked response!", - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 9, - "completion_tokens": 12, - "total_tokens": 21, - }, - }, - ) - ) - - try: - response = await litellm.acompletion(model=model, messages=messages) - - # verify we had a response - assert response.choices[0].message.content == "Hello from mocked response!" - - # Verify the mock was called - assert ( - mock_route.called - ), "Mock route was not called - request may have bypassed respx" - finally: - # Clean up to avoid affecting other tests - litellm.disable_aiohttp_transport = False - - -def build_database_url(username, password, host, dbname): - username_enc = urllib.parse.quote_plus(username) - password_enc = urllib.parse.quote_plus(password) - dbname_enc = urllib.parse.quote_plus(dbname) - return f"postgresql://{username_enc}:{password_enc}@{host}/{dbname_enc}" - - -def test_build_database_url(): - url = build_database_url("user@name", "p@ss:word", "localhost", "db/name") - assert url == "postgresql://user%40name:p%40ss%3Aword@localhost/db%2Fname" - - -def test_bedrock_llama(): - litellm._turn_on_debug() - from litellm.types.utils import CallTypes - from litellm.utils import return_raw_request - - model = "bedrock/invoke/us.meta.llama4-scout-17b-instruct-v1:0" - - request = return_raw_request( - endpoint=CallTypes.completion, - kwargs={ - "model": model, - "messages": [ - {"role": "user", "content": "hi"}, - ], - }, - ) - print(request) - - assert ( - request["raw_request_body"]["prompt"] - == "<|begin_of_text|><|start_header_id|>user<|end_header_id|>\n\nhi<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n" - ) - - -def _mocked_openai_chat_response(model: str) -> httpx.Response: - return httpx.Response( - status_code=200, - json={ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": model, - "choices": [ - { - "index": 0, - "message": { - "role": "assistant", - "content": "Hello from mocked response!", - }, - "finish_reason": "stop", - } - ], - "usage": { - "prompt_tokens": 9, - "completion_tokens": 12, - "total_tokens": 21, - }, - }, - ) - - -def test_return_raw_request_does_not_call_provider(respx_mock: respx.MockRouter): - """Regression for #33952: return_raw_request must transform without contacting the provider. - - Previously return_raw_request invoked the real endpoint with a fake key and relied on the - provider rejecting it, which sent an unintended inference request and (in the async proxy - route) blocked the event loop on provider I/O. - """ - from litellm.types.utils import CallTypes - from litellm.utils import return_raw_request - - model = "gpt-4o" - route = respx_mock.post("https://api.openai.com/v1/chat/completions").mock( - return_value=_mocked_openai_chat_response(model) - ) - - request = return_raw_request( - endpoint=CallTypes.completion, - kwargs={ - "model": model, - "messages": [{"role": "user", "content": "hi"}], - }, - ) - - assert route.call_count == 0 - assert request.get("error") is None - assert request["raw_request_body"]["model"] == model - assert request["raw_request_body"]["messages"] == [ - {"role": "user", "content": "hi"} - ] - - -def test_completion_forwards_verbosity_in_raw_request(respx_mock: respx.MockRouter): - """Regression test: completion() must forward the verbosity param to the provider request body.""" - from litellm.types.utils import CallTypes - from litellm.utils import return_raw_request - - model = "gpt-5.2" - messages = [{"role": "user", "content": "hi"}] - respx_mock.post("https://api.openai.com/v1/chat/completions").mock( - return_value=_mocked_openai_chat_response(model) - ) - - request = return_raw_request( - endpoint=CallTypes.completion, - kwargs={ - "model": model, - "messages": messages, - "verbosity": "high", - }, - ) - - assert request["raw_request_body"]["verbosity"] == "high" - assert request["raw_request_body"]["model"] == model - assert request["raw_request_body"]["messages"] == messages - - -@pytest.mark.asyncio -async def test_acompletion_forwards_verbosity_to_provider_request( - respx_mock: respx.MockRouter, monkeypatch -): - """Regression test: acompletion() must forward the verbosity param to the provider request body.""" - original_disable_aiohttp = litellm.disable_aiohttp_transport - try: - litellm.disable_aiohttp_transport = True - monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") - litellm.in_memory_llm_clients_cache.flush_cache() - - model = "gpt-5.2" - messages = [{"role": "user", "content": "hi"}] - mock_route = respx_mock.post("https://api.openai.com/v1/chat/completions").mock( - return_value=_mocked_openai_chat_response(model) - ) - - response = await litellm.acompletion( - model=model, - messages=messages, - verbosity="low", - api_key="fake-openai-api-key", - ) - - assert response.choices[0].message.content == "Hello from mocked response!" - assert mock_route.called - request_body = json.loads(respx_mock.calls[0].request.read()) - assert request_body["verbosity"] == "low" - assert request_body["model"] == model - assert request_body["messages"] == messages - finally: - litellm.disable_aiohttp_transport = original_disable_aiohttp - litellm.in_memory_llm_clients_cache.flush_cache() - - -def test_responses_api_bridge_check_strips_responses_prefix(): - """Test that responses_api_bridge_check strips 'responses/' prefix and sets mode.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 4096} - - model_info, model = responses_api_bridge_check( - model="responses/gpt-4-responses", - custom_llm_provider="openai", - ) - - assert model == "gpt-4-responses" - assert model_info["mode"] == "responses" - - -def test_responses_api_bridge_check_gpt_5_4_pro(): - """Test that gpt-5.4-pro routes through responses API bridge, not chat completions. - - Regression test for https://github.com/BerriAI/litellm/issues/23014 - gpt-5.4-pro is a responses-only model and must not be sent to /v1/chat/completions. - """ - from litellm.main import responses_api_bridge_check - - for model_name in ["gpt-5.4-pro", "gpt-5.4-pro-2026-03-05"]: - model_info, model = responses_api_bridge_check( - model=model_name, - custom_llm_provider="openai", - ) - assert ( - model_info.get("mode") == "responses" - ), f"{model_name} should have mode='responses', got '{model_info.get('mode')}'" - - -def test_responses_api_bridge_check_gpt_5_4_tools_plus_reasoning_routes_to_responses(): - """gpt-5.4 with both tools and reasoning_effort should route to Responses API.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.4", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort="xhigh", - ) - - assert model == "gpt-5.4" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_gpt_6_astra_tools_with_default_reasoning_routes_to_responses(): - from litellm.main import responses_api_bridge_check - - model_info, model = responses_api_bridge_check( - model="gpt-6-astra", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - ) - - assert model == "gpt-6-astra" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_gpt_5_5_tools_plus_reasoning_routes_to_responses(): - """gpt-5.5+ with both tools and reasoning_effort should route to Responses API.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.5-pro", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort="xhigh", - ) - - assert model == "gpt-5.5-pro" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_azure_gpt_5_4_tools_plus_reasoning_routes_to_responses(): - """Azure gpt-5.4 with both tools and reasoning_effort should route to Responses API.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.4", - custom_llm_provider="azure", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort="high", - ) - - assert model == "gpt-5.4" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_azure_gpt_5_4_tools_with_default_reasoning_routes_to_responses(): - """ - Azure gpt-5.4 with tools and UNSET reasoning_effort must bridge: OpenAI enables - reasoning by default for gpt-5.4+, and Chat Completions rejects function tools - whenever reasoning is on. - """ - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.4", - custom_llm_provider="azure", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - ) - - assert model == "gpt-5.4" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_gpt_5_4_tools_with_default_reasoning_routes_to_responses(): - """ - gpt-5.4 with tools and UNSET reasoning_effort must bridge: OpenAI enables reasoning - by default for gpt-5.4+, and Chat Completions rejects function tools whenever - reasoning is on ("use /v1/responses or set reasoning_effort to 'none'"). - """ - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.4", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - ) - - assert model == "gpt-5.4" - assert model_info.get("mode") == "responses" - - -@pytest.mark.parametrize( - "model_name, expected_mode", - [ - pytest.param("gpt-5.6-sol", "responses", id="above-boundary-bridges"), - pytest.param("gpt-5.1", None, id="below-boundary-stays-chat"), - ], -) -def test_responses_api_bridge_check_gpt_5_6_tools_with_default_reasoning_routes_to_responses( - monkeypatch, model_name, expected_mode -): - """ - gpt-5.6 must bridge on function tools alone. The bridge used to require an explicit - reasoning_effort, so a gpt-5.6 call carrying tools and no effort was rejected with - "Function tools with reasoning_effort are not supported for gpt-5.6-sol in - /v1/chat/completions". - - Paired with a model below the gpt-5.4 boundary, which must still stay on chat. The - gate parses the version and drops any suffix, so the family members bridge - identically and only the boundary distinguishes behaviour. - """ - import litellm - from litellm.main import responses_api_bridge_check - - monkeypatch.delenv("OPENAI_BASE_URL", raising=False) - monkeypatch.delenv("OPENAI_API_BASE", raising=False) - monkeypatch.setattr(litellm, "api_base", None) - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model=model_name, - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - ) - - assert model == model_name - assert model_info.get("mode") == expected_mode - - -def test_responses_api_bridge_check_gpt_5_4_tools_with_reasoning_none_stays_chat(): - """ - Explicit reasoning_effort "none" is OpenAI's documented escape hatch that keeps - function tools servable on Chat Completions; the bridge must not fire. - """ - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.4", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort="none", - ) - - assert model == "gpt-5.4" - assert model_info.get("mode") != "responses" - - -def test_responses_api_bridge_check_reasoning_none_with_summary_still_routes_to_responses(): - """A reasoning summary is Responses-only regardless of effort value.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.4", - custom_llm_provider="openai", - reasoning_effort="none", - reasoning_summary="detailed", - ) - - assert model == "gpt-5.4" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_gpt_5_4_custom_tools_only_stays_chat(): - """ - Chat Completions serves custom (grammar) tools natively with reasoning on; only - FUNCTION tools trigger the OpenAI rejection. Custom-only requests must stay on chat - so responses keep the native custom tool_call shape instead of the bridge's - function-shaped mapping. - """ - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "custom", "custom": {"name": "ApplyPatch", "description": "V4A patch"}}], - reasoning_effort=None, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") != "responses" - - -def test_responses_api_bridge_check_gpt_5_4_mixed_function_and_custom_tools_routes_to_responses(): - """One function tool in the mix is enough to make chat unservable with reasoning on.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[ - {"type": "custom", "custom": {"name": "ApplyPatch"}}, - {"type": "function", "function": {"name": "shell"}}, - ], - reasoning_effort=None, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_gpt_5_4_flat_function_tool_routes_to_responses(): - """Responses-style flat function tool defs still count as function tools.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "name": "shell", "parameters": {"type": "object"}}], - reasoning_effort=None, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") == "responses" - - -@pytest.mark.parametrize( - "custom_llm_provider, model_name, api_base", - [ - pytest.param("openai", "gpt-5.6", None, id="openai"), - pytest.param("azure_ai", "gpt-6-astra", "https://myproject.services.ai.azure.com", id="azure-ai-foundry"), - ], -) -def test_responses_api_bridge_check_function_tool_without_body_stays_chat( - monkeypatch, custom_llm_provider, model_name, api_base -): - import litellm - from litellm.main import responses_api_bridge_check - - monkeypatch.delenv("OPENAI_BASE_URL", raising=False) - monkeypatch.delenv("OPENAI_API_BASE", raising=False) - monkeypatch.setattr(litellm, "api_base", None) - - model_info, model = responses_api_bridge_check( - model=model_name, - custom_llm_provider=custom_llm_provider, - tools=[{"type": "function"}], - reasoning_effort=None, - api_base=api_base, - ) - - assert model == model_name - assert model_info.get("mode") != "responses" - - -def test_responses_api_bridge_check_dict_effort_none_stays_chat(): - """The escape hatch must honor litellm's dict form: {"effort": "none"} means reasoning off.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort={"effort": "none"}, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") != "responses" - - -def test_responses_api_bridge_check_dict_effort_active_routes_to_responses(): - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort={"effort": "low"}, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_dict_effort_none_with_summary_routes_to_responses(): - """A summary inside the dict form is Responses-only even when effort is none.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort={"effort": "none", "summary": "concise"}, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") == "responses" - - -@pytest.mark.parametrize("blank_api_base", [None, "", " ", "\t"]) -def test_responses_api_bridge_check_blank_api_base_is_default_openai(blank_api_base): - """ - A blank api_base (None, empty, or whitespace) resolves to the default OpenAI - endpoint downstream, which enforces the reasoning+tools constraint, so gpt-5.4+ - function-tool requests with unset reasoning_effort must still auto-bridge. - """ - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - api_base=blank_api_base, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_custom_api_base_with_unset_effort_stays_chat(): - """ - Chat-only OpenAI-compatible backends registered under the openai provider with a - custom api_base and gpt-5.4+ model names serve tools-without-reasoning fine and - have no /responses route; the unset-effort arm must not reroute them. - """ - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - api_base="http://vllm.internal:8000/v1", - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") != "responses" - - -def test_responses_api_bridge_check_custom_api_base_via_global_with_unset_effort_stays_chat(monkeypatch): - """ - A custom base set through the litellm.api_base global (not the call arg) is resolved the - same way the chat handler resolves it, so the unset-effort arm must not reroute a chat-only - backend to a /responses route it lacks. Regression guard: the gate previously inspected only - the call-level api_base and bridged these requests. - """ - import litellm - from litellm.main import responses_api_bridge_check - - monkeypatch.setattr(litellm, "api_base", "http://vllm.internal:8000/v1") - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - api_base=None, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") != "responses" - - -@pytest.mark.parametrize("env_var", ["OPENAI_BASE_URL", "OPENAI_API_BASE"]) -def test_responses_api_bridge_check_custom_api_base_via_env_with_unset_effort_stays_chat(monkeypatch, env_var): - """ - A custom base set via OPENAI_BASE_URL/OPENAI_API_BASE env is resolved identically to the chat - handler, so the unset-effort arm leaves the request on chat instead of bridging it. - """ - import litellm - from litellm.main import responses_api_bridge_check - - monkeypatch.setattr(litellm, "api_base", None) - monkeypatch.delenv("OPENAI_BASE_URL", raising=False) - monkeypatch.delenv("OPENAI_API_BASE", raising=False) - monkeypatch.setenv(env_var, "http://vllm.internal:8000/v1") - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - api_base=None, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") != "responses" - - -@pytest.mark.parametrize( - "api_base", - [ - "https://southcentralus.privatelink.api.openai.com/v1", - "https://privatelink.corp.api.openai.com/v1", - "https://api.openai.com:443/v1", - "https://api.openai.com/v1/", - "HTTPS://API.OPENAI.COM/v1", - ], -) -def test_responses_api_bridge_check_openai_backed_custom_api_base_with_unset_effort_routes_to_responses(api_base): - """ - A custom api_base whose host is api.openai.com or a subdomain of it (a PrivateLink hostname, a - port-qualified or trailing-slash default) still reaches the real OpenAI backend, which rejects - function tools with reasoning on Chat Completions, so the unset-effort arm must bridge exactly as - it does for the literal default URL. Regression guard for GH #39353. - """ - from litellm.main import responses_api_bridge_check - - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - api_base=api_base, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") == "responses" - - -@pytest.mark.parametrize( - "api_base", - [ - "https://api.openai.com.evil.example/v1", - "https://notapi.openai.com/v1", - "https://gateway.example/v1?upstream=api.openai.com", - "https://openai.internal.example/api.openai.com/v1", - ], -) -def test_responses_api_bridge_check_lookalike_custom_api_base_with_unset_effort_stays_chat(api_base): - """Only the host decides: api.openai.com appearing elsewhere in the URL is still a foreign backend.""" - from litellm.main import responses_api_bridge_check - - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - api_base=api_base, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") != "responses" - - -def test_responses_api_bridge_check_privatelink_api_base_via_env_with_unset_effort_routes_to_responses(monkeypatch): - """A PrivateLink base set through OPENAI_BASE_URL resolves the way the chat handler's does and still bridges.""" - import litellm - from litellm.main import responses_api_bridge_check - - monkeypatch.setattr(litellm, "api_base", None) - monkeypatch.delenv("OPENAI_API_BASE", raising=False) - monkeypatch.setenv("OPENAI_BASE_URL", "https://southcentralus.privatelink.api.openai.com/v1") - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - api_base=None, - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_custom_api_base_with_explicit_effort_still_routes(): - """Explicit reasoning_effort keeps its pre-existing bridging behavior on any api_base.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.6", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort="high", - api_base="http://vllm.internal:8000/v1", - ) - - assert model == "gpt-5.6" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_azure_with_api_base_and_unset_effort_routes(): - """Azure OpenAI always sets api_base and does enforce the constraint; keep bridging.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.4", - custom_llm_provider="azure", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - api_base="https://myresource.openai.azure.com", - ) - - assert model == "gpt-5.4" - assert model_info.get("mode") == "responses" - - -_FOUNDRY_API_BASE: Final = "https://myproject.services.ai.azure.com" -_FOUNDRY_FUNCTION_TOOL: Final = ({"type": "function", "function": {"name": "get_weather"}},) - - -@pytest.mark.parametrize( - "model_name, api_base, reasoning_effort", - [ - pytest.param("gpt-6-astra", _FOUNDRY_API_BASE, None, id="gpt-6-unset-effort"), - pytest.param("gpt-6-astra", _FOUNDRY_API_BASE, "low", id="gpt-6-explicit-effort"), - pytest.param("gpt-6-astra", "https://myresource.openai.azure.com", None, id="gpt-6-azure-openai-host"), - pytest.param("gpt-5.6-sol", _FOUNDRY_API_BASE, "low", id="gpt-5.6-explicit-effort"), - pytest.param("gpt-5.6-sol", _FOUNDRY_API_BASE, {"effort": "high"}, id="gpt-5.6-explicit-effort-dict"), - ], -) -def test_responses_api_bridge_check_azure_ai_foundry_rejected_tools_route_to_responses( - model_name, api_base, reasoning_effort -): - from litellm.main import responses_api_bridge_check - - model_info, model = responses_api_bridge_check( - model=model_name, - custom_llm_provider="azure_ai", - tools=_FOUNDRY_FUNCTION_TOOL, - reasoning_effort=reasoning_effort, - api_base=api_base, - ) - - assert model == model_name - assert model_info.get("mode") == "responses" - - -@pytest.mark.parametrize( - "model_name, api_base, reasoning_effort", - [ - pytest.param("gpt-6-astra", _FOUNDRY_API_BASE, "none", id="explicit-none-stays-chat"), - pytest.param("gpt-5.6-sol", _FOUNDRY_API_BASE, None, id="gpt-5.6-unset-effort-stays-chat"), - pytest.param("gpt-5.6-sol", _FOUNDRY_API_BASE, "none", id="gpt-5.6-explicit-none-stays-chat"), - pytest.param("gpt-5.5", _FOUNDRY_API_BASE, "high", id="gpt-5.5-explicit-effort-stays-chat"), - pytest.param("gpt-5.4-mini", _FOUNDRY_API_BASE, None, id="gpt-5.4-mini-unset-effort-stays-chat"), - pytest.param("gpt-5.4-mini", _FOUNDRY_API_BASE, "low", id="gpt-5.4-mini-explicit-effort-stays-chat"), - pytest.param("gpt-6-astra", "https://myproject.models.ai.azure.com", None, id="serverless-host-stays-chat"), - pytest.param("Mistral-large-2411", _FOUNDRY_API_BASE, None, id="non-gpt-5-model-stays-chat"), - pytest.param("claude-opus-4-1", _FOUNDRY_API_BASE, None, id="claude-on-foundry-stays-chat"), - ], -) -def test_responses_api_bridge_check_azure_ai_without_foundry_responses_route_stays_chat( - model_name, api_base, reasoning_effort -): - from litellm.main import responses_api_bridge_check - - model_info, model = responses_api_bridge_check( - model=model_name, - custom_llm_provider="azure_ai", - tools=_FOUNDRY_FUNCTION_TOOL, - reasoning_effort=reasoning_effort, - api_base=api_base, - ) - - assert model == model_name - assert model_info.get("mode") != "responses" - - -def test_responses_api_bridge_check_older_gpt_5_tools_without_reasoning_stays_chat(): - """Pre-5.4 GPT-5 names keep the old boundary: tools alone never bridge.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.1", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort=None, - ) - - assert model == "gpt-5.1" - assert model_info.get("mode") != "responses" - - -def test_responses_api_bridge_check_gpt_5_4_reasoning_summary_without_tools_routes_to_responses(): - """gpt-5.4+ with reasoning_effort + reasoningSummary but no tools should bridge (AI SDK).""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5.4", - custom_llm_provider="openai", - tools=None, - reasoning_effort="medium", - reasoning_summary="auto", - ) - - assert model == "gpt-5.4" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_gpt_5_reasoning_summary_routes_to_responses(): - """Bare ``gpt-5`` with reasoning_effort + reasoningSummary should bridge (not 5.4+).""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5", - custom_llm_provider="openai", - tools=None, - reasoning_effort="medium", - reasoning_summary="auto", - ) - - assert model == "gpt-5" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_gpt_5_tools_without_summary_stays_chat(): - """gpt-5 with tools + reasoning_effort but no summary should stay on chat.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 128000} - model_info, model = responses_api_bridge_check( - model="gpt-5", - custom_llm_provider="openai", - tools=[{"type": "function", "function": {"name": "get_capital"}}], - reasoning_effort="medium", - reasoning_summary=None, - ) - - assert model == "gpt-5" - assert model_info.get("mode") != "responses" - - -@patch("litellm.completion_extras.responses_api_bridge.completion") -def test_gpt_5_4_responses_bridge_preserves_reasoning_summary_dict( - mock_responses_completion, -): - """When routed to Responses, preserve reasoning_effort summary dict.""" - mock_responses_completion.return_value = MagicMock() - - import litellm - - litellm.completion( - model="gpt-5.4", - messages=[{"role": "user", "content": "What is the capital of France?"}], - tools=[ - { - "type": "function", - "function": { - "name": "get_capital", - "description": "Get the capital of a country", - "parameters": { - "type": "object", - "properties": {"country": {"type": "string"}}, - }, - }, - } - ], - reasoning_effort={"effort": "xhigh", "summary": "detailed"}, - api_key="fake-key", - ) - - assert mock_responses_completion.called is True - optional_params = mock_responses_completion.call_args.kwargs["optional_params"] - assert optional_params["reasoning_effort"] == { - "effort": "xhigh", - "summary": "detailed", - } - - -@pytest.mark.parametrize("reasoning_effort", ["high", {"effort": "high"}]) -def test_responses_bridge_preserves_reasoning_effort_with_drop_params( - reasoning_effort, - restore_model_registry, - respx_mock: respx.MockRouter, - monkeypatch: pytest.MonkeyPatch, -): - monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) - response_body: Final = { - "id": "resp_test", - "object": "response", - "created_at": 1734366691, - "status": "completed", - "model": "test-responses-bridge", - "output": [ - { - "type": "message", - "id": "msg_1", - "status": "completed", - "role": "assistant", - "content": [{"type": "output_text", "text": "Done.", "annotations": []}], - } - ], - "parallel_tool_calls": True, - "usage": { - "input_tokens": 1, - "output_tokens": 1, - "total_tokens": 2, - "output_tokens_details": {"reasoning_tokens": 0}, - }, - "error": None, - "incomplete_details": None, - "instructions": None, - "metadata": None, - "temperature": None, - "tool_choice": "auto", - "tools": [], - "top_p": None, - "max_output_tokens": None, - "previous_response_id": None, - "reasoning": None, - "truncation": None, - "user": None, - } - response_route: Final = respx_mock.post("https://api.perplexity.ai/v1/responses").respond(json=response_body) - model: Final = "perplexity/test-responses-bridge" - litellm.register_model( - { - model: { - "litellm_provider": "perplexity", - "mode": "responses", - "supports_reasoning": False, - "input_cost_per_token": 0.0, - "output_cost_per_token": 0.0, - } - }, - persist_across_reloads=False, - ) - - litellm.completion( - model=model, - messages=[{"role": "user", "content": "hello"}], - reasoning_effort=reasoning_effort, - drop_params=True, - api_key="fake-key", - api_base="https://api.perplexity.ai", - ) - - request_body: Final = json.loads(response_route.calls[0].request.content) - assert request_body["reasoning"] == {"effort": "high"} - - -_FOUNDRY_RESPONSES_FUNCTION_CALL_BODY: Final = { - "id": "resp_foundry", - "object": "response", - "created_at": 1789852145, - "status": "completed", - "model": "gpt-6-astra", - "output": [ - { - "id": "fc_1", - "type": "function_call", - "status": "completed", - "arguments": '{"city":"Paris"}', - "call_id": "call_1", - "name": "get_weather", - } - ], - "parallel_tool_calls": True, - "usage": { - "input_tokens": 53, - "output_tokens": 18, - "total_tokens": 71, - "output_tokens_details": {"reasoning_tokens": 0}, - }, - "error": None, - "incomplete_details": None, - "instructions": None, - "metadata": {}, - "temperature": 1.0, - "tool_choice": "auto", - "tools": [], - "top_p": 1.0, - "max_output_tokens": 200, - "previous_response_id": None, - "reasoning": {"effort": "medium", "summary": None}, - "truncation": "disabled", - "user": None, -} - - -def test_completion_bridges_azure_ai_foundry_gpt_5_4_plus_function_tools_to_responses( - respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch -): - monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) - responses_route: Final = respx_mock.post(f"{_FOUNDRY_API_BASE}/openai/v1/responses").respond( - json=_FOUNDRY_RESPONSES_FUNCTION_CALL_BODY - ) - - response: Final = litellm.completion( - model="azure_ai/gpt-6-astra", - messages=[{"role": "user", "content": "What is the weather in Paris? Use the tool."}], - tools=[ - { - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather for a city", - "parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}, - }, - } - ], - max_tokens=200, - api_base=_FOUNDRY_API_BASE, - api_key="fake-foundry-key", - ) - - assert [str(call.request.url) for call in respx_mock.calls] == [f"{_FOUNDRY_API_BASE}/openai/v1/responses"] - request: Final = responses_route.calls[0].request - request_body: Final = json.loads(request.content) - assert request_body["tools"][0]["type"] == "function" - assert request_body["tools"][0]["name"] == "get_weather" - assert request.headers["api-key"] == "fake-foundry-key" - assert response.choices[0].finish_reason == "tool_calls" - assert response.choices[0].message.tool_calls[0].function.name == "get_weather" - - -@pytest.mark.parametrize( - "model, model_info, expected_model_param, expected_base_model_param", - [ - ("gemini/gemini-3.1-pro", None, "gemini-3.1-pro", None), - ( - "gemini/gemini-3.1-pro", - {"base_model": "gemini-3.1-pro-preview"}, - "gemini-3.1-pro", - "gemini-3.1-pro-preview", - ), - ], -) -def test_completion_optional_params_base_model( - model: str, - model_info: dict | None, - expected_model_param: str, - expected_base_model_param: str | None, -): - """``model_info.base_model`` must reach ``get_optional_params`` as ``base_model`` - (an additive capability hint), without overwriting ``model`` with the label. - - Regression for #29618: overwriting ``model`` with a friendly ``base_model`` - label made Bedrock drop ``tools``/``tool_choice`` under ``drop_params``.""" - with patch("litellm.main.get_optional_params") as mock_get_optional_params: - mock_get_optional_params.return_value = MagicMock() - - import litellm - - kwargs = { - "model": model, - "messages": [{"role": "user", "content": "What is the capital of France?"}], - "api_key": "fake-key", - "mock_response": "Hey, how's it going?", - } - if model_info is not None: - kwargs["model_info"] = model_info - - litellm.completion(**kwargs) - - assert mock_get_optional_params.called is True - call_kwargs = mock_get_optional_params.call_args.kwargs - assert call_kwargs["model"] == expected_model_param - assert call_kwargs["base_model"] == expected_base_model_param - - -@patch("litellm.completion_extras.responses_api_bridge.completion") -def test_gpt_5_4_responses_bridge_merges_reasoning_summary_kwarg_without_tools( - mock_responses_completion, -): - """reasoningSummary without tools should route and merge into reasoning_effort dict.""" - mock_responses_completion.return_value = MagicMock() - - import litellm - - litellm.completion( - model="gpt-5.4", - messages=[{"role": "user", "content": "ok"}], - reasoning_effort="medium", - reasoningSummary="auto", - api_key="fake-key", - ) - - assert mock_responses_completion.called is True - optional_params = mock_responses_completion.call_args.kwargs["optional_params"] - assert optional_params["reasoning_effort"] == { - "effort": "medium", - "summary": "auto", - } - assert "reasoningSummary" not in optional_params - assert "reasoning_summary" not in optional_params - - -@patch("litellm.completion_extras.responses_api_bridge.completion") -def test_responses_bridge_preserves_reasoning_summary_without_effort( - mock_responses_completion, -): - """Reasoning summary should survive responses routing even without effort.""" - mock_responses_completion.return_value = MagicMock() - - import litellm - - with patch.object(litellm, "route_all_chat_openai_to_responses", True): - litellm.completion( - model="gpt-4o", - messages=[{"role": "user", "content": "ok"}], - reasoningSummary="auto", - api_key="fake-key", - ) - - assert mock_responses_completion.called is True - optional_params = mock_responses_completion.call_args.kwargs["optional_params"] - assert optional_params["reasoning_effort"] == {"summary": "auto"} - assert "reasoningSummary" not in optional_params - assert "reasoning_summary" not in optional_params - - -@patch("litellm.completion_extras.responses_api_bridge.completion") -def test_gpt_5_responses_bridge_tools_and_reasoning_summary( - mock_responses_completion, -): - """Bare gpt-5 with tools + reasoningSummary should bridge (OpenCode-style).""" - mock_responses_completion.return_value = MagicMock() - - import litellm - - litellm.completion( - model="gpt-5", - messages=[{"role": "user", "content": "ok"}], - tools=[ - { - "type": "function", - "function": { - "name": "apply_patch", - "parameters": {"type": "object", "properties": {}}, - }, - } - ], - tool_choice="auto", - reasoning_effort="medium", - reasoningSummary="auto", - stream=True, - api_key="fake-key", - ) - - assert mock_responses_completion.called is True - optional_params = mock_responses_completion.call_args.kwargs["optional_params"] - assert optional_params.get("reasoning_effort") == { - "effort": "medium", - "summary": "auto", - } - - -def test_responses_api_bridge_check_handles_exception(): - """Test that responses_api_bridge_check handles exceptions and still processes responses/ models.""" - from litellm.main import responses_api_bridge_check - - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.side_effect = Exception("Model not found") - - model_info, model = responses_api_bridge_check( - model="responses/custom-model", custom_llm_provider="custom" - ) - - assert model == "custom-model" - assert model_info["mode"] == "responses" - - -def test_responses_api_bridge_check_global_flag_routes_openai(): - """When route_all_chat_openai_to_responses is True, any OpenAI model routes to responses.""" - from litellm.main import responses_api_bridge_check - - with patch.object(litellm, "route_all_chat_openai_to_responses", True): - model_info, model = responses_api_bridge_check( - model="gpt-4o", - custom_llm_provider="openai", - ) - - assert model == "gpt-4o" - assert model_info.get("mode") == "responses" - - -def test_responses_api_bridge_check_global_flag_does_not_affect_azure(): - """route_all_chat_openai_to_responses should not affect Azure models.""" - from litellm.main import responses_api_bridge_check - - with patch.object(litellm, "route_all_chat_openai_to_responses", True): - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 4096} - model_info, model = responses_api_bridge_check( - model="gpt-4o", - custom_llm_provider="azure", - ) - - assert model_info.get("mode") != "responses" - - -def test_responses_api_bridge_check_global_flag_default_false(): - """By default, route_all_chat_openai_to_responses is False and doesn't affect routing.""" - from litellm.main import responses_api_bridge_check - - with patch.object(litellm, "route_all_chat_openai_to_responses", False): - with patch("litellm.main._get_model_info_helper") as mock_get_model_info: - mock_get_model_info.return_value = {"max_tokens": 4096} - model_info, model = responses_api_bridge_check( - model="gpt-4o", - custom_llm_provider="openai", - ) - - assert model_info.get("mode") != "responses" - - -@pytest.mark.asyncio -async def test_async_mock_delay(): - """Use asyncio await for mock delay on acompletion""" - import time - - from litellm import acompletion - - start_time = time.time() - result = await acompletion( - model="gpt-3.5-turbo", - messages=[{"role": "user", "content": "Hey, how's it going?"}], - mock_delay=0.01, - mock_response="Hello world", - ) - end_time = time.time() - delay = end_time - start_time - assert delay >= 0.01 - - -def test_stream_chunk_builder_keeps_tool_calls_carried_only_by_a_later_choice_of_a_multi_choice_chunk(): - from litellm import stream_chunk_builder - from litellm.types.utils import ( - ChatCompletionDeltaToolCall, - Delta, - Function, - ModelResponseStream, - StreamingChoices, - ) - - def chunk(choices: list[StreamingChoices]) -> ModelResponseStream: - return ModelResponseStream( - id="chatcmpl-multi-choice", - created=1751934860, - model="gpt-4.1-mini", - object="chat.completion.chunk", - choices=choices, - ) - - chunks = [ - chunk( - [ - StreamingChoices(index=0, delta=Delta(role="assistant", content="hello")), - StreamingChoices( - index=1, - delta=Delta( - role="assistant", - tool_calls=[ - ChatCompletionDeltaToolCall( - id="call_1", - index=0, - type="function", - function=Function(name="lookup_fruit", arguments='{"fruit":'), - ) - ], - ), - ), - ] - ), - chunk( - [ - StreamingChoices(index=0, delta=Delta(content=" world"), finish_reason="stop"), - StreamingChoices( - index=1, - delta=Delta( - tool_calls=[ChatCompletionDeltaToolCall(index=0, function=Function(arguments='"kiwi"}'))] - ), - finish_reason="tool_calls", - ), - ] - ), - ] - - response = stream_chunk_builder(chunks=chunks) - - tool_calls = response.choices[0].message.tool_calls - assert tool_calls is not None - assert [(call.id, call.function.name, call.function.arguments) for call in tool_calls] == [ - ("call_1", "lookup_fruit", '{"fruit":"kiwi"}') - ] - - -def test_stream_chunk_builder_thinking_blocks(): - from litellm import stream_chunk_builder - from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices - - chunks = [ - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - reasoning_content="I need to summar", - thinking_blocks=[ - { - "type": "thinking", - "thinking": "I need to summar", - "signature": None, - } - ], - provider_specific_fields={ - "thinking_blocks": [ - { - "type": "thinking", - "thinking": "I need to summar", - "signature": None, - } - ] - }, - content="", - role="assistant", - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - reasoning_content="ize the previous agent's thinking process into a", - thinking_blocks=[ - { - "type": "thinking", - "thinking": "ize the previous agent's thinking process into a", - "signature": None, - } - ], - provider_specific_fields={ - "thinking_blocks": [ - { - "type": "thinking", - "thinking": "ize the previous agent's thinking process into a", - "signature": None, - } - ] - }, - content="", - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - reasoning_content=" short description. Based on the input data provide", - thinking_blocks=[ - { - "type": "thinking", - "thinking": " short description. Based on the input data provide", - "signature": None, - } - ], - provider_specific_fields={ - "thinking_blocks": [ - { - "type": "thinking", - "thinking": " short description. Based on the input data provide", - "signature": None, - } - ] - }, - content="", - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - reasoning_content="d, it seems the agent was planning to refine their search", - thinking_blocks=[ - { - "type": "thinking", - "thinking": "d, it seems the agent was planning to refine their search", - "signature": None, - } - ], - provider_specific_fields={ - "thinking_blocks": [ - { - "type": "thinking", - "thinking": "d, it seems the agent was planning to refine their search", - "signature": None, - } - ] - }, - content="", - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - reasoning_content=" to focus more on technical aspects of home automation and home", - thinking_blocks=[ - { - "type": "thinking", - "thinking": " to focus more on technical aspects of home automation and home", - "signature": None, - } - ], - provider_specific_fields={ - "thinking_blocks": [ - { - "type": "thinking", - "thinking": " to focus more on technical aspects of home automation and home", - "signature": None, - } - ] - }, - content="", - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - reasoning_content=" energy system management.\n\nI'll create a brief", - thinking_blocks=[ - { - "type": "thinking", - "thinking": " energy system management.\n\nI'll create a brief", - "signature": None, - } - ], - provider_specific_fields={ - "thinking_blocks": [ - { - "type": "thinking", - "thinking": " energy system management.\n\nI'll create a brief", - "signature": None, - } - ] - }, - content="", - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - reasoning_content=" summary of what the agent was doing.", - thinking_blocks=[ - { - "type": "thinking", - "thinking": " summary of what the agent was doing.", - "signature": None, - } - ], - provider_specific_fields={ - "thinking_blocks": [ - { - "type": "thinking", - "thinking": " summary of what the agent was doing.", - "signature": None, - } - ] - }, - content="", - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=0, - delta=Delta( - reasoning_content="", - thinking_blocks=[ - { - "type": "thinking", - "thinking": "", - "signature": "ErUBCkYIBRgCIkAKBSMkB2+MBF643wiWxlERsGXVdlhbPx9lnTIbygzjFIeZ5uhTV+HNWDon9vQV4hmXvAKwQfwS8vkNFB366l05Egzt2U18IpRrZRyQn1UaDDdYvKHYP8Ps1IbWjSIw8eSYOU9gtqNcwR6D0wY7iOPx2GliDEatLI5rSs96CByoTIoADL2M5bX8KP0jEpbHKh0ccYryigdH/3J8EiFt/BmGUceVASP5l9r22dFWiBgC", - } - ], - provider_specific_fields={ - "thinking_blocks": [ - { - "type": "thinking", - "thinking": "", - "signature": "ErUBCkYIBRgCIkAKBSMkB2+MBF643wiWxlERsGXVdlhbPx9lnTIbygzjFIeZ5uhTV+HNWDon9vQV4hmXvAKwQfwS8vkNFB366l05Egzt2U18IpRrZRyQn1UaDDdYvKHYP8Ps1IbWjSIw8eSYOU9gtqNcwR6D0wY7iOPx2GliDEatLI5rSs96CByoTIoADL2M5bX8KP0jEpbHKh0ccYryigdH/3J8EiFt/BmGUceVASP5l9r22dFWiBgC", - } - ] - }, - content="", - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=1, - delta=Delta( - provider_specific_fields=None, - content='{"a', - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=1, - delta=Delta( - provider_specific_fields=None, - content='gent_doing"', - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=1, - delta=Delta( - provider_specific_fields=None, - content=': "Re', - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=1, - delta=Delta( - provider_specific_fields=None, - content="searching", - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=1, - delta=Delta( - provider_specific_fields=None, - content=" technic", - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=1, - delta=Delta( - provider_specific_fields=None, - content="al aspect", - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=1, - delta=Delta( - provider_specific_fields=None, - content="s of home au", - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason=None, - index=1, - delta=Delta( - provider_specific_fields=None, - content='tomation"}', - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - citations=None, - ), - ModelResponseStream( - id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", - created=1751934860, - model="claude-3-7-sonnet-latest", - object="chat.completion.chunk", - system_fingerprint=None, - choices=[ - StreamingChoices( - finish_reason="tool_calls", - index=0, - delta=Delta( - provider_specific_fields=None, - content=None, - role=None, - function_call=None, - tool_calls=None, - audio=None, - ), - logprobs=None, - ) - ], - provider_specific_fields=None, - ), - ] - - response = stream_chunk_builder(chunks=chunks) - print(response) - - assert response is not None - assert response.choices[0].message.content is not None - assert response.choices[0].message.thinking_blocks is not None - - -from litellm.llms.openai.openai import OpenAIChatCompletion - - -def throw_retryable_error(*_, **__): - raise RuntimeError("BOOM") - - -@pytest.mark.asyncio -async def test_retrying() -> None: - litellm.num_retries = 10 - with ( - patch.object( - OpenAIChatCompletion, - "make_openai_chat_completion_request", - side_effect=throw_retryable_error, - ) as mock_request, - pytest.raises(litellm.InternalServerError, match="LiteLLM Retried: 10 times"), - ): - await litellm.acompletion( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "Hello"}], - ) - - -def test_anthropic_disable_url_suffix_env_var(): - """Test that LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX prevents /v1/messages suffix.""" - import os - from unittest.mock import MagicMock, patch - - from litellm import completion - - # Test with environment variable disabled (default behavior) - with patch.dict(os.environ, {"ANTHROPIC_API_BASE": "https://api.example.com"}): - actual_api_base = None - - with patch("litellm.main.anthropic_chat_completions") as mock_anthropic: - - def capture_completion(**kwargs): - nonlocal actual_api_base - actual_api_base = kwargs.get("api_base") - mock_response = MagicMock() - mock_response.choices = [MagicMock()] - return mock_response - - mock_anthropic.completion = capture_completion - - # This should append /v1/messages - completion( - model="anthropic/claude-3-sonnet", - messages=[{"role": "user", "content": "test"}], - api_key="test-key", - ) - - # Verify the api_base has /v1/messages appended - assert actual_api_base.endswith("/v1/messages") - assert actual_api_base == "https://api.example.com/v1/messages" - - # Test with environment variable enabled - with patch.dict( - os.environ, - { - "ANTHROPIC_API_BASE": "https://api.example.com/custom/path", - "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX": "true", - }, - ): - actual_api_base = None - - with patch("litellm.main.anthropic_chat_completions") as mock_anthropic: - - def capture_completion(**kwargs): - nonlocal actual_api_base - actual_api_base = kwargs.get("api_base") - mock_response = MagicMock() - mock_response.choices = [MagicMock()] - return mock_response - - mock_anthropic.completion = capture_completion - - # This should NOT append /v1/messages - completion( - model="anthropic/claude-3-sonnet", - messages=[{"role": "user", "content": "test"}], - api_key="test-key", - ) - - # Verify the api_base does not have /v1/messages appended - assert actual_api_base == "https://api.example.com/custom/path" - assert not actual_api_base.endswith("/v1/messages") - - -def test_anthropic_text_disable_url_suffix_env_var(): - """Test that LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX prevents /v1/complete suffix for anthropic_text.""" - import os - from unittest.mock import MagicMock, patch - - from litellm import completion - - # Test with environment variable disabled (default behavior) - with patch.dict(os.environ, {"ANTHROPIC_API_BASE": "https://api.example.com"}): - actual_api_base = None - - with patch("litellm.main.base_llm_http_handler") as mock_handler: - - def capture_completion(**kwargs): - nonlocal actual_api_base - actual_api_base = kwargs.get("api_base") - return MagicMock() - - mock_handler.completion = capture_completion - - # This should append /v1/complete - completion( - model="anthropic_text/claude-instant-1", - messages=[{"role": "user", "content": "test"}], - api_key="test-key", - ) - - # Verify the api_base has /v1/complete appended - assert actual_api_base.endswith("/v1/complete") - assert actual_api_base == "https://api.example.com/v1/complete" - - # Test with environment variable enabled - with patch.dict( - os.environ, - { - "ANTHROPIC_API_BASE": "https://api.example.com/custom/complete", - "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX": "true", - }, - ): - actual_api_base = None - - with patch("litellm.main.base_llm_http_handler") as mock_handler: - - def capture_completion(**kwargs): - nonlocal actual_api_base - actual_api_base = kwargs.get("api_base") - return MagicMock() - - mock_handler.completion = capture_completion - - # This should NOT append /v1/complete - completion( - model="anthropic_text/claude-instant-1", - messages=[{"role": "user", "content": "test"}], - api_key="test-key", - ) - - # Verify the api_base does not have /v1/complete appended - assert actual_api_base == "https://api.example.com/custom/complete" - assert not actual_api_base.endswith("/v1/complete") - - -def test_image_edit_merges_headers_and_extra_headers(): - from litellm.images.main import base_llm_http_handler - - combined_headers = { - "x-test-header-one": "value-1", - "x-test-header-two": "value-2", - } - - mock_image_edit_config = MagicMock() - mock_image_edit_config.get_supported_openai_params.return_value = set() - mock_image_edit_config.map_openai_params.side_effect = lambda **kwargs: dict( - kwargs["image_edit_optional_params"] - ) - - with ( - patch( - "litellm.images.main.ProviderConfigManager.get_provider_image_edit_config", - return_value=mock_image_edit_config, - ) as mock_config, - patch.object( - base_llm_http_handler, - "image_edit_handler", - return_value="ok", - ) as mock_handler, - ): - response = litellm.image_edit( - image=MagicMock(name="image"), - prompt="test", - model="azure/gpt-image-1", - headers={"x-test-header-one": "value-1"}, - extra_headers={ - "x-test-header-two": "value-2", - }, - ) - - assert response == "ok" - mock_config.assert_called_once() - - handler_kwargs = mock_handler.call_args.kwargs - assert handler_kwargs["extra_headers"] == combined_headers - assert "extra_headers" not in handler_kwargs["image_edit_optional_request_params"] - - -@pytest.mark.parametrize("metadata_key", ("metadata", "litellm_metadata")) -@pytest.mark.parametrize("input_tokens", (51234, 0)) -def test_mock_completion_usage_reports_admission_input_tokens(metadata_key: str, input_tokens: int): - response = litellm.completion( - model="anthropic/claude-sonnet-5", - messages=[{"role": "user", "content": "hello"}], - mock_response="ok", - api_key="mock", - **{metadata_key: {"user_api_key_budget_reservation": {"reserved_cost": 1.0, "input_tokens": input_tokens}}}, - ) - - assert response.usage.prompt_tokens == input_tokens - assert response.usage.total_tokens == input_tokens + response.usage.completion_tokens - - -def test_mock_completion_usage_falls_back_to_default_without_admission_count(): - response = litellm.completion( - model="anthropic/claude-sonnet-5", - messages=[{"role": "user", "content": "hello"}], - mock_response="ok", - api_key="mock", - metadata={"user_api_key_budget_reservation": {"reserved_cost": 1.0}}, - ) - - assert response.usage.prompt_tokens == litellm_main.DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT - - -_AZURE_AI_CUSTOM_PRICED_DEPLOYMENT: Final = { - "model_name": "azure-ai-custom-priced", - "litellm_params": { - "model": "azure_ai/gpt-5.6", - "api_key": "mock", - "api_base": "https://example.services.ai.azure.com", - "mock_response": "ok", - "input_cost_per_token": 3e-6, - "output_cost_per_token": 7e-6, - "cache_read_input_token_cost": 1e-7, - "cache_creation_input_token_cost": 5e-7, - }, - "model_info": {"id": "azure-ai-custom-priced-deployment-id"}, -} - - -def _expected_custom_price(response: litellm.ModelResponse) -> float: - params: Final = _AZURE_AI_CUSTOM_PRICED_DEPLOYMENT["litellm_params"] - return ( - response.usage.prompt_tokens * params["input_cost_per_token"] - + response.usage.completion_tokens * params["output_cost_per_token"] - ) - - -@pytest.mark.asyncio -@pytest.mark.parametrize("use_async", (False, True)) -async def test_mock_completion_prices_azure_ai_router_deployment_with_custom_pricing(use_async: bool): - router: Final = litellm.Router(model_list=[_AZURE_AI_CUSTOM_PRICED_DEPLOYMENT]) - messages: Final = [{"role": "user", "content": "hello"}] - - response: Final = ( - await router.acompletion(model="azure-ai-custom-priced", messages=messages) - if use_async - else router.completion(model="azure-ai-custom-priced", messages=messages) - ) - - assert response._hidden_params["response_cost"] == pytest.approx(_expected_custom_price(response)) - assert response._hidden_params["custom_llm_provider"] == "azure_ai" - - -@pytest.mark.parametrize( - ("model", "expected_provider"), - (("anthropic/claude-sonnet-5", "anthropic"), ("no-such-provider-model", None)), -) -def test_mock_completion_infers_provider_when_called_directly_without_one(model: str, expected_provider: str | None): - response: Final = litellm.mock_completion( - model=model, - messages=[{"role": "user", "content": "hello"}], - mock_response="ok", - ) - - assert response.choices[0].message.content == "ok" - assert response._hidden_params.get("custom_llm_provider") == expected_provider - - -_ADMISSION_INPUT_TOKENS: Final = 51234 - - -def _admission_metadata(input_tokens: int) -> dict[str, object]: # mutable-ok: logging writes into metadata - return {"user_api_key_budget_reservation": {"reserved_cost": 1.0, "input_tokens": input_tokens}} - - -_ADMISSION_METADATA: Final = _admission_metadata(_ADMISSION_INPUT_TOKENS) -_MOCK_STREAM_MESSAGES: Final = [{"role": "user", "content": "hello " * 200}] -_STREAM_CHUNK_BUILDER_TOKEN_COUNTER: Final = "litellm.litellm_core_utils.streaming_chunk_builder_utils.token_counter" - - -def _prompt_token_counter_calls(token_counter: MagicMock) -> list[object]: - return [call for call in token_counter.call_args_list if call.kwargs.get("messages") is not None] - - -def _client_usage_chunks(chunks: list[ModelResponseStream]) -> list[Usage]: - return [chunk.usage for chunk in chunks if getattr(chunk, "usage", None) is not None] - - -@pytest.mark.parametrize("n", (None, 2)) -def test_mock_completion_stream_usage_reports_admission_input_tokens_without_tokenizer_fallback(n: int | None): - with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - chunks: Final = list( - litellm.completion( - model="openai/gpt-5.4-mini", - messages=_MOCK_STREAM_MESSAGES, - mock_response="ok", - api_key="mock", - stream=True, - n=n, - stream_options={"include_usage": True}, - metadata=_ADMISSION_METADATA, - ) - ) - - usage_chunks: Final = _client_usage_chunks(chunks) - assert len(usage_chunks) == 1 - assert usage_chunks[0].prompt_tokens == _ADMISSION_INPUT_TOKENS - assert usage_chunks[0].completion_tokens == litellm_main.DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT - assert usage_chunks[0].total_tokens == _ADMISSION_INPUT_TOKENS + usage_chunks[0].completion_tokens - assert _prompt_token_counter_calls(token_counter) == [] - assert all(chunk.choices for chunk in chunks[:-1]) - assert {chunk.id for chunk in chunks} == {chunks[0].id} - - -@pytest.mark.asyncio -@pytest.mark.parametrize("n", (None, 2)) -async def test_mock_acompletion_stream_usage_reports_admission_input_tokens_without_tokenizer_fallback( - n: int | None, -): - with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - response: Final = await litellm.acompletion( - model="openai/gpt-5.4-mini", - messages=_MOCK_STREAM_MESSAGES, - mock_response="ok", - api_key="mock", - stream=True, - n=n, - stream_options={"include_usage": True}, - litellm_metadata=_ADMISSION_METADATA, - ) - chunks: Final = [chunk async for chunk in response] - - usage_chunks: Final = _client_usage_chunks(chunks) - assert len(usage_chunks) == 1 - assert usage_chunks[0].prompt_tokens == _ADMISSION_INPUT_TOKENS - assert usage_chunks[0].total_tokens == _ADMISSION_INPUT_TOKENS + usage_chunks[0].completion_tokens - assert _prompt_token_counter_calls(token_counter) == [] - assert all(chunk.choices for chunk in chunks[:-1]) - assert {chunk.id for chunk in chunks} == {chunks[0].id} - - -def test_mock_completion_stream_without_include_usage_hides_usage_chunk_but_logs_admission_count(): - with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - chunks: Final = list( - litellm.completion( - model="openai/gpt-5.4-mini", - messages=_MOCK_STREAM_MESSAGES, - mock_response="ok", - api_key="mock", - stream=True, - metadata=_ADMISSION_METADATA, - ) - ) - - assert _client_usage_chunks(chunks) == [] - assert all(len(chunk.choices) == 1 for chunk in chunks) - assert chunks[-1]._hidden_params["usage"].prompt_tokens == _ADMISSION_INPUT_TOKENS - assert _prompt_token_counter_calls(token_counter) == [] - - -def test_mock_completion_stream_with_empty_stream_options_completes_and_logs_admission_count(): - with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - chunks: Final = list( - litellm.completion( - model="openai/gpt-5.4-mini", - messages=_MOCK_STREAM_MESSAGES, - mock_response="ok", - api_key="mock", - stream=True, - stream_options={}, - metadata=_ADMISSION_METADATA, - ) - ) - - assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" - assert _client_usage_chunks(chunks) == [] - assert _prompt_token_counter_calls(token_counter) == [] - - -@pytest.mark.asyncio -async def test_mock_acompletion_stream_with_empty_stream_options_completes_and_logs_admission_count(): - with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - response: Final = await litellm.acompletion( - model="openai/gpt-5.4-mini", - messages=_MOCK_STREAM_MESSAGES, - mock_response="ok", - api_key="mock", - stream=True, - stream_options={}, - litellm_metadata=_ADMISSION_METADATA, - ) - chunks: Final = [chunk async for chunk in response] - - assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" - assert _client_usage_chunks(chunks) == [] - assert _prompt_token_counter_calls(token_counter) == [] - - -def test_mock_completion_stream_without_admission_count_falls_back_to_tokenizer(): - expected_prompt_tokens: Final = litellm.token_counter(model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES) - with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - chunks: Final = list( - litellm.completion( - model="openai/gpt-5.4-mini", - messages=_MOCK_STREAM_MESSAGES, - mock_response="ok", - api_key="mock", - stream=True, - stream_options={"include_usage": True}, - metadata={"user_api_key_budget_reservation": {"reserved_cost": 1.0}}, - ) - ) - - usage_chunks: Final = _client_usage_chunks(chunks) - assert len(usage_chunks) == 1 - assert usage_chunks[0].prompt_tokens == expected_prompt_tokens - assert usage_chunks[0].total_tokens == expected_prompt_tokens + usage_chunks[0].completion_tokens - assert len(_prompt_token_counter_calls(token_counter)) >= 1 - - -@pytest.mark.asyncio -async def test_mock_acompletion_stream_without_admission_count_falls_back_to_tokenizer(): - expected_prompt_tokens: Final = litellm.token_counter(model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES) - with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - response: Final = await litellm.acompletion( - model="openai/gpt-5.4-mini", - messages=_MOCK_STREAM_MESSAGES, - mock_response="ok", - api_key="mock", - stream=True, - stream_options={"include_usage": True}, - ) - chunks: Final = [chunk async for chunk in response] - - usage_chunks: Final = _client_usage_chunks(chunks) - assert len(usage_chunks) == 1 - assert usage_chunks[0].prompt_tokens == expected_prompt_tokens - assert len(_prompt_token_counter_calls(token_counter)) >= 1 - - -def _usage_triple(usage: Usage) -> tuple[int, int, int]: - return (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) - - -@pytest.mark.parametrize("input_tokens", (_ADMISSION_INPUT_TOKENS, 0)) -def test_mock_completion_stream_and_non_stream_report_the_same_admission_usage(input_tokens: int): - metadata: Final = _admission_metadata(input_tokens) - non_stream: Final = litellm.completion( - model="openai/gpt-5.4-mini", - messages=_MOCK_STREAM_MESSAGES, - mock_response="ok", - api_key="mock", - metadata=metadata, - ) - with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - chunks: Final = list( - litellm.completion( - model="openai/gpt-5.4-mini", - messages=_MOCK_STREAM_MESSAGES, - mock_response="ok", - api_key="mock", - stream=True, - stream_options={"include_usage": True}, - metadata=metadata, - ) - ) - - assert _usage_triple(non_stream.usage) == _usage_triple(_client_usage_chunks(chunks)[0]) - assert non_stream.usage.prompt_tokens == input_tokens - assert _prompt_token_counter_calls(token_counter) == [] - - -@pytest.mark.asyncio -async def test_mock_acompletion_stream_reports_zero_admission_input_tokens_without_tokenizer_fallback(): - with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - response: Final = await litellm.acompletion( - model="openai/gpt-5.4-mini", - messages=[{"role": "user", "content": ""}], - mock_response="ok", - api_key="mock", - stream=True, - stream_options={"include_usage": True}, - litellm_metadata=_admission_metadata(0), - ) - chunks: Final = [chunk async for chunk in response] - - usage_chunks: Final = _client_usage_chunks(chunks) - assert len(usage_chunks) == 1 - assert _usage_triple(usage_chunks[0]) == (0, usage_chunks[0].completion_tokens, usage_chunks[0].completion_tokens) - assert _prompt_token_counter_calls(token_counter) == [] - - -def test_mock_text_completion_stream_and_non_stream_report_the_same_zero_admission_usage(): - metadata: Final = _admission_metadata(0) - non_stream: Final = litellm.text_completion( - model="openai/gpt-5.4-mini", prompt="", mock_response="ok", api_key="mock", metadata=metadata - ) - chunks: Final = list( - litellm.text_completion( - model="openai/gpt-5.4-mini", - prompt="", - mock_response="ok", - api_key="mock", - stream=True, - stream_options={"include_usage": True}, - metadata=metadata, - ) - ) - - stream_usages: Final = tuple(chunk.usage for chunk in chunks if getattr(chunk, "usage", None) is not None) - assert len(stream_usages) == 1 - assert _usage_triple(non_stream.usage) == _usage_triple(stream_usages[0]) - assert non_stream.usage.prompt_tokens == 0 - - -def test_mock_completion_stream_with_model_response(): - """Test that mock_completion correctly handles stream=True with a ModelResponse as mock_response.""" - from litellm import completion - from litellm.types.utils import Choices, Message, ModelResponse, Usage - - # Create a ModelResponse object - mock_model_response = ModelResponse( - id="chatcmpl-test-123", - created=1234567890, - model="gpt-4o-mini", - object="chat.completion", - choices=[ - Choices( - finish_reason="stop", - index=0, - message=Message( - content="This is a test response", - role="assistant", - ), - ) - ], - usage=Usage( - prompt_tokens=10, - completion_tokens=20, - total_tokens=30, - ), - ) - - # Call completion with stream=True and mock_response as ModelResponse - response = completion( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "Hello"}], - stream=True, - mock_response=mock_model_response, - ) - - # Verify that the response is a stream - assert response is not None - - # Collect all chunks from the stream - chunks = [] - for chunk in response: - chunks.append(chunk) - print(f"Chunk: {chunk}") - - # Verify we got chunks - assert len(chunks) > 0 - - # Verify the content is streamed correctly - accumulated_content = "" - for chunk in chunks: - if ( - hasattr(chunk.choices[0].delta, "content") - and chunk.choices[0].delta.content - ): - accumulated_content += chunk.choices[0].delta.content - - assert "This is a test response" in accumulated_content or len(chunks) > 0 - - -@pytest.mark.asyncio -async def test_async_mock_completion_stream_with_model_response(): - """Test that async mock_completion correctly handles stream=True with a ModelResponse as mock_response.""" - from litellm import acompletion - from litellm.types.utils import Choices, Message, ModelResponse, Usage - - # Create a ModelResponse object - mock_model_response = ModelResponse( - id="chatcmpl-test-456", - created=1234567890, - model="gpt-4o-mini", - object="chat.completion", - choices=[ - Choices( - finish_reason="stop", - index=0, - message=Message( - content="This is an async test response", - role="assistant", - ), - ) - ], - usage=Usage( - prompt_tokens=15, - completion_tokens=25, - total_tokens=40, - ), - ) - - # Call acompletion with stream=True and mock_response as ModelResponse - response = await acompletion( - model="gpt-4o-mini", - messages=[{"role": "user", "content": "Hello async"}], - stream=True, - mock_response=mock_model_response, - ) - - # Verify that the response is a stream - assert response is not None - - # Collect all chunks from the stream - chunks = [] - async for chunk in response: - chunks.append(chunk) - print(f"Async Chunk: {chunk}") - - # Verify we got chunks - assert len(chunks) > 0 - - # Verify the content is streamed correctly - accumulated_content = "" - for chunk in chunks: - if ( - hasattr(chunk.choices[0].delta, "content") - and chunk.choices[0].delta.content - ): - accumulated_content += chunk.choices[0].delta.content - - assert "This is an async test response" in accumulated_content or len(chunks) > 0 - - -class TestCallTypesOCR: - """Test that OCR call types are properly defined in CallTypes enum. - - Fixes https://github.com/BerriAI/litellm/issues/17381 - """ - - def test_ocr_call_type_exists(self): - """Test that CallTypes.ocr exists and has correct value.""" - from litellm.types.utils import CallTypes - - assert hasattr(CallTypes, "ocr") - assert CallTypes.ocr.value == "ocr" - - def test_aocr_call_type_exists(self): - """Test that CallTypes.aocr exists and has correct value.""" - from litellm.types.utils import CallTypes - - assert hasattr(CallTypes, "aocr") - assert CallTypes.aocr.value == "aocr" - - def test_ocr_call_type_from_string(self): - """Test that CallTypes can be constructed from 'ocr' string.""" - from litellm.types.utils import CallTypes - - call_type = CallTypes("ocr") - assert call_type == CallTypes.ocr - - def test_aocr_call_type_from_string(self): - """Test that CallTypes can be constructed from 'aocr' string. - - This is the actual use case that was failing - the OCR endpoint - uses route_type='aocr' and guardrails try to instantiate - CallTypes('aocr'). - """ - from litellm.types.utils import CallTypes - - call_type = CallTypes("aocr") - assert call_type == CallTypes.aocr - - -def test_stream_chunk_builder_text_completion_combines_text_and_usage(): - from litellm.main import stream_chunk_builder_text_completion - from litellm.types.utils import TextCompletionResponse - - chunks = [ - TextCompletionResponse( - id="cmpl-1", - object="text_completion", - created=1, - model="gpt-3.5-turbo-instruct", - choices=[{"text": "Hello", "index": 0, "logprobs": None, "finish_reason": None}], - ), - TextCompletionResponse( - id="cmpl-1", - object="text_completion", - created=1, - model="gpt-3.5-turbo-instruct", - choices=[{"text": " world", "index": 0, "logprobs": None, "finish_reason": "stop"}], - ), - ] - - response = stream_chunk_builder_text_completion( - chunks=chunks, messages=[{"role": "user", "content": "say hello"}] - ) - - assert response.choices[0].text == "Hello world" - assert response.choices[0].finish_reason == "stop" - assert response.usage.prompt_tokens > 0 - assert response.usage.completion_tokens > 0 - assert response.usage.total_tokens == response.usage.prompt_tokens + response.usage.completion_tokens - - -def test_completion_forwards_store_and_prompt_cache_key_to_openai(): - """ - Regression test for https://github.com/BerriAI/litellm/issues/33184 - - store and prompt_cache_key are documented OpenAI chat completion params that - were accepted as supported but silently dropped before the provider request - was built, because they were not named parameters of completion() and - get_optional_params() the way safety_identifier is. - """ - from openai import OpenAI - - client = OpenAI(api_key="fake-api-key") - - with patch.object(client.chat.completions.with_raw_response, "create") as mock_client: - try: - litellm.completion( - model="openai/gpt-4o", - messages=[{"role": "user", "content": "Hello"}], - store=False, - prompt_cache_key="test-cache-key", - client=client, - ) - except Exception as e: - print(e) - - mock_client.assert_called_once() - request_body = mock_client.call_args.kwargs - assert request_body["store"] is False - assert request_body["prompt_cache_key"] == "test-cache-key" - - -@pytest.mark.asyncio -async def test_acompletion_forwards_store_and_prompt_cache_key_to_openai(): - """ - Async variant of the store/prompt_cache_key forwarding regression test for - https://github.com/BerriAI/litellm/issues/33184 - """ - from openai import AsyncOpenAI - - client = AsyncOpenAI(api_key="fake-api-key") - - with patch.object(client.chat.completions.with_raw_response, "create") as mock_client: - try: - await litellm.acompletion( - model="openai/gpt-4o", - messages=[{"role": "user", "content": "Hello"}], - store=False, - prompt_cache_key="test-cache-key", - client=client, - ) - except Exception as e: - print(e) - - mock_client.assert_called_once() - request_body = mock_client.call_args.kwargs - assert request_body["store"] is False - assert request_body["prompt_cache_key"] == "test-cache-key" - - -def test_completion_omits_store_and_prompt_cache_key_when_not_passed(): - """ - When store and prompt_cache_key are not passed, they must not appear in the - outbound request body (guards against always forwarding None defaults). - """ - from openai import OpenAI - - client = OpenAI(api_key="fake-api-key") - - with patch.object(client.chat.completions.with_raw_response, "create") as mock_client: - try: - litellm.completion( - model="openai/gpt-4o", - messages=[{"role": "user", "content": "Hello"}], - client=client, - ) - except Exception as e: - print(e) - - mock_client.assert_called_once() - request_body = mock_client.call_args.kwargs - assert "store" not in request_body - assert "prompt_cache_key" not in request_body - - -def test_completion_forwards_store_and_prompt_cache_key_to_mcp_gateway(): - """ - Regression test for the MCP gateway early-return in completion(): store and - prompt_cache_key are named params, so they no longer travel via **kwargs and - must be forwarded explicitly like safety_identifier and service_tier. - """ - with patch.object( - import_module("litellm.responses.mcp.chat_completions_handler"), "acompletion_with_mcp" - ) as mock_mcp: - result = litellm.completion( - model="openai/gpt-4o", - messages=[{"role": "user", "content": "Hello"}], - tools=[{"type": "mcp", "server_url": "litellm_proxy"}], - store=False, - prompt_cache_key="test-cache-key", - ) - - result.close() - mock_mcp.assert_called_once() - call_kwargs = mock_mcp.call_args.kwargs - assert call_kwargs["store"] is False - assert call_kwargs["prompt_cache_key"] == "test-cache-key" - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "aws_credential_kwargs", - [ - { - "aws_session_name": "litellm-gcp", - "aws_role_name": "arn:aws:iam::123456789012:role/litellm-bedrock-role", - "aws_web_identity_token": "oidc/google/108963886734710037768", - }, - { - "aws_access_key_id": "AKIASTATICKEYFORTEST", - "aws_secret_access_key": "static-secret-key", - "aws_session_token": "static-session-token", - }, - ], - ids=["web_identity", "static_keys"], -) -async def test_acompletion_forwards_aws_credentials_through_responses_bridge( - respx_mock: respx.MockRouter, monkeypatch, aws_credential_kwargs: dict -): - from botocore.credentials import Credentials - - from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM - - original_disable_aiohttp = litellm.disable_aiohttp_transport - try: - litellm.disable_aiohttp_transport = True - monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") - litellm.in_memory_llm_clients_cache.flush_cache() - monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) - monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) - - get_credentials_mock = MagicMock(return_value=Credentials("fake-key", "fake-secret")) - monkeypatch.setattr(BaseAWSLLM, "get_credentials", get_credentials_mock) - - respx_mock.post("https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses").respond( - json={ - "id": "resp_123", - "object": "response", - "created_at": 1760144904, - "status": "completed", - "model": "openai.gpt-5.4", - "output": [ - { - "type": "message", - "id": "msg_1", - "role": "assistant", - "status": "completed", - "content": [{"type": "output_text", "text": "ok", "annotations": []}], - } - ], - } - ) - - response = await litellm.acompletion( - model="bedrock_mantle/openai.gpt-5.4", - messages=[{"role": "user", "content": "hi"}], - api_base="https://bedrock-mantle.us-east-2.api.aws/v1", - aws_region_name="us-east-2", - num_retries=0, - **aws_credential_kwargs, - ) - - assert response.choices[0].message.content == "ok" - credential_kwargs = get_credentials_mock.call_args.kwargs - assert credential_kwargs["aws_region_name"] == "us-east-2" - for key, value in aws_credential_kwargs.items(): - assert credential_kwargs[key] == value - authorization = respx_mock.calls.last.request.headers["Authorization"] - assert authorization.startswith("AWS4-HMAC-SHA256") - assert "fake-key" in authorization - finally: - litellm.disable_aiohttp_transport = original_disable_aiohttp - litellm.in_memory_llm_clients_cache.flush_cache() - - -_GEMINI_RESPONSE_BODY = { - "candidates": [{"content": {"parts": [{"text": "hello"}], "role": "model"}, "finishReason": "STOP"}], - "usageMetadata": {"promptTokenCount": 2, "candidatesTokenCount": 1, "totalTokenCount": 3}, -} - - -def _gemini_client_returning_a_reply(): - """An injected HTTP client whose post() answers like generativelanguage does.""" - from litellm.llms.custom_httpx.http_handler import HTTPHandler - - client = HTTPHandler() - request = httpx.Request("POST", "https://generativelanguage.googleapis.com/") - post = MagicMock(return_value=httpx.Response(200, json=_GEMINI_RESPONSE_BODY, request=request)) - return client, post - - -@pytest.fixture -def restore_model_registry(): - """litellm.model_cost and the provider name sets are module-global. - - register_model merges into the existing entry in place, hence the deep copy. - """ - model_cost = copy.deepcopy(litellm.model_cost) - openai_models = set(litellm.open_ai_chat_completion_models) - yield - litellm.model_cost.clear() - litellm.model_cost.update(model_cost) - litellm.open_ai_chat_completion_models.clear() - litellm.open_ai_chat_completion_models.update(openai_models) - - -def test_openai_model_name_does_not_outrank_explicit_provider(): - """`gemini/gpt-4o` goes to Google, not to litellm's OpenAI handler. - - completion() checks `model in litellm.open_ai_chat_completion_models` ahead of - the gemini branch, so the call used to reach the OpenAI handler carrying - VertexGeminiConfig, whose transform_request raises NotImplementedError. - """ - assert "gpt-4o" in litellm.open_ai_chat_completion_models - client, post = _gemini_client_returning_a_reply() - - with patch.object(client, "post", new=post): - response = litellm.completion( - model="gemini/gpt-4o", - messages=[{"role": "user", "content": "hello"}], - api_key="test-api-key", - client=client, - ) - - assert "generativelanguage.googleapis.com" in post.call_args.kwargs["url"] - assert "models/gpt-4o" in post.call_args.kwargs["url"] - assert response.choices[0].message.content == "hello" - - -def test_mislabelled_pricing_entry_does_not_reroute_provider(restore_model_registry): - """register_model is the other way into the same failure. - - An entry claiming litellm_provider "openai" adds its name to - open_ai_chat_completion_models, so one mislabelled price reroutes every later - call to that model in the process. - """ - litellm.register_model( - { - "gemini-2.5-pro": { - "litellm_provider": "openai", - "mode": "chat", - "input_cost_per_token": 1e-06, - "output_cost_per_token": 4e-06, - } - } - ) - assert "gemini-2.5-pro" in litellm.open_ai_chat_completion_models - client, post = _gemini_client_returning_a_reply() - - with patch.object(client, "post", new=post): - response = litellm.completion( - model="gemini/gemini-2.5-pro", - messages=[{"role": "user", "content": "hello"}], - api_key="test-api-key", - client=client, - ) - - assert "generativelanguage.googleapis.com" in post.call_args.kwargs["url"] - assert response.choices[0].message.content == "hello" - - -def test_openai_model_without_a_provider_still_routes_to_openai(): - from openai import OpenAI - - client = OpenAI(api_key="fake-key") - raw_response = client.chat.completions.with_raw_response - with patch.object(raw_response, "create") as mock_create, contextlib.suppress(Exception): - litellm.completion( - model="gpt-4o", - messages=[{"role": "user", "content": "hello"}], - client=client, - ) - - mock_create.assert_called() - - -def _openai_chat_create_kwargs(client, **completion_kwargs): - with patch.object(client.chat.completions.with_raw_response, "create") as mock_client: - with contextlib.suppress(Exception): - litellm.completion( - messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], - cache_control_injection_points=[{"location": "message", "role": "system"}], - client=client, - **completion_kwargs, - ) - - mock_client.assert_called_once() - return mock_client.call_args.kwargs - - -@pytest.fixture -def _no_openai_api_base_override(monkeypatch): - monkeypatch.delenv("OPENAI_BASE_URL", raising=False) - monkeypatch.delenv("OPENAI_API_BASE", raising=False) - monkeypatch.setattr(litellm, "api_base", None) - - -@pytest.mark.usefixtures("_no_openai_api_base_override") -def test_completion_custom_api_base_sends_no_prompt_cache_breakpoint_for_gpt_5_6(): - from openai import OpenAI - - client = OpenAI(api_key="fake-api-key", base_url="http://127.0.0.1:9/v1") - request_body = _openai_chat_create_kwargs(client, model="gpt-5.6", api_base="http://127.0.0.1:9/v1") - - assert request_body["messages"][0] == {"role": "system", "content": "sys", "cache_control": {"type": "ephemeral"}} - assert "prompt_cache_breakpoint" not in json.dumps(request_body["messages"]) - assert "prompt_cache_options" not in json.dumps(request_body) - - -@pytest.mark.usefixtures("_no_openai_api_base_override") -def test_completion_custom_base_url_sends_no_prompt_cache_breakpoint_for_gpt_5_6(): - from openai import OpenAI - - client = OpenAI(api_key="fake-api-key", base_url="http://127.0.0.1:9/v1") - request_body = _openai_chat_create_kwargs(client, model="gpt-5.6", base_url="http://127.0.0.1:9/v1") - - assert request_body["messages"][0] == {"role": "system", "content": "sys", "cache_control": {"type": "ephemeral"}} - assert "prompt_cache_breakpoint" not in json.dumps(request_body["messages"]) - assert "prompt_cache_options" not in json.dumps(request_body) - - -@pytest.mark.asyncio -@pytest.mark.usefixtures("_no_openai_api_base_override") -async def test_acompletion_custom_base_url_sends_no_prompt_cache_breakpoint_for_gpt_5_6(): - from openai import AsyncOpenAI - - client = AsyncOpenAI(api_key="fake-api-key", base_url="http://127.0.0.1:9/v1") - with patch.object(client.chat.completions.with_raw_response, "create") as mock_create: - with contextlib.suppress(Exception): - await litellm.acompletion( - model="gpt-5.6", - messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], - cache_control_injection_points=[{"location": "message", "role": "system"}], - client=client, - base_url="http://127.0.0.1:9/v1", - ) - - mock_create.assert_called_once() - request_body = mock_create.call_args.kwargs - - assert request_body["messages"][0] == {"role": "system", "content": "sys", "cache_control": {"type": "ephemeral"}} - assert "prompt_cache_breakpoint" not in json.dumps(request_body["messages"]) - assert "prompt_cache_options" not in json.dumps(request_body) - - -@pytest.mark.usefixtures("_no_openai_api_base_override") -def test_completion_default_api_base_sends_prompt_cache_breakpoint_for_gpt_5_6(): - from openai import OpenAI - - client = OpenAI(api_key="fake-api-key") - request_body = _openai_chat_create_kwargs(client, model="gpt-5.6") - - assert request_body["messages"][0]["content"] == [ - {"type": "text", "text": "sys", "prompt_cache_breakpoint": {"mode": "explicit"}} - ] - assert request_body["extra_body"]["prompt_cache_options"] == {"mode": "explicit"} - - -_SUBSCRIPTION_OAUTH_CREDENTIAL = "Bearer sk-ant-oat01-fake-subscription-token-for-testing-0123456789" - - -def _scoped_headers_for_oauth_request(): - from litellm.types.utils import ProviderSpecificHeader - - return [ - ProviderSpecificHeader( - custom_llm_provider="anthropic,bedrock,vertex_ai", - extra_headers={"anthropic-version": "2023-06-01"}, - ), - ProviderSpecificHeader( - custom_llm_provider="anthropic", - extra_headers={"authorization": _SUBSCRIPTION_OAUTH_CREDENTIAL}, - ), - ] - - -def _run_anthropic_hop_with_shared_headers(shared_headers): - litellm.completion( - model="anthropic/claude-3-5-sonnet-20240620", - messages=[{"role": "user", "content": "Say OK"}], - extra_headers=shared_headers, - provider_specific_header=_scoped_headers_for_oauth_request(), - api_key="sk-fake-anthropic-key", - mock_response="OK", - ) - - -def test_completion_does_not_mutate_caller_supplied_headers(): - shared_headers = {"x-tenant": "acme"} - - _run_anthropic_hop_with_shared_headers(shared_headers) - - assert shared_headers == {"x-tenant": "acme"} - - -def test_anthropic_oauth_credential_does_not_persist_into_next_provider_hop(): - shared_headers = {"x-tenant": "acme"} - - _run_anthropic_hop_with_shared_headers(shared_headers) - - leaked = [name for name, value in shared_headers.items() if value == _SUBSCRIPTION_OAUTH_CREDENTIAL] - assert leaked == [] - assert "anthropic-version" not in shared_headers - - -STREAM_COST_MODEL = "gpt-4o" -STREAMED_USAGE = {"prompt_tokens": 137, "completion_tokens": 42, "total_tokens": 179} - - -def _text_chunk(content, finish_reason=None, usage=None): - chunk = { - "id": "chatcmpl-stream-cost", - "object": "chat.completion.chunk", - "created": 1700000000, - "model": STREAM_COST_MODEL, - "choices": [ - { - "index": 0, - "delta": {"role": "assistant", "content": content}, - "finish_reason": finish_reason, - } - ], - } - if usage is not None: - chunk["usage"] = usage - return chunk - - -def _priced_at(prompt_tokens, completion_tokens): - prices = litellm.model_cost[STREAM_COST_MODEL] - return ( - prompt_tokens * prices["input_cost_per_token"] - + completion_tokens * prices["output_cost_per_token"] - ) - - -@pytest.fixture -def local_cost_map(monkeypatch): - """The prices these tests assert are the checked-in ones. Setting the environment - variable alone does not reload the map, so pin the map itself. - - Prices are read through two separate lru_caches, so pinning ``model_cost`` is not - enough on its own: an entry warmed against the network-fetched map keeps its old - prices and billing reads those while the assertions read the pinned map. - ``_invalidate_model_cost_lowercase_map`` clears both caches, where - ``get_model_info.cache_clear`` reaches only one. Invalidate on the way in and out - so entries never leak across tests in either direction.""" - from litellm.utils import _invalidate_model_cost_lowercase_map - - monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") - monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) - _invalidate_model_cost_lowercase_map() - yield - _invalidate_model_cost_lowercase_map() - - -def test_a_streamed_response_bills_the_usage_the_provider_reported(local_cost_map): - rebuilt = litellm.stream_chunk_builder( - chunks=[ - _text_chunk("Hello"), - _text_chunk(" there"), - _text_chunk(None, finish_reason="stop", usage=STREAMED_USAGE), - ], - messages=[{"role": "user", "content": "hi"}], - ) - - assert rebuilt.choices[0].message.content == "Hello there" - assert rebuilt.usage.prompt_tokens == STREAMED_USAGE["prompt_tokens"] - assert rebuilt.usage.completion_tokens == STREAMED_USAGE["completion_tokens"] - - cost = litellm.completion_cost(completion_response=rebuilt, model=STREAM_COST_MODEL) - - assert cost == pytest.approx(_priced_at(137, 42)) - - -def test_streaming_and_not_streaming_bill_the_same_usage_the_same(local_cost_map): - rebuilt = litellm.stream_chunk_builder( - chunks=[ - _text_chunk("Hello"), - _text_chunk(" there"), - _text_chunk(None, finish_reason="stop", usage=STREAMED_USAGE), - ], - messages=[{"role": "user", "content": "hi"}], - ) - whole = litellm.ModelResponse( - id="chatcmpl-stream-cost", - model=STREAM_COST_MODEL, - object="chat.completion", - created=1700000000, - choices=[ - { - "index": 0, - "message": {"role": "assistant", "content": "Hello there"}, - "finish_reason": "stop", - } - ], - usage=STREAMED_USAGE, - ) - - assert litellm.completion_cost( - completion_response=rebuilt, model=STREAM_COST_MODEL - ) == pytest.approx(litellm.completion_cost(completion_response=whole, model=STREAM_COST_MODEL)) - - -def test_a_stream_that_reported_no_usage_is_still_billed(local_cost_map): - rebuilt = litellm.stream_chunk_builder( - chunks=[ - _text_chunk("Hello"), - _text_chunk(" there"), - _text_chunk(None, finish_reason="stop"), - ], - messages=[{"role": "user", "content": "hi"}], - ) - - assert rebuilt.usage.prompt_tokens > 0 - assert rebuilt.usage.completion_tokens > 0 - - cost = litellm.completion_cost(completion_response=rebuilt, model=STREAM_COST_MODEL) - - assert cost > 0 - assert cost == pytest.approx( - _priced_at(rebuilt.usage.prompt_tokens, rebuilt.usage.completion_tokens) - ) - - -@pytest.mark.asyncio -async def test_acompletion_resolves_provider_from_api_base(): - response = await litellm.acompletion( - model="deepseek-chat", - api_base="https://api.deepseek.com/v1", - api_key="fake-key", - messages=[{"role": "user", "content": "hi"}], - mock_response="resolved", - ) - - assert response.choices[0].message.content == "resolved" - - -@dataclass(frozen=True, slots=True) -class _RecordedSpeechSuccess: - call_type: str | None - spend_metadata: Mapping[str, object] - response_cost: float | None - logged_response_cost: float | None - - -def _record_speech_success(payload: dict[str, object]) -> _RecordedSpeechSuccess: - call_type: Final = payload.get("call_type") - response_cost: Final = payload.get("response_cost") - logging_payload: Final = payload.get("standard_logging_object") - logged_cost: Final = logging_payload.get("response_cost") if isinstance(logging_payload, dict) else None - return _RecordedSpeechSuccess( - call_type=call_type if isinstance(call_type, str) else None, - spend_metadata=get_litellm_metadata_from_kwargs(payload), - response_cost=response_cost if isinstance(response_cost, float) else None, - logged_response_cost=logged_cost if isinstance(logged_cost, float) else None, - ) - - -class _SuccessEventRecorder(CustomLogger): - def __init__(self) -> None: - super().__init__() - self.events: list[_RecordedSpeechSuccess] = [] # mutable-ok: test recorder of success-callback events - - async def async_log_success_event( - self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object - ) -> None: - self.events.append(_record_speech_success(kwargs)) - - -async def _wait_for_success_event(recorder: _SuccessEventRecorder, call_type: str) -> _RecordedSpeechSuccess: - for _ in range(100): - if (event := next((e for e in recorder.events if e.call_type == call_type), None)) is not None: - return event - await asyncio.sleep(0.05) - pytest.fail(f"no {call_type} success event; got {[e.call_type for e in recorder.events]}") - - -def _gemini_tts_generate_content_response() -> dict[str, object]: - return { - "candidates": [ - { - "content": { - "parts": [ - { - "inlineData": { - "mimeType": "audio/L16;codec=pcm;rate=24000", - "data": base64.b64encode(b"pcm-audio-bytes").decode(), - } - } - ], - "role": "model", - }, - "finishReason": "STOP", - "index": 0, - } - ], - "usageMetadata": { - "promptTokenCount": 5, - "candidatesTokenCount": 60, - "totalTokenCount": 65, - "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 5}], - "candidatesTokensDetails": [{"modality": "AUDIO", "tokenCount": 60}], - }, - "modelVersion": "gemini-2.5-flash-preview-tts", - } - - -@pytest.mark.asyncio -async def test_aspeech_gemini_bridge_keeps_proxy_metadata_for_spend_tracking( - respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch -) -> None: - monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) - monkeypatch.delenv("GEMINI_API_KEY", raising=False) - monkeypatch.delenv("GOOGLE_API_KEY", raising=False) - recorder: Final = _SuccessEventRecorder() - monkeypatch.setattr(litellm, "callbacks", [recorder]) - mock_route: Final = respx_mock.post( - url__regex=r"https://generativelanguage\.googleapis\.com/v1beta/models/gemini-2\.5-flash-preview-tts:generateContent.*" - ).mock(return_value=httpx.Response(200, json=_gemini_tts_generate_content_response())) - - await litellm.aspeech( - model="gemini/gemini-2.5-flash-preview-tts", - input="spend tracking check", - voice="Kore", - api_key="fake-gemini-key", - metadata={"user_api_key": "hashed-virtual-key", "user_api_key_user_id": "user-1"}, - ) - - assert mock_route.called - assert mock_route.calls.last.request.headers["x-goog-api-key"] == "fake-gemini-key" - speech_event: Final = await _wait_for_success_event(recorder, call_type="aspeech") - assert speech_event.spend_metadata["user_api_key"] == "hashed-virtual-key" - assert speech_event.spend_metadata["user_api_key_user_id"] == "user-1" - expected_prompt_cost, expected_completion_cost = litellm.cost_per_token( - model="gemini/gemini-2.5-flash-preview-tts", - usage_object=Usage(prompt_tokens=5, completion_tokens=60, total_tokens=65), - ) - expected_cost: Final = expected_prompt_cost + expected_completion_cost - assert expected_cost > 0 - assert speech_event.response_cost == pytest.approx(expected_cost) - assert speech_event.logged_response_cost == pytest.approx(expected_cost) - - -def _stream_builder_text_chunk(model: str, content: str, finish_reason: str | None = None) -> ModelResponseStream: - return ModelResponseStream( - id="chatcmpl-cost", - created=1724900000, - model=model, - object="chat.completion.chunk", - choices=[StreamingChoices(finish_reason=finish_reason, index=0, delta=Delta(content=content, role="assistant"))], - ) - - -def test_stream_chunk_builder_sets_hidden_response_cost_for_known_model(): - chunks: Final = [ - _stream_builder_text_chunk("gpt-4o", "Hello "), - _stream_builder_text_chunk("gpt-4o", "world.", finish_reason="stop"), - ] - - response: Final = litellm.stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "hi"}]) - - assert response is not None - prompt_cost, completion_cost = litellm.cost_per_token(model="gpt-4o", usage_object=response.usage) - expected_cost: Final = prompt_cost + completion_cost - assert expected_cost > 0 - assert response._hidden_params["response_cost"] == pytest.approx(expected_cost) - - -def test_stream_chunk_builder_unknown_model_leaves_response_cost_unset(): - chunks: Final = [ - _stream_builder_text_chunk("totally-unknown-model-xyz", "Hello "), - _stream_builder_text_chunk("totally-unknown-model-xyz", "world.", finish_reason="stop"), - ] - - response: Final = litellm.stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "hi"}]) - - assert response is not None - assert response._hidden_params.get("response_cost") is None - assert response.choices[0].message.content == "Hello world." - - -def test_stream_chunk_builder_prices_proxy_alias_via_model_map(): - chunks: Final = [ - _stream_builder_text_chunk("claude-opus-5", "Hello "), - _stream_builder_text_chunk("claude-opus-5", "world.", finish_reason="stop"), - ] - for chunk in chunks: - chunk._hidden_params = {"custom_llm_provider": "openai"} - - response: Final = litellm.stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "hi"}]) - - assert response is not None - assert response._hidden_params["custom_llm_provider"] == "openai" - prompt_cost, completion_cost = litellm.cost_per_token(model="claude-opus-5", usage_object=response.usage) - expected_cost: Final = prompt_cost + completion_cost - assert expected_cost > 0 - assert response._hidden_params["response_cost"] == pytest.approx(expected_cost) - - -def _stream_builder_logging_obj(model: str = "gpt-4o", custom_llm_provider: str = "openai") -> LiteLLMLogging: - logging_obj: Final = LiteLLMLogging( - model=model, - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=datetime.now(), - litellm_call_id="test-call-id", - function_id="test-function-id", - ) - logging_obj.update_environment_variables( - model=model, - user=None, - optional_params={}, - litellm_params={"custom_llm_provider": custom_llm_provider}, - custom_llm_provider=custom_llm_provider, - ) - return logging_obj - - -def test_stream_chunk_builder_stamps_streaming_usage_cost_by_default(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", False) - chunks: Final = [ - _stream_builder_text_chunk("gpt-4o", "Hello "), - _stream_builder_text_chunk("gpt-4o", "world.", finish_reason="stop"), - ] - - response: Final = litellm.stream_chunk_builder( - chunks=chunks, messages=[{"role": "user", "content": "hi"}], logging_obj=_stream_builder_logging_obj() - ) - - assert response is not None - usage_cost: Final = getattr(response.usage, "cost", None) - assert usage_cost is not None - assert usage_cost > 0 - assert response._hidden_params["response_cost"] == pytest.approx(usage_cost) - - -def test_stream_chunk_builder_skips_stamp_when_cost_is_unpriceable(): - import time as time_module - - from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging - - logging_obj: Final = LiteLLMLogging( - model="us.anthropic.claude-opus-5", - messages=[{"role": "user", "content": "hi"}], - stream=True, - call_type="completion", - start_time=time_module.time(), - litellm_call_id="stream-builder-alias-unpriceable", - function_id="1", - ) - logging_obj.model_call_details["custom_llm_provider"] = "bedrock" - logging_obj.optional_params = {} - usage_chunk: Final = _stream_builder_text_chunk("bedrock-claude-opus-5", "") - usage_chunk.usage = Usage(prompt_tokens=40, completion_tokens=5, total_tokens=45) - chunks: Final = [ - _stream_builder_text_chunk("bedrock-claude-opus-5", "Hello ", finish_reason="stop"), - usage_chunk, - ] - - response: Final = litellm.stream_chunk_builder( - chunks=chunks, messages=[{"role": "user", "content": "hi"}], logging_obj=logging_obj - ) - - assert response is not None - assert getattr(response.usage, "cost", None) is None - assert response._hidden_params.get("response_cost") is None - - -def test_stream_chunk_builder_keeps_provider_reported_usage_cost(): - usage_chunk: Final = _stream_builder_text_chunk("gpt-4o", "") - usage_chunk.usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15, cost=0.5) - chunks: Final = [ - _stream_builder_text_chunk("gpt-4o", "Hello "), - _stream_builder_text_chunk("gpt-4o", "world.", finish_reason="stop"), - usage_chunk, - ] - - response: Final = litellm.stream_chunk_builder( - chunks=chunks, messages=[{"role": "user", "content": "hi"}], logging_obj=_stream_builder_logging_obj() - ) - - assert response is not None - assert getattr(response.usage, "cost", None) == pytest.approx(0.5) - assert response._hidden_params["response_cost"] == pytest.approx(0.5) - - -def test_stream_chunk_builder_prices_alias_from_openai_sdk_usage_chunk(): - from openai.types.completion_usage import CompletionUsage - - usage_chunk: Final = _stream_builder_text_chunk("mantle-claude", "") - usage_chunk.usage = CompletionUsage(prompt_tokens=20, completion_tokens=60, total_tokens=80, cost=0.000704) - assert type(usage_chunk.usage) is CompletionUsage - chunks: Final = [ - _stream_builder_text_chunk("mantle-claude", "Hello "), - _stream_builder_text_chunk("mantle-claude", "world.", finish_reason="stop"), - usage_chunk, - ] - - response: Final = litellm.stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "hi"}]) - - assert response is not None - assert response.usage.prompt_tokens == 20 - assert response.usage.completion_tokens == 60 - assert getattr(response.usage, "cost", None) == pytest.approx(0.000704) - assert response._hidden_params["response_cost"] == pytest.approx(0.000704) - - -def test_stream_chunk_builder_leaves_xai_reported_cost_to_the_calculator(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(litellm, "cost_margin_config", {"xai": 0.5}) - usage_chunk: Final = _stream_builder_text_chunk("grok-4", "") - usage_chunk.usage = Usage(prompt_tokens=5, completion_tokens=2, total_tokens=7, cost=0.42) - chunks: Final = [ - _stream_builder_text_chunk("grok-4", "Hello "), - _stream_builder_text_chunk("grok-4", "world.", finish_reason="stop"), - usage_chunk, - ] - logging_obj: Final = _stream_builder_logging_obj(model="grok-4", custom_llm_provider="xai") - - response: Final = litellm.stream_chunk_builder( - chunks=chunks, messages=[{"role": "user", "content": "hi"}], logging_obj=logging_obj - ) - - assert response is not None - assert getattr(response.usage, "cost", None) == pytest.approx(0.42) - assert response._hidden_params.get("response_cost") is None - assert logging_obj._response_cost_calculator(result=response) == pytest.approx(0.63) - - -def test_speech_mistral_dispatches_and_decodes_audio(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv("MISTRAL_API_KEY", "sk-mistral-test") - audio_bytes: Final = b"ID3-fake-mp3-bytes" - mock_route: Final = respx_mock.post("https://api.mistral.ai/v1/audio/speech").mock( - return_value=httpx.Response(200, json={"audio_data": base64.b64encode(audio_bytes).decode()}) - ) - - response: Final = litellm.speech( - model="mistral/voxtral-mini-tts-2603", - input="hello from litellm", - voice="en_paul_neutral", - response_format="wav", - speed=2, - instructions="sound cheerful", - ) - - assert mock_route.called - request_body: Final = json.loads(mock_route.calls.last.request.content) - assert request_body == { - "model": "voxtral-mini-tts-2603", - "input": "hello from litellm", - "voice_id": "en_paul_neutral", - "response_format": "wav", - } - assert mock_route.calls.last.request.headers["authorization"] == "Bearer sk-mistral-test" - assert response.content == audio_bytes - - -def test_speech_mistral_routes_to_configured_api_base(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setenv("MISTRAL_API_KEY", "sk-mistral-test") - audio_bytes: Final = b"ID3-gateway-bytes" - gateway_route: Final = respx_mock.post("https://mistral.gateway.internal/v1/audio/speech").mock( - return_value=httpx.Response(200, json={"audio_data": base64.b64encode(audio_bytes).decode()}) - ) - - response: Final = litellm.speech( - model="mistral/voxtral-mini-tts-2603", - input="hello from litellm", - voice="en_paul_neutral", - api_base="https://mistral.gateway.internal", - ) - - assert gateway_route.called - assert response.content == audio_bytes - - -FOUNDRY_HOST: Final = "https://my-project.services.ai.azure.com" - - -def test_azure_ai_transcription_on_a_foundry_host_uses_the_azure_openai_deployment_route( - respx_mock: respx.MockRouter, -): - route: Final = respx_mock.post( - url__regex=r"https://my-project\.services\.ai\.azure\.com/openai/deployments/whisper-1/audio/transcriptions\?api-version=.+" - ).mock(return_value=httpx.Response(200, json={"text": "hello"})) - - response: Final = litellm.transcription( - model="azure_ai/whisper-1", - file=("tone.wav", b"RIFF\x00\x00\x00\x00WAVE", "audio/wav"), - api_base=FOUNDRY_HOST, - api_key="fake-key", - ) - - assert route.called - assert response.text == "hello" - - -def test_azure_ai_speech_on_a_foundry_host_uses_the_azure_openai_deployment_route( - respx_mock: respx.MockRouter, -): - route: Final = respx_mock.post( - url__regex=r"https://my-project\.services\.ai\.azure\.com/openai/deployments/tts-1/audio/speech\?api-version=.+" - ).mock(return_value=httpx.Response(200, content=b"mp3-bytes")) - - response: Final = litellm.speech( - model="azure_ai/tts-1", - input="hello", - voice="alloy", - api_base=FOUNDRY_HOST, - api_key="fake-key", - ) - - assert route.called - assert response.content == b"mp3-bytes" - - -FORWARDED_CLIENT_HEADERS: Final = {"x-forwarded-for": "10.0.0.1", "x-amzn-trace-id": "Root=1-lit7694"} - - -def _chat_completion_json() -> Mapping[str, object]: - return { - "id": "chatcmpl-lit7694", - "object": "chat.completion", - "created": 1, - "model": "gpt-5.4", - "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], - "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, - } - - -def _chat_completion_sse() -> bytes: - chunk: Final = { - "id": "chatcmpl-lit7694", - "object": "chat.completion.chunk", - "created": 1, - "model": "gpt-5.4", - "choices": [{"index": 0, "delta": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], - } - return f"data: {json.dumps(chunk)}\n\ndata: [DONE]\n\n".encode() - - -@pytest.mark.parametrize("stream", [False, True]) -def test_bridged_responses_with_openai_http_handler_keeps_forwarded_headers_out_of_the_body( - respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, stream: bool -): - monkeypatch.setenv("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER", "true") - route: Final = respx_mock.post("https://api.openai.com/v1/chat/completions").mock( - return_value=httpx.Response(200, content=_chat_completion_sse(), headers={"content-type": "text/event-stream"}) - if stream - else httpx.Response(200, json=_chat_completion_json()) - ) - - response: Final = litellm.responses( - model="openai/gpt-5.4", - input="Reply with the single word ok", - stream=stream, - use_chat_completions_api=True, - headers=dict(FORWARDED_CLIENT_HEADERS), - api_key="sk-test", - ) - if stream: - list(response) - - assert route.called - request: Final = route.calls.last.request - body: Final = json.loads(request.content) - assert "extra_headers" not in body - assert body["model"] == "gpt-5.4" - assert {k: request.headers[k] for k in FORWARDED_CLIENT_HEADERS} == FORWARDED_CLIENT_HEADERS - - -@pytest.mark.parametrize("http2_on", [True, False]) -def test_aiohttp_openai_warns_only_when_http2_enabled( - monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, http2_on: bool -): - from litellm.main import base_llm_aiohttp_handler - - monkeypatch.setattr(litellm, "http2", http2_on) - monkeypatch.delenv("LITELLM_HTTP2", raising=False) - - handler_completion: Final = MagicMock(return_value=MagicMock()) - monkeypatch.setattr(base_llm_aiohttp_handler, "completion", handler_completion) - - with caplog.at_level(logging.WARNING, logger="LiteLLM"): - litellm.completion( - model="aiohttp_openai/gpt-4o", - messages=[{"role": "user", "content": "hi"}], - api_key="sk-test", - ) - - assert handler_completion.called - warned: Final = "aiohttp_openai/ always uses aiohttp" in caplog.text - assert warned is http2_on - - -@pytest.mark.parametrize("tool_choice", [{"type": "bogus"}, {"name": "lookup_fruit"}, {"type": "file_search"}]) -def test_completion_rejects_untranslatable_tool_choice_with_a_400(tool_choice): - with pytest.raises(litellm.BadRequestError) as exc_info: - litellm.completion( - model="anthropic/claude-haiku-4-5", - messages=[{"role": "user", "content": "Which fruit is red?"}], - tools=[{"type": "function", "function": {"name": "lookup_fruit", "parameters": {"type": "object"}}}], - tool_choice=tool_choice, - api_key="sk-unused", - ) - assert exc_info.value.status_code == 400 - assert f"tool_choice={tool_choice}" in str(exc_info.value) diff --git a/tests/unit/batches/test_batch_utils.py b/tests/unit/batches/test_batch_utils.py index d1572f4a7c9..dd95addac40 100644 --- a/tests/unit/batches/test_batch_utils.py +++ b/tests/unit/batches/test_batch_utils.py @@ -2072,3 +2072,348 @@ def test_chat_rows_from_mistral_still_use_token_pricing(monkeypatch): ) assert result.cost == pytest.approx((10 * 0.001 + 5 * 0.002) / 2) assert result.usage.total_tokens == 15 + + +GROUNDED_USAGE_METADATA = { + "promptTokenCount": 19, + "candidatesTokenCount": 59, + "thoughtsTokenCount": 406, + "toolUsePromptTokenCount": 73, + "totalTokenCount": 557, + "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 19}], + "candidatesTokensDetails": [{"modality": "TEXT", "tokenCount": 59}], + "toolUsePromptTokensDetails": [{"modality": "TEXT", "tokenCount": 73}], + "trafficType": "ON_DEMAND", +} + + +PASSTHROUGH_OUTPUT_URI = ( + "gs://litellm-bucket/litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/u/" + "predictions.jsonl" +) + + +UNGROUNDED_USAGE_METADATA = { + "promptTokenCount": 20, + "candidatesTokenCount": 48, + "thoughtsTokenCount": 195, + "toolUsePromptTokenCount": 73, + "totalTokenCount": 336, + "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 20}], + "trafficType": "ON_DEMAND", +} + + +def _native_vertex_row(usage_metadata: dict, *, grounded: bool, model_version: str | None = "gemini-2.5-flash"): + candidate = {"content": {"role": "model", "parts": [{"text": "ok"}]}, "finishReason": "STOP"} + grounding = {"groundingMetadata": {"webSearchQueries": ["q"]}} if grounded else {} + response = {"candidates": [{**candidate, **grounding}], "usageMetadata": usage_metadata} + return { + "request": {"contents": [{"role": "user", "parts": [{"text": "q"}]}], "tools": [{"googleSearch": {}}]}, + "status": "", + "response": {**response, **({"modelVersion": model_version} if model_version else {})}, + "processed_time": "2026-09-23T19:02:00.000+00:00", + } + + +def _capture_cost_calls(monkeypatch, prompt_cost=0.5, completion_cost=0.25) -> list: + import litellm.cost_calculator as cc + + calls: list = [] + + def _calc(**kw): + calls.append(kw) + return (prompt_cost, completion_cost) + + monkeypatch.setattr(cc, "batch_cost_calculator", _calc) + return calls + + +def test_vertex_native_cost_bills_embedding_rows(monkeypatch): + monkeypatch.setitem(litellm.model_cost, "vertex_ai/gemini-embedding-2", {"input_cost_per_token_batches": 1e-7}) + rows = [ + { + "key": "id_1", + "status": "", + "request": {"content": {"parts": [{"text": "hello world"}]}}, + "response": {"embedding": {"values": [0.1, 0.2]}, "usageMetadata": {"promptTokenCount": 2}}, + }, + { + "key": "id_2", + "status": "", + "request": {"content": {"parts": [{"text": "hello"}]}}, + "response": {"embedding": {"values": [0.3]}, "tokenCount": "3"}, + }, + {"key": "id_3", "status": "INVALID_ARGUMENT", "request": {"content": {"parts": [{"text": ""}]}}}, + ] + + result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-embedding-2") + + assert (result.successful_requests, result.failed_requests) == (2, 1) + assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (5, 0, 5) + assert result.cost == pytest.approx(5 * 1e-7) + assert result.models == ["gemini-embedding-2"] + + +@pytest.mark.asyncio +async def test_native_vertex_rows_route_to_vertex_cost_path_without_flag(monkeypatch): + monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False) + monkeypatch.setattr( + bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run") + ) + calls = _capture_cost_calls(monkeypatch) + rows = [ + _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True), + _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False), + ] + + result = await bu.calculate_batch_cost_and_usage( + file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash" + ) + + assert result.cost == pytest.approx(1.5) + assert (result.successful_requests, result.failed_requests) == (2, 0) + assert result.models == ["gemini-2.5-flash"] + assert {(call["model"], call["custom_llm_provider"]) for call in calls} == {("gemini-2.5-flash", "vertex_ai")} + + +@pytest.mark.asyncio +async def test_openai_shaped_vertex_rows_keep_the_generic_path_without_flag(monkeypatch): + monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False) + monkeypatch.setattr( + bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run") + ) + _capture_cost_calls(monkeypatch) + rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)] + + result = await bu.calculate_batch_cost_and_usage( + file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash" + ) + + assert result.successful_requests == 1 + + +@pytest.mark.asyncio +async def test_native_vertex_rows_on_another_provider_keep_the_generic_path(monkeypatch): + monkeypatch.setattr( + bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run") + ) + _capture_cost_calls(monkeypatch) + + result = await bu.calculate_batch_cost_and_usage( + file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)], + custom_llm_provider="openai", + ) + + assert result.successful_requests == 0 + + +@pytest.mark.asyncio +async def test_handle_completed_batch_routes_native_rows_without_flag(monkeypatch): + monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False) + raw_rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)] + + async def fake_fetch(batch, custom_llm_provider, litellm_params=None): + return _vertex_jsonl(raw_rows) + + monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch) + monkeypatch.setattr( + bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run") + ) + calls = _capture_cost_calls(monkeypatch, prompt_cost=0.7, completion_cost=0.3) + deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6} + + result = await bu._handle_completed_batch( + _batch(PASSTHROUGH_OUTPUT_URI), + custom_llm_provider="vertex_ai", + model_name="gemini-2.5-flash", + model_info=deployment_model_info, + ) + + assert result.cost == pytest.approx(1.0) + assert result.usage.total_tokens == 557 + assert [call["model_info"] for call in calls] == [deployment_model_info] + + +def test_native_vertex_usage_is_billed_like_the_online_path(monkeypatch): + calls = _capture_cost_calls(monkeypatch) + grounded = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True) + ungrounded = _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False) + + result = bu.calculate_vertex_ai_batch_cost_and_usage([grounded, ungrounded], "gemini-2.5-flash") + + grounded_usage, ungrounded_usage = (call["usage"] for call in calls) + assert grounded_usage.prompt_tokens == 19 + assert grounded_usage.completion_tokens == 59 + 406 + assert grounded_usage.completion_tokens_details.reasoning_tokens == 406 + assert ungrounded_usage.prompt_tokens == 20 + 73 + assert ungrounded_usage.completion_tokens == 48 + 195 + assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == ( + 19 + 93, + 465 + 243, + 557 + 336, + ) + + +def test_native_vertex_rows_are_priced_by_model_version_without_a_model_name(monkeypatch): + calls = _capture_cost_calls(monkeypatch) + rows = [ + _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"), + _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-pro"), + _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None), + ] + + result = bu.calculate_vertex_ai_batch_cost_and_usage(rows) + + assert [call["model"] for call in calls] == ["gemini-2.5-flash", "gemini-2.5-pro"] + assert result.models == ["gemini-2.5-flash", "gemini-2.5-pro"] + assert result.cost == pytest.approx(1.5) + assert result.successful_requests == 3 + assert result.usage.total_tokens == 557 + 336 + 336 + + +def test_native_vertex_rows_without_usage_metadata_count_as_failed(monkeypatch): + _capture_cost_calls(monkeypatch) + rows = [ + {"request": {"contents": []}, "status": "Error: bad request", "processed_time": "t"}, + {"request": {"contents": []}, "response": {"candidates": []}}, + _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True), + ] + + result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash") + + assert (result.successful_requests, result.failed_requests) == (1, 2) + assert result.usage.total_tokens == 557 + + +def test_native_vertex_batch_whose_rows_all_failed_still_names_the_deployment_model(monkeypatch): + calls = _capture_cost_calls(monkeypatch) + rows = [{"request": {"contents": []}, "status": "Error: quota exceeded", "processed_time": "t"}] * 2 + + result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash") + + assert result.models == ["gemini-2.5-flash"] + assert (result.successful_requests, result.failed_requests, result.cost) == (0, 2, 0.0) + assert calls == [] + + +def test_native_vertex_rows_are_priced_with_the_deployment_model_info(monkeypatch): + calls = _capture_cost_calls(monkeypatch) + deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6} + + bu.calculate_vertex_ai_batch_cost_and_usage( + [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)], + "gemini-2.5-flash", + model_info=deployment_model_info, + ) + + assert [call["model_info"] for call in calls] == [deployment_model_info] + + +@pytest.mark.asyncio +async def test_native_vertex_rows_keep_the_deployment_model_info_through_the_batch_entrypoint(monkeypatch): + calls = _capture_cost_calls(monkeypatch) + deployment_model_info = {"input_cost_per_token_batches": 1e-6} + + await bu.calculate_batch_cost_and_usage( + file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)], + custom_llm_provider="vertex_ai", + model_name="gemini-2.5-flash", + model_info=deployment_model_info, + ) + + assert [call["model_info"] for call in calls] == [deployment_model_info] + + +def test_native_vertex_rows_are_priced_by_the_deployment_model_over_model_version(monkeypatch): + calls = _capture_cost_calls(monkeypatch) + rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-pro")] + + result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash") + + assert [call["model"] for call in calls] == ["gemini-2.5-flash"] + assert result.models == ["gemini-2.5-flash"] + + +def test_native_vertex_rows_that_fail_response_validation_count_as_failed(monkeypatch): + calls = _capture_cost_calls(monkeypatch) + rows = [ + {"request": {"contents": []}, "response": {"candidates": "nope", "usageMetadata": GROUNDED_USAGE_METADATA}}, + _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True), + ] + + result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash") + + assert (result.successful_requests, result.failed_requests) == (1, 1) + assert result.usage.total_tokens == 557 + assert len(calls) == 1 + + +@pytest.mark.parametrize("wildcard_model", ["*", "vertex_ai/*"]) +def test_native_vertex_rows_under_a_wildcard_deployment_are_priced_by_model_version(monkeypatch, wildcard_model): + calls = _capture_cost_calls(monkeypatch) + rows = [ + _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"), + _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None), + ] + + result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, wildcard_model) + + assert [call["model"] for call in calls] == ["gemini-2.5-flash", wildcard_model] + assert result.cost == pytest.approx(1.5) + assert (result.successful_requests, result.failed_requests) == (2, 0) + assert result.usage.total_tokens == 557 + 336 + + +def test_native_vertex_row_without_model_version_under_a_wildcard_deployment_bills_its_explicit_prices(): + deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6} + with_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash") + without_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version=None) + + twin = bu.calculate_vertex_ai_batch_cost_and_usage([with_version], "vertex_ai/*", model_info=deployment_model_info) + both = bu.calculate_vertex_ai_batch_cost_and_usage( + [with_version, without_version], "vertex_ai/*", model_info=deployment_model_info + ) + + assert twin.cost > 0 + assert both.cost == pytest.approx(2 * twin.cost) + assert (both.successful_requests, both.failed_requests) == (2, 0) + + +def test_native_vertex_row_the_cost_map_cannot_price_is_billed_at_zero_and_the_rest_still_bills(monkeypatch): + import litellm.cost_calculator as cc + + def _calc(**kw): + if kw["model"] == "gemini-unpriced": + raise ValueError("no pricing") + return (0.5, 0.25) + + monkeypatch.setattr(cc, "batch_cost_calculator", _calc) + rows = [ + _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-unpriced"), + _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-flash"), + ] + + result = bu.calculate_vertex_ai_batch_cost_and_usage(rows) + + assert result.cost == pytest.approx(0.75) + assert (result.successful_requests, result.failed_requests) == (2, 0) + assert result.usage.total_tokens == 557 + 336 + assert result.models == ["gemini-unpriced", "gemini-2.5-flash"] + + +@pytest.mark.asyncio +async def test_flag_sends_every_vertex_row_down_the_native_path_when_a_model_is_known(monkeypatch): + monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False) + monkeypatch.setattr( + bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run") + ) + calls = _capture_cost_calls(monkeypatch) + rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)] + + result = await bu.calculate_batch_cost_and_usage( + file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash" + ) + + assert calls == [] + assert (result.successful_requests, result.failed_requests) == (0, 1) diff --git a/tests/unit/chat_completions/test_dispatch.py b/tests/unit/chat_completions/test_dispatch.py index 2807ed7f8f7..40b1c0ef019 100644 --- a/tests/unit/chat_completions/test_dispatch.py +++ b/tests/unit/chat_completions/test_dispatch.py @@ -20,6 +20,8 @@ from litellm.rust_bridge.chat_completions.entrypoints import ( ) from litellm.rust_bridge.configuration import Rollout from litellm.types.utils import ModelResponse +from litellm.chat_completions import dispatch +from litellm.rust_bridge.catalog import Rules MESSAGES: Final = [{"role": "user", "content": "hi"}] PYTHON_RULES: Final = () @@ -256,3 +258,100 @@ async def test_public_acompletion_routes_through_dispatch(monkeypatch: pytest.Mo NATIVE_ACOMPLETION.reset() assert result is expected assert [request.model for request in captured] == ["gpt-4o"] + + +@pytest.mark.asyncio +async def test_public_completion_calls_keep_the_python_result() -> None: + sync_response: Final = litellm.completion(model="openai/test-model", messages=MESSAGES, mock_response="ok") + async_response: Final = await litellm.acompletion(model="openai/test-model", messages=MESSAGES, mock_response="ok") + + assert isinstance(sync_response, ModelResponse) + assert isinstance(async_response, ModelResponse) + assert sync_response.choices[0].message.content == "ok" + assert async_response.choices[0].message.content == "ok" + + +def test_sync_completion_request_projects_public_arguments() -> None: + rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),) + expected: Final = ModelResponse() + + def native( + request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] + ) -> ModelResponse: + assert request.model == "test-model" + assert request.messages == MESSAGES + assert request.custom_llm_provider == "openai" + assert request.stream is True + return expected + + binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None) + binding.override(native) + response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision + ("test-model", MESSAGES), + {"custom_llm_provider": "openai", "stream": True}, + python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"), + binding=binding, + native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + rules=rules, + ) + + assert response is expected + + +@pytest.mark.asyncio +async def test_async_completion_falls_back_after_native_declines() -> None: + from litellm.rust_bridge.bindings import native_exception_types + + native_types: Final = native_exception_types() + if native_types is None: + pytest.skip("native bridge is unavailable") + declined, _ = native_types + expected: Final = ModelResponse() + rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_OPT_OUT),) + + async def native( + request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] + ) -> ModelResponse: + raise declined("unsupported") + + async def python(*args: object, **kwargs: object) -> ModelResponse: + return expected + + binding: Final[NativeBinding[NativeAcompletion]] = NativeBinding("acompletion", validate=lambda _: None) + binding.override(native) + response: Final = await dispatch._ADISPATCH.arun( # pyright: ignore[reportPrivateUsage] # test an explicit route decision + ("test-model", MESSAGES), + {}, + python=python, + binding=binding, + native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + rules=rules, + ) + + assert response is expected + + +def test_internal_acompletion_marker_bypasses_native() -> None: + rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),) + expected: Final = ModelResponse() + + def python(*args: object, **kwargs: object) -> ModelResponse: + return expected + + def native( + request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] + ) -> ModelResponse: + pytest.fail("acompletion's inner completion call must stay on Python") + + binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None) + binding.override(native) + response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision + ("test-model", MESSAGES), + {"custom_llm_provider": "openai", "acompletion": True}, + python=python, + binding=binding, + native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + rules=rules, + ) + + assert response is expected diff --git a/tests/test_litellm/a2a_protocol/__init__.py b/tests/unit/completion_extras/litellm_responses_transformation/__init__.py similarity index 100% rename from tests/test_litellm/a2a_protocol/__init__.py rename to tests/unit/completion_extras/litellm_responses_transformation/__init__.py diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py similarity index 100% rename from tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py rename to tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py similarity index 100% rename from tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py rename to tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py diff --git a/tests/unit/conftest.py b/tests/unit/conftest.py index 202ecb80d7b..ecea4723bf4 100644 --- a/tests/unit/conftest.py +++ b/tests/unit/conftest.py @@ -1,7 +1,14 @@ +import asyncio +import base64 +import importlib import os -from collections.abc import Iterator +from collections.abc import Coroutine, Iterator +from dataclasses import dataclass, field +from pathlib import Path from typing import Final +import boto3 +import httpx import pytest from pytest_socket import enable_socket, socket_allow_hosts @@ -10,6 +17,17 @@ os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" import litellm # noqa: E402 # litellm reads LITELLM_LOCAL_MODEL_COST_MAP at import import litellm.router as litellm_router_module # noqa: E402 # same import-time dependency import litellm.utils as litellm_utils_module # noqa: E402 # same import-time dependency +from litellm._logging import ALL_LOGGERS # noqa: E402 # same import-time dependency +from litellm.anthropic_beta_headers_manager import reload_beta_headers_config # noqa: E402 # same import-time dependency +from litellm.litellm_core_utils.prompt_templates import factory as prompt_factory_module # noqa: E402 # same import-time dependency +from litellm.litellm_core_utils.prompt_templates import ( # noqa: E402 # same import-time dependency + image_handling as image_handling_module, +) +from litellm.llms.gemini.chat import transformation as gemini_chat_transformation_module # noqa: E402 # same import-time dependency +from litellm.llms.custom_httpx.async_client_cleanup import ( # noqa: E402 # same import-time dependency + close_litellm_async_clients, +) +from litellm.proxy.db import tool_registry_writer as tool_registry_writer_module # noqa: E402 # same import-time dependency LOOPBACK_HOSTS: Final = ["127.0.0.1", "::1", "localhost"] AMBIENT_AZURE_CREDENTIAL_ENV_VARS: Final = ( @@ -20,6 +38,66 @@ AMBIENT_AZURE_CREDENTIAL_ENV_VARS: Final = ( "AZURE_USERNAME", "AZURE_PASSWORD", ) +AMBIENT_AWS_ENV_VARS: Final = ( + "AWS_PROFILE", + "AWS_DEFAULT_PROFILE", + "AWS_CONTAINER_CREDENTIALS_FULL_URI", + "AWS_CONTAINER_CREDENTIALS_RELATIVE_URI", + "AWS_SESSION_TOKEN", + "AWS_ROLE_ARN", + "AWS_WEB_IDENTITY_TOKEN_FILE", + "AWS_BEARER_TOKEN_BEDROCK", + "AWS_REGION_NAME", + "AWS_DEFAULT_REGION", +) +MODULES_WITH_AWS_AUTH_HANDLERS: Final = ( + "litellm.main", + "litellm.files.main", + "litellm.rerank_api.main", + "litellm.realtime_api.main", +) +CALLBACK_LISTS: Final = ( + "callbacks", + "success_callback", + "failure_callback", + "input_callback", + "_async_success_callback", + "_async_failure_callback", + "_async_input_callback", +) +RESET_TO_NONE_GLOBALS: Final = ("model_fallbacks", "cache") +RESTORED_GLOBALS: Final = ( + "disable_aiohttp_transport", + "force_ipv4", + "drop_params", + "secret_manager_client", + "_key_management_system", + "_key_management_settings", + "api_base", + "num_retries", + "modify_params", + "ssl_verify", + "credential_list", + "model_group_settings", + "default_internal_user_params", + "default_team_params", + "prometheus_emit_stream_label", + "vector_store_registry", + "model_cost", + "cost_margin_config", + "cost_discount_config", + "disable_hf_tokenizer_download", + "disable_copilot_system_to_assistant", + "cohere_models", + "anthropic_models", + "token_counter", + "initialized_langfuse_clients", +) +MODULE_LEVEL_CLIENTS: Final = ("module_level_client", "module_level_aclient") +SESSION_CLIENTS: Final = ("base_llm_aiohttp_handler", "httpx_client", "aclient", "client") +ONE_PIXEL_PNG: Final = base64.b64decode( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg==" +) def _allow_loopback_only() -> None: @@ -29,11 +107,116 @@ def _allow_loopback_only() -> None: _allow_loopback_only() +def pytest_collectstart() -> None: + _allow_loopback_only() + + @pytest.hookimpl(trylast=True) def pytest_runtest_setup() -> None: _allow_loopback_only() +def _run_coroutine_if_needed(result: object) -> None: + if not asyncio.iscoroutine(result): + return + coroutine: Final[Coroutine[object, object, object]] = result + try: + asyncio.run(coroutine) + except RuntimeError: + try: + loop: Final = asyncio.get_running_loop() + except RuntimeError: + coroutine.close() + return + loop.create_task(coroutine) + + +def _close_handler_if_needed(handler: object) -> None: + close: Final = getattr(handler, "close", None) + if not callable(close): + return + _run_coroutine_if_needed(close()) + + +def _reset_aws_auth_caches() -> None: + modules: Final = tuple(importlib.import_module(name) for name in MODULES_WITH_AWS_AUTH_HANDLERS) + flushes: Final = ( + getattr(getattr(getattr(module, attr_name), "iam_cache", None), "flush_cache", None) + for module in modules + for attr_name in dir(module) + ) + for flush in filter(callable, flushes): + flush() + boto3.DEFAULT_SESSION = None + + +def _flush_client_caches() -> None: + litellm.in_memory_llm_clients_cache.flush_cache() + image_handling_module.in_memory_cache.flush_cache() + _reset_aws_auth_caches() + + +@pytest.fixture(scope="session") +def isolated_aws_config_files(tmp_path_factory: pytest.TempPathFactory) -> tuple[Path, Path]: + aws_dir: Final = tmp_path_factory.mktemp("aws-config") + credentials: Final = aws_dir / "credentials" + config: Final = aws_dir / "config" + credentials.write_text("", encoding="utf-8") + config.write_text("", encoding="utf-8") + return credentials, config + + +@pytest.fixture(autouse=True) +def isolate_host_environment(isolated_aws_config_files: tuple[Path, Path]) -> Iterator[None]: + credentials, config = isolated_aws_config_files + with pytest.MonkeyPatch.context() as environment: + environment.setenv("AWS_SHARED_CREDENTIALS_FILE", str(credentials)) + environment.setenv("AWS_CONFIG_FILE", str(config)) + environment.setenv("AWS_EC2_METADATA_DISABLED", "true") + for name in AMBIENT_AWS_ENV_VARS: + environment.delenv(name, raising=False) + environment.delenv("PROXY_BASE_URL", raising=False) + environment.setenv("LITELLM_CLI_DISABLE_KEYRING", "1") + yield + + +@pytest.fixture(autouse=True) +def isolate_litellm_globals() -> Iterator[None]: + original_callbacks: Final = {name: list(getattr(litellm, name) or []) for name in CALLBACK_LISTS} + original_reset: Final = {name: getattr(litellm, name) for name in RESET_TO_NONE_GLOBALS} + original_restored: Final = {name: getattr(litellm, name) for name in RESTORED_GLOBALS if hasattr(litellm, name)} + original_clients: Final = {name: litellm.__dict__[name] for name in MODULE_LEVEL_CLIENTS if name in litellm.__dict__} + original_loggers: Final = { + logger: (logger.level, logger.disabled, logger.propagate, list(logger.handlers), list(logger.filters)) + for logger in ALL_LOGGERS + } + original_tool_policy_registry: Final = tool_registry_writer_module._tool_policy_registry + _flush_client_caches() + for name in CALLBACK_LISTS: + setattr(litellm, name, []) + for name in RESET_TO_NONE_GLOBALS: + setattr(litellm, name, None) + for name in MODULE_LEVEL_CLIENTS: + litellm.__dict__.pop(name, None) + tool_registry_writer_module._tool_policy_registry = None + yield + _flush_client_caches() + leaked_clients: Final = tuple(litellm.__dict__.pop(name, None) for name in MODULE_LEVEL_CLIENTS) + for name, client in zip(MODULE_LEVEL_CLIENTS, leaked_clients): + if client is not original_clients.get(name): + _close_handler_if_needed(client) + litellm.__dict__.update(original_clients) + for name, value in (original_callbacks | original_reset | original_restored).items(): + setattr(litellm, name, value) + for logger, (level, disabled, propagate, handlers, filters) in original_loggers.items(): + logger.setLevel(level) + logger.disabled = disabled + logger.propagate = propagate + logger.handlers = handlers + logger.filters = filters + tool_registry_writer_module._tool_policy_registry = original_tool_policy_registry + + @pytest.fixture(autouse=True) def isolate_router_model_cost_state() -> Iterator[None]: original_live_routers: Final = frozenset(litellm_router_module._live_routers) @@ -41,6 +224,7 @@ def isolate_router_model_cost_state() -> Iterator[None]: model_key: dict(model_value) for model_key, model_value in litellm_utils_module._runtime_registered_model_cost.items() } + litellm_utils_module._invalidate_model_cost_lowercase_map() yield for router in tuple(litellm_router_module._live_routers): litellm_router_module._live_routers.discard(router) @@ -61,6 +245,47 @@ def local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: litellm.get_model_info.cache_clear() +@pytest.fixture +def local_beta_headers_config(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.setenv("LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS", "True") + reload_beta_headers_config() + yield + monkeypatch.delenv("LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS", raising=False) + reload_beta_headers_config() + + +@dataclass(slots=True) +class AsyncOnlyImageFetch: + fetched: list[str] = field(default_factory=list) # mutable-ok: tests assert on the URLs fetched, in order + base64_png: str = base64.b64encode(ONE_PIXEL_PNG).decode() + data_url: str = "data:image/png;base64," + base64.b64encode(ONE_PIXEL_PNG).decode() + + +@pytest.fixture +def async_only_image_fetch(monkeypatch: pytest.MonkeyPatch) -> AsyncOnlyImageFetch: + fetch: Final = AsyncOnlyImageFetch() + + def forbid_sync_fetch(client: object, url: str, **kwargs: object) -> httpx.Response: + raise litellm.ImageFetchError(f"sync image fetch ran on the event loop: {url}") + + async def serve_png(client: object, url: str, **kwargs: object) -> httpx.Response: + fetch.fetched.append(url) + return httpx.Response( + 200, content=ONE_PIXEL_PNG, headers={"content-type": "image/png"}, request=httpx.Request("GET", url) + ) + + def forbid_sync_convert(url: str, *args: object, **kwargs: object) -> str: + if url.startswith(("http://", "https://")): + raise litellm.ImageFetchError(f"sync convert_url_to_base64 ran on the request path: {url}") + return url + + monkeypatch.setattr(image_handling_module, "safe_get", forbid_sync_fetch) + monkeypatch.setattr(image_handling_module, "async_safe_get", serve_png) + for module in (image_handling_module, prompt_factory_module, gemini_chat_transformation_module): + monkeypatch.setattr(module, "convert_url_to_base64", forbid_sync_convert) + return fetch + + @pytest.fixture def no_ambient_azure_credentials(monkeypatch: pytest.MonkeyPatch) -> None: for name in AMBIENT_AZURE_CREDENTIAL_ENV_VARS: @@ -68,4 +293,9 @@ def no_ambient_azure_credentials(monkeypatch: pytest.MonkeyPatch) -> None: def pytest_sessionfinish() -> None: + for name in MODULE_LEVEL_CLIENTS: + _close_handler_if_needed(litellm.__dict__.pop(name, None)) + for name in SESSION_CLIENTS: + _close_handler_if_needed(getattr(litellm, name, None)) + _run_coroutine_if_needed(close_litellm_async_clients()) enable_socket() diff --git a/tests/test_litellm/a2a_protocol/providers/__init__.py b/tests/unit/containers/__init__.py similarity index 100% rename from tests/test_litellm/a2a_protocol/providers/__init__.py rename to tests/unit/containers/__init__.py diff --git a/tests/test_litellm/containers/test_azure_container_transformation.py b/tests/unit/containers/test_azure_container_transformation.py similarity index 100% rename from tests/test_litellm/containers/test_azure_container_transformation.py rename to tests/unit/containers/test_azure_container_transformation.py diff --git a/tests/test_litellm/containers/test_container_api.py b/tests/unit/containers/test_container_api.py similarity index 100% rename from tests/test_litellm/containers/test_container_api.py rename to tests/unit/containers/test_container_api.py diff --git a/tests/test_litellm/containers/test_container_handler_url.py b/tests/unit/containers/test_container_handler_url.py similarity index 100% rename from tests/test_litellm/containers/test_container_handler_url.py rename to tests/unit/containers/test_container_handler_url.py diff --git a/tests/test_litellm/containers/test_container_integration.py b/tests/unit/containers/test_container_integration.py similarity index 100% rename from tests/test_litellm/containers/test_container_integration.py rename to tests/unit/containers/test_container_integration.py diff --git a/tests/test_litellm/containers/test_container_proxy_ownership.py b/tests/unit/containers/test_container_proxy_ownership.py similarity index 100% rename from tests/test_litellm/containers/test_container_proxy_ownership.py rename to tests/unit/containers/test_container_proxy_ownership.py diff --git a/tests/test_litellm/containers/test_container_regional_api_base.py b/tests/unit/containers/test_container_regional_api_base.py similarity index 100% rename from tests/test_litellm/containers/test_container_regional_api_base.py rename to tests/unit/containers/test_container_regional_api_base.py diff --git a/tests/test_litellm/containers/test_container_transformation.py b/tests/unit/containers/test_container_transformation.py similarity index 100% rename from tests/test_litellm/containers/test_container_transformation.py rename to tests/unit/containers/test_container_transformation.py diff --git a/tests/test_litellm/containers/test_container_utils.py b/tests/unit/containers/test_container_utils.py similarity index 100% rename from tests/test_litellm/containers/test_container_utils.py rename to tests/unit/containers/test_container_utils.py diff --git a/tests/test_litellm/containers/test_endpoint_factory.py b/tests/unit/containers/test_endpoint_factory.py similarity index 100% rename from tests/test_litellm/containers/test_endpoint_factory.py rename to tests/unit/containers/test_endpoint_factory.py diff --git a/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/__init__.py b/tests/unit/embeddings/__init__.py similarity index 100% rename from tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/__init__.py rename to tests/unit/embeddings/__init__.py diff --git a/tests/test_litellm/embeddings/test_dispatch.py b/tests/unit/embeddings/test_dispatch.py similarity index 100% rename from tests/test_litellm/embeddings/test_dispatch.py rename to tests/unit/embeddings/test_dispatch.py diff --git a/tests/test_litellm/a2a_protocol/providers/pydantic_ai_agents/__init__.py b/tests/unit/expected_fine_tuning_api/__init__.py similarity index 100% rename from tests/test_litellm/a2a_protocol/providers/pydantic_ai_agents/__init__.py rename to tests/unit/expected_fine_tuning_api/__init__.py diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_cancel_expected_output.json b/tests/unit/expected_fine_tuning_api/azure_cancel_expected_output.json similarity index 100% rename from tests/test_litellm/expected_fine_tuning_api/azure_cancel_expected_output.json rename to tests/unit/expected_fine_tuning_api/azure_cancel_expected_output.json diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_cancel_raw_response.json b/tests/unit/expected_fine_tuning_api/azure_cancel_raw_response.json similarity index 100% rename from tests/test_litellm/expected_fine_tuning_api/azure_cancel_raw_response.json rename to tests/unit/expected_fine_tuning_api/azure_cancel_raw_response.json diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_cancel_request.json b/tests/unit/expected_fine_tuning_api/azure_cancel_request.json similarity index 100% rename from tests/test_litellm/expected_fine_tuning_api/azure_cancel_request.json rename to tests/unit/expected_fine_tuning_api/azure_cancel_request.json diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_create_expected_output.json b/tests/unit/expected_fine_tuning_api/azure_create_expected_output.json similarity index 100% rename from tests/test_litellm/expected_fine_tuning_api/azure_create_expected_output.json rename to tests/unit/expected_fine_tuning_api/azure_create_expected_output.json diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_create_raw_response.json b/tests/unit/expected_fine_tuning_api/azure_create_raw_response.json similarity index 100% rename from tests/test_litellm/expected_fine_tuning_api/azure_create_raw_response.json rename to tests/unit/expected_fine_tuning_api/azure_create_raw_response.json diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_create_request.json b/tests/unit/expected_fine_tuning_api/azure_create_request.json similarity index 100% rename from tests/test_litellm/expected_fine_tuning_api/azure_create_request.json rename to tests/unit/expected_fine_tuning_api/azure_create_request.json diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_list_raw_response.json b/tests/unit/expected_fine_tuning_api/azure_list_raw_response.json similarity index 100% rename from tests/test_litellm/expected_fine_tuning_api/azure_list_raw_response.json rename to tests/unit/expected_fine_tuning_api/azure_list_raw_response.json diff --git a/tests/test_litellm/expected_fine_tuning_api/azure_list_request.json b/tests/unit/expected_fine_tuning_api/azure_list_request.json similarity index 100% rename from tests/test_litellm/expected_fine_tuning_api/azure_list_request.json rename to tests/unit/expected_fine_tuning_api/azure_list_request.json diff --git a/tests/test_litellm/batches/__init__.py b/tests/unit/experimental_mcp_client/__init__.py similarity index 100% rename from tests/test_litellm/batches/__init__.py rename to tests/unit/experimental_mcp_client/__init__.py diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py similarity index 100% rename from tests/test_litellm/experimental_mcp_client/test_mcp_client.py rename to tests/unit/experimental_mcp_client/test_mcp_client.py diff --git a/tests/test_litellm/experimental_mcp_client/test_tools.py b/tests/unit/experimental_mcp_client/test_tools.py similarity index 100% rename from tests/test_litellm/experimental_mcp_client/test_tools.py rename to tests/unit/experimental_mcp_client/test_tools.py diff --git a/tests/test_litellm/chat_completions/__init__.py b/tests/unit/files/__init__.py similarity index 100% rename from tests/test_litellm/chat_completions/__init__.py rename to tests/unit/files/__init__.py diff --git a/tests/test_litellm/files/test_main.py b/tests/unit/files/test_main.py similarity index 100% rename from tests/test_litellm/files/test_main.py rename to tests/unit/files/test_main.py diff --git a/tests/test_litellm/completion_extras/__init__.py b/tests/unit/fixtures/__init__.py similarity index 100% rename from tests/test_litellm/completion_extras/__init__.py rename to tests/unit/fixtures/__init__.py diff --git a/tests/test_litellm/containers/__init__.py b/tests/unit/fixtures/together_ai_sync/__init__.py similarity index 100% rename from tests/test_litellm/containers/__init__.py rename to tests/unit/fixtures/together_ai_sync/__init__.py diff --git a/tests/test_litellm/fixtures/together_ai_sync/deprecations.md b/tests/unit/fixtures/together_ai_sync/deprecations.md similarity index 100% rename from tests/test_litellm/fixtures/together_ai_sync/deprecations.md rename to tests/unit/fixtures/together_ai_sync/deprecations.md diff --git a/tests/test_litellm/fixtures/together_ai_sync/models_serverless.json b/tests/unit/fixtures/together_ai_sync/models_serverless.json similarity index 100% rename from tests/test_litellm/fixtures/together_ai_sync/models_serverless.json rename to tests/unit/fixtures/together_ai_sync/models_serverless.json diff --git a/tests/test_litellm/endpoints/__init__.py b/tests/unit/google_genai/__init__.py similarity index 100% rename from tests/test_litellm/endpoints/__init__.py rename to tests/unit/google_genai/__init__.py diff --git a/tests/test_litellm/google_genai/test_google_genai_adapter.py b/tests/unit/google_genai/test_google_genai_adapter.py similarity index 100% rename from tests/test_litellm/google_genai/test_google_genai_adapter.py rename to tests/unit/google_genai/test_google_genai_adapter.py diff --git a/tests/test_litellm/google_genai/test_google_genai_adapter_fixes.py b/tests/unit/google_genai/test_google_genai_adapter_fixes.py similarity index 100% rename from tests/test_litellm/google_genai/test_google_genai_adapter_fixes.py rename to tests/unit/google_genai/test_google_genai_adapter_fixes.py diff --git a/tests/test_litellm/google_genai/test_google_genai_handler.py b/tests/unit/google_genai/test_google_genai_handler.py similarity index 76% rename from tests/test_litellm/google_genai/test_google_genai_handler.py rename to tests/unit/google_genai/test_google_genai_handler.py index bf037c59854..5361d91718d 100644 --- a/tests/test_litellm/google_genai/test_google_genai_handler.py +++ b/tests/unit/google_genai/test_google_genai_handler.py @@ -2,99 +2,13 @@ """ Test to verify the Google GenAI generate_content handler functionality """ -import json from unittest.mock import AsyncMock, MagicMock, patch import pytest -import litellm from litellm.google_genai.adapters.handler import GenerateContentToCompletionHandler from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter -from litellm.types.utils import ModelResponse - - -def test_non_stream_response_when_stream_requested_sync(): - """ - Test that when a non-stream response is returned but streaming was requested, - the sync handler correctly transforms it to generate_content format. - """ - from litellm.types.utils import Choices - - # Mock a non-stream response (ModelResponse with valid choices) - mock_response = ModelResponse( - id="test-123", - choices=[ - Choices( - index=0, - message={"role": "assistant", "content": "Hello, world!"}, - finish_reason="stop", - ) - ], - created=1234567890, - model="gpt-3.5-turbo", - object="chat.completion", - ) - - # Create an instance of the adapter - adapter = GoogleGenAIAdapter() - - # Test the adapter's translate_completion_to_generate_content method directly - result = adapter.translate_completion_to_generate_content(mock_response) - - # Verify the result is a valid Google GenAI format response - assert "candidates" in result - assert isinstance(result["candidates"], list) - assert len(result["candidates"]) > 0 - candidate = result["candidates"][0] - assert "content" in candidate - assert "parts" in candidate["content"] - assert isinstance(candidate["content"]["parts"], list) - assert len(candidate["content"]["parts"]) > 0 - assert "text" in candidate["content"]["parts"][0] - assert candidate["content"]["parts"][0]["text"] == "Hello, world!" - - -@pytest.mark.asyncio -async def test_non_stream_response_when_stream_requested_async(): - """ - Test that when a non-stream response is returned but streaming was requested, - the async handler correctly transforms it to generate_content format. - """ - from litellm.types.utils import Choices - - # Mock a non-stream response (ModelResponse with valid choices) - mock_response = ModelResponse( - id="test-123", - choices=[ - Choices( - index=0, - message={"role": "assistant", "content": "Hello, world!"}, - finish_reason="stop", - ) - ], - created=1234567890, - model="gpt-3.5-turbo", - object="chat.completion", - ) - - # Create an instance of the adapter - adapter = GoogleGenAIAdapter() - - # Test the adapter's translate_completion_to_generate_content method directly - result = adapter.translate_completion_to_generate_content(mock_response) - - # Verify the result is a valid Google GenAI format response - assert "candidates" in result - assert isinstance(result["candidates"], list) - assert len(result["candidates"]) > 0 - candidate = result["candidates"][0] - assert "content" in candidate - assert "parts" in candidate["content"] - assert isinstance(candidate["content"]["parts"], list) - assert len(candidate["content"]["parts"]) > 0 - assert "text" in candidate["content"]["parts"][0] - assert candidate["content"]["parts"][0]["text"] == "Hello, world!" def test_stream_response_when_stream_requested_sync(): diff --git a/tests/test_litellm/google_genai/test_google_genai_main.py b/tests/unit/google_genai/test_google_genai_main.py similarity index 100% rename from tests/test_litellm/google_genai/test_google_genai_main.py rename to tests/unit/google_genai/test_google_genai_main.py diff --git a/tests/test_litellm/google_genai/test_google_genai_streaming_iterator.py b/tests/unit/google_genai/test_google_genai_streaming_iterator.py similarity index 100% rename from tests/test_litellm/google_genai/test_google_genai_streaming_iterator.py rename to tests/unit/google_genai/test_google_genai_streaming_iterator.py diff --git a/tests/test_litellm/google_genai/test_google_genai_transformation.py b/tests/unit/google_genai/test_google_genai_transformation.py similarity index 100% rename from tests/test_litellm/google_genai/test_google_genai_transformation.py rename to tests/unit/google_genai/test_google_genai_transformation.py diff --git a/tests/test_litellm/endpoints/speech/__init__.py b/tests/unit/images/__init__.py similarity index 100% rename from tests/test_litellm/endpoints/speech/__init__.py rename to tests/unit/images/__init__.py diff --git a/tests/test_litellm/images/test_image_edit_extra_params.py b/tests/unit/images/test_image_edit_extra_params.py similarity index 100% rename from tests/test_litellm/images/test_image_edit_extra_params.py rename to tests/unit/images/test_image_edit_extra_params.py diff --git a/tests/test_litellm/images/test_image_edit_utils.py b/tests/unit/images/test_image_edit_utils.py similarity index 100% rename from tests/test_litellm/images/test_image_edit_utils.py rename to tests/unit/images/test_image_edit_utils.py diff --git a/tests/test_litellm/images/test_image_generation_extra_headers.py b/tests/unit/images/test_image_generation_extra_headers.py similarity index 100% rename from tests/test_litellm/images/test_image_generation_extra_headers.py rename to tests/unit/images/test_image_generation_extra_headers.py diff --git a/tests/test_litellm/endpoints/speech/speech_to_completion_bridge/__init__.py b/tests/unit/interactions/__init__.py similarity index 100% rename from tests/test_litellm/endpoints/speech/speech_to_completion_bridge/__init__.py rename to tests/unit/interactions/__init__.py diff --git a/tests/test_litellm/interactions/test_agents_http_handler.py b/tests/unit/interactions/test_agents_http_handler.py similarity index 100% rename from tests/test_litellm/interactions/test_agents_http_handler.py rename to tests/unit/interactions/test_agents_http_handler.py diff --git a/tests/test_litellm/interactions/test_agents_main_and_utils.py b/tests/unit/interactions/test_agents_main_and_utils.py similarity index 100% rename from tests/test_litellm/interactions/test_agents_main_and_utils.py rename to tests/unit/interactions/test_agents_main_and_utils.py diff --git a/tests/test_litellm/interactions/test_background_cost_polling.py b/tests/unit/interactions/test_background_cost_polling.py similarity index 100% rename from tests/test_litellm/interactions/test_background_cost_polling.py rename to tests/unit/interactions/test_background_cost_polling.py diff --git a/tests/test_litellm/interactions/test_gemini_interactions_transformation.py b/tests/unit/interactions/test_gemini_interactions_transformation.py similarity index 100% rename from tests/test_litellm/interactions/test_gemini_interactions_transformation.py rename to tests/unit/interactions/test_gemini_interactions_transformation.py diff --git a/tests/test_litellm/interactions/test_interactions_streaming_iterator.py b/tests/unit/interactions/test_interactions_streaming_iterator.py similarity index 100% rename from tests/test_litellm/interactions/test_interactions_streaming_iterator.py rename to tests/unit/interactions/test_interactions_streaming_iterator.py diff --git a/tests/unit/interactions/test_litellm_responses_bridge.py b/tests/unit/interactions/test_litellm_responses_bridge.py new file mode 100644 index 00000000000..3abd0a6ca98 --- /dev/null +++ b/tests/unit/interactions/test_litellm_responses_bridge.py @@ -0,0 +1,80 @@ +""" +Tests for LiteLLM Responses bridge provider. + +Inherits from BaseInteractionsTest to run the same test suite against +the litellm_responses bridge provider, which calls litellm.responses() internally. +""" + + +from litellm.interactions.litellm_responses_transformation.transformation import ( + LiteLLMResponsesInteractionsConfig, +) +from litellm.types.interactions import Turn + + +class TestBridgeInputTransformation: + """Regression tests for translating Interactions input into Responses API input. + + The bridge used to pass Google content parts through raw ({"type": "text"}), + which the Responses API rejects with a 400, and it dropped the role encoded + in step types and in the legacy "model" turn role. + """ + + def test_step_input_maps_roles_and_content_types(self): + transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input( + [ + {"type": "user_input", "content": [{"type": "text", "text": "I like apples."}]}, + {"type": "model_output", "content": [{"type": "text", "text": "I like oranges."}]}, + {"type": "user_input", "content": [{"type": "text", "text": "What did you say?"}]}, + ] + ) + assert transformed == [ + {"role": "user", "content": [{"type": "input_text", "text": "I like apples."}]}, + {"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]}, + {"role": "user", "content": [{"type": "input_text", "text": "What did you say?"}]}, + ] + + def test_legacy_turn_input_maps_model_role_to_assistant(self): + transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input( + [ + {"role": "user", "content": [{"type": "text", "text": "I like apples."}]}, + {"role": "model", "content": [{"type": "text", "text": "I like oranges."}]}, + ] + ) + assert transformed == [ + {"role": "user", "content": [{"type": "input_text", "text": "I like apples."}]}, + {"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]}, + ] + + def test_turn_pydantic_model_with_string_content(self): + transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input( + [Turn(role="model", content="I like oranges.")] + ) + assert transformed == [ + {"role": "assistant", "content": [{"type": "output_text", "text": "I like oranges."}]} + ] + + def test_string_input_passes_through(self): + transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input("Hello") + assert transformed == "Hello" + + def test_content_list_input_becomes_single_user_message(self): + transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input( + [{"type": "text", "text": "Hello"}, "world"] + ) + assert transformed == [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "Hello"}, + {"type": "input_text", "text": "world"}, + ], + } + ] + + def test_non_text_content_passes_through_unchanged(self): + image_part = {"type": "image", "data": "base64data", "mime_type": "image/png"} + transformed = LiteLLMResponsesInteractionsConfig._transform_interactions_input_to_responses_input( + [{"type": "user_input", "content": [image_part]}] + ) + assert transformed == [{"role": "user", "content": [image_part]}] diff --git a/tests/test_litellm/interactions/test_openapi_compliance.py b/tests/unit/interactions/test_openapi_compliance.py similarity index 99% rename from tests/test_litellm/interactions/test_openapi_compliance.py rename to tests/unit/interactions/test_openapi_compliance.py index 2665f8703a6..d3f1183cea6 100644 --- a/tests/test_litellm/interactions/test_openapi_compliance.py +++ b/tests/unit/interactions/test_openapi_compliance.py @@ -4,7 +4,7 @@ OpenAPI compliance tests for Google Interactions API. Validates that our SDK requests/responses match the OpenAPI spec at: https://ai.google.dev/static/api/interactions.openapi.json -Run with: pytest tests/test_litellm/interactions/test_openapi_compliance.py -v +Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v """ import json diff --git a/tests/test_litellm/files/__init__.py b/tests/unit/llms/aiml/__init__.py similarity index 100% rename from tests/test_litellm/files/__init__.py rename to tests/unit/llms/aiml/__init__.py diff --git a/tests/test_litellm/llms/anthropic/__init__.py b/tests/unit/llms/aiml/image_generation/__init__.py similarity index 100% rename from tests/test_litellm/llms/anthropic/__init__.py rename to tests/unit/llms/aiml/image_generation/__init__.py diff --git a/tests/test_litellm/llms/aiml/image_generation/test_aiml_image_generation_transformation.py b/tests/unit/llms/aiml/image_generation/test_aiml_image_generation_transformation.py similarity index 100% rename from tests/test_litellm/llms/aiml/image_generation/test_aiml_image_generation_transformation.py rename to tests/unit/llms/aiml/image_generation/test_aiml_image_generation_transformation.py diff --git a/tests/unit/llms/anthropic/batches/test_transformation.py b/tests/unit/llms/anthropic/batches/test_transformation.py index eacd2c9d03b..419fc7740eb 100644 --- a/tests/unit/llms/anthropic/batches/test_transformation.py +++ b/tests/unit/llms/anthropic/batches/test_transformation.py @@ -616,7 +616,7 @@ def test_transform_response_reraises_unexpected_error(config): # automatically. See base_batches_config_test.py. # --------------------------------------------------------------------------- # -from tests.test_litellm.llms.base_llm.batches.base_batches_config_test import ( # noqa: E402 +from tests.unit.llms.base_llm.batches.base_batches_config_test import ( # noqa: E402 BatchesConfigContractTests, ) diff --git a/tests/test_litellm/llms/anthropic/batches/__init__.py b/tests/unit/llms/anthropic/chat/__init__.py similarity index 100% rename from tests/test_litellm/llms/anthropic/batches/__init__.py rename to tests/unit/llms/anthropic/chat/__init__.py diff --git a/tests/test_litellm/llms/anthropic/chat/conftest.py b/tests/unit/llms/anthropic/chat/conftest.py similarity index 100% rename from tests/test_litellm/llms/anthropic/chat/conftest.py rename to tests/unit/llms/anthropic/chat/conftest.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/__init__.py b/tests/unit/llms/anthropic/chat/guardrail_translation/__init__.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/__init__.py rename to tests/unit/llms/anthropic/chat/guardrail_translation/__init__.py diff --git a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py similarity index 100% rename from tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py rename to tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py b/tests/unit/llms/anthropic/chat/test_anthropic_chat_handler.py similarity index 100% rename from tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py rename to tests/unit/llms/anthropic/chat/test_anthropic_chat_handler.py diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/unit/llms/anthropic/chat/test_anthropic_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py rename to tests/unit/llms/anthropic/chat/test_anthropic_chat_transformation.py diff --git a/tests/test_litellm/llms/anthropic/chat/test_code_interpreter_results_extraction.py b/tests/unit/llms/anthropic/chat/test_code_interpreter_results_extraction.py similarity index 100% rename from tests/test_litellm/llms/anthropic/chat/test_code_interpreter_results_extraction.py rename to tests/unit/llms/anthropic/chat/test_code_interpreter_results_extraction.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/__init__.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/__init__.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/__init__.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/__init__.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_prompt_cache_key.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_handler_prompt_cache_key.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_prompt_cache_key.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_handler_prompt_cache_key.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_reasoning_effort_normalization.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_handler_reasoning_effort_normalization.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_reasoning_effort_normalization.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_handler_reasoning_effort_normalization.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_combined_chunk.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_compaction.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_compaction.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_compaction.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_compaction.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_empty_choices.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_empty_choices.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_empty_choices.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_empty_choices.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_message_id.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_message_id.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_message_id.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_message_id.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_mid_stream_error.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_mid_stream_error.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_mid_stream_error.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_mid_stream_error.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_stop_reason.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_stop_reason.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_stop_reason.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_stop_reason.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_tool_args.py b/tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_tool_args.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_tool_args.py rename to tests/unit/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_tool_args.py diff --git a/tests/test_litellm/llms/anthropic/files/__init__.py b/tests/unit/llms/anthropic/experimental_pass_through/context_management/__init__.py similarity index 100% rename from tests/test_litellm/llms/anthropic/files/__init__.py rename to tests/unit/llms/anthropic/experimental_pass_through/context_management/__init__.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/test_clear_tool_uses.py b/tests/unit/llms/anthropic/experimental_pass_through/context_management/test_clear_tool_uses.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/test_clear_tool_uses.py rename to tests/unit/llms/anthropic/experimental_pass_through/context_management/test_clear_tool_uses.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/test_compact.py b/tests/unit/llms/anthropic/experimental_pass_through/context_management/test_compact.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/test_compact.py rename to tests/unit/llms/anthropic/experimental_pass_through/context_management/test_compact.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/test_dispatcher.py b/tests/unit/llms/anthropic/experimental_pass_through/context_management/test_dispatcher.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/test_dispatcher.py rename to tests/unit/llms/anthropic/experimental_pass_through/context_management/test_dispatcher.py diff --git a/tests/test_litellm/llms/azure/batches/__init__.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/__init__.py similarity index 100% rename from tests/test_litellm/llms/azure/batches/__init__.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/__init__.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_advisor_integration.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_advisor_integration.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_advisor_integration.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_advisor_integration.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_effort.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_effort.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_effort.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_effort.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_encrypted_reasoning.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_encrypted_reasoning.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_encrypted_reasoning.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_encrypted_reasoning.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_per_turn_control.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_per_turn_control.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_per_turn_control.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_per_turn_control.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_speed.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_speed.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_speed.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_speed.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_structured_outputs.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_structured_outputs.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_structured_outputs.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_structured_outputs.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_content_after_stop_reason.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_content_after_stop_reason.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_content_after_stop_reason.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_content_after_stop_reason.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mid_conversation_system.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_mid_conversation_system.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mid_conversation_system.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_mid_conversation_system.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_auto_summary_messages.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_reasoning_auto_summary_messages.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_auto_summary_messages.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_reasoning_auto_summary_messages.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_request_optional_param_utils.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_request_optional_param_utils.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_request_optional_param_utils.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_request_optional_param_utils.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_response_cache.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_response_cache.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_response_cache.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_response_cache.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_sse_wrapper.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_sse_wrapper.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_sse_wrapper.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_sse_wrapper.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py b/tests/unit/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py rename to tests/unit/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py diff --git a/tests/test_litellm/llms/azure/vector_stores/__init__.py b/tests/unit/llms/anthropic/experimental_pass_through/responses_adapters/__init__.py similarity index 100% rename from tests/test_litellm/llms/azure/vector_stores/__init__.py rename to tests/unit/llms/anthropic/experimental_pass_through/responses_adapters/__init__.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_handler.py b/tests/unit/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_handler.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_handler.py rename to tests/unit/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_handler.py diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_streaming_iterator.py b/tests/unit/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_streaming_iterator.py similarity index 72% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_streaming_iterator.py rename to tests/unit/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_streaming_iterator.py index bfe2d6b7cea..392ecc2bcdd 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_streaming_iterator.py +++ b/tests/unit/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_streaming_iterator.py @@ -4,18 +4,25 @@ Tests for AnthropicResponsesStreamWrapper """ import asyncio +import json import os import sys from types import SimpleNamespace +import pytest + sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../../.."))) +import litellm +from litellm.exceptions import MidStreamFallbackError from litellm.litellm_core_utils.prompt_templates.common_utils import ( encrypted_reasoning_signature, ) +from litellm.llms.anthropic.experimental_pass_through.messages.utils import INCOMPLETE_STREAM_ERROR_MESSAGE from litellm.llms.anthropic.experimental_pass_through.responses_adapters.streaming_iterator import ( AnthropicResponsesStreamWrapper, ) +from litellm.types.llms.openai import ResponseFailedEvent, ResponsesAPIResponse def _process_all(events: list) -> list: @@ -132,6 +139,7 @@ class TestReasoningItemWithoutSummaryText: {"type": "response.output_item.added", "item": {"type": "message", "id": "msg_1"}}, {"type": "response.output_text.delta", "item_id": "msg_1", "delta": "Hello"}, {"type": "response.output_item.done", "item": {"type": "message", "id": "msg_1"}}, + {"type": "response.completed"}, ] def test_reasoning_without_summary_emits_no_thinking_block(self): @@ -144,6 +152,8 @@ class TestReasoningItemWithoutSummaryText: ("content_block_start", 0), ("content_block_delta", 0), ("content_block_stop", 0), + ("message_delta", None), + ("message_stop", None), ] assert chunks[1]["content_block"] == {"type": "text", "text": ""} @@ -166,6 +176,8 @@ class TestReasoningItemWithoutSummaryText: ("content_block_start", 1), ("content_block_delta", 1), ("content_block_stop", 1), + ("message_delta", None), + ("message_stop", None), ] assert chunks[1]["content_block"] == {"type": "thinking", "thinking": "", "signature": ""} assert "".join(c["delta"]["thinking"] for c in chunks[2:4]) == "Weighing options" @@ -215,6 +227,8 @@ class TestEncryptedReasoningIsStreamedForReplay: ("content_block_start", 1), ("content_block_delta", 1), ("content_block_stop", 1), + ("message_delta", None), + ("message_stop", None), ] assert chunks[1]["content_block"] == { "type": "redacted_thinking", @@ -234,9 +248,7 @@ class TestEncryptedReasoningIsStreamedForReplay: ] chunks = _process_all(events) - thinking = "".join( - c["delta"]["thinking"] for c in chunks if c.get("delta", {}).get("type") == "thinking_delta" - ) + thinking = "".join(c["delta"]["thinking"] for c in chunks if c.get("delta", {}).get("type") == "thinking_delta") assert thinking == "First.\n\nSecond." assert [c["type"] for c in chunks].count("content_block_start") == 1 @@ -283,6 +295,7 @@ class TestToolUseBlockClosedExactlyOnce: "type": "response.output_item.done", "item": {"type": "message", "id": "chatcmpl-123", "status": "completed"}, }, + {"type": "response.completed"}, ] def test_one_content_block_stop_per_content_block_start(self): @@ -302,6 +315,8 @@ class TestToolUseBlockClosedExactlyOnce: ("content_block_delta", 0), ("content_block_delta", 0), ("content_block_stop", 0), + ("message_delta", None), + ("message_stop", None), ] assert chunks[1]["content_block"] == { "type": "tool_use", @@ -452,3 +467,158 @@ class TestRefusalStreamEvents: message_delta = next(c for c in chunks if c["type"] == "message_delta") assert message_delta["delta"]["stop_reason"] == "max_tokens" assert "stop_details" not in message_delta["delta"] + + +def _collect(stream) -> list: + async def _run() -> list: + wrapper = AnthropicResponsesStreamWrapper(responses_stream=stream, model="m") + return [chunk async for chunk in wrapper] + + return asyncio.run(_run()) + + +class TestUpstreamFailureEndsStreamWithErrorEvent: + """A provider failure must reach the Anthropic client as an ``error`` event that + ends the stream, never as a fabricated ``end_turn`` or a silent close.""" + + def test_response_failed_event_emits_error_event_and_stops_pulling_upstream(self): + failed = SimpleNamespace( + status="failed", + output=[], + usage=None, + error={"code": "rate_limit_exceeded", "message": "Rate limit reached for gpt-5.5, try again in 20s."}, + ) + + async def _gen(): + yield {"type": "response.created"} + yield {"type": "response.failed", "response": failed} + raise AssertionError("upstream was pulled again after the failure") + + async def _run() -> list: + wrapper = AnthropicResponsesStreamWrapper(responses_stream=_gen(), model="m") + return [frame async for frame in wrapper.async_anthropic_sse_wrapper()] + + frames = asyncio.run(_run()) + assert [frame.split(b"\n", 1)[0] for frame in frames] == [b"event: message_start", b"event: error"] + error_payload = json.loads(frames[1].split(b"data: ", 1)[1]) + assert error_payload["type"] == "error" + assert error_payload["error"] == { + "type": "rate_limit_error", + "message": "Rate limit reached for gpt-5.5, try again in 20s.", + } + + def test_raised_mid_stream_fallback_error_is_unwrapped_to_the_provider_failure(self): + rate_limit = litellm.RateLimitError(message="You have no credits remaining.", llm_provider="openai", model="m") + wrapped = MidStreamFallbackError( + message=str(rate_limit), + model="m", + llm_provider="openai", + original_exception=rate_limit, + is_pre_first_chunk=True, + ) + + async def _gen(): + yield {"type": "response.created"} + raise wrapped + + chunks = _collect(_gen()) + assert [chunk["type"] for chunk in chunks] == ["message_start", "error"] + assert chunks[1]["error"] == {"type": "rate_limit_error", "message": rate_limit.message} + + def test_sync_upstream_transport_error_after_content_becomes_api_error_event(self): + def _events(): + yield {"type": "response.created"} + yield {"type": "response.output_item.added", "item": {"type": "message", "id": "msg_1"}} + yield {"type": "response.output_text.delta", "item_id": "msg_1", "delta": "Hi"} + raise ConnectionResetError("Response payload is not completed") + + chunks = _collect(_events()) + assert [chunk["type"] for chunk in chunks] == [ + "message_start", + "content_block_start", + "content_block_delta", + "error", + ] + assert chunks[-1]["error"] == {"type": "api_error", "message": "Response payload is not completed"} + + def test_error_event_message_is_redacted_before_it_reaches_the_client(self): + async def _gen(): + yield {"type": "response.created"} + raise RuntimeError("upstream failed with key sk-proj-abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGHIJ") + + chunks = _collect(_gen()) + assert chunks[-1]["type"] == "error" + assert "sk-proj-" not in chunks[-1]["error"]["message"] + assert chunks[-1]["error"]["message"].startswith("upstream failed with key") + + @pytest.mark.parametrize( + ("raised", "expected_error"), + [ + ( + MidStreamFallbackError(message="boom", model="m", llm_provider="openai"), + {"type": "api_error", "message": "litellm.MidStreamFallbackError: boom"}, + ), + ( + type("StringStatusError", (Exception,), {"status_code": "429"})("throttled"), + {"type": "rate_limit_error", "message": "throttled"}, + ), + ( + type("NonErrorStatusError", (Exception,), {"status_code": 200})("odd status"), + {"type": "api_error", "message": "odd status"}, + ), + ], + ids=["mid-stream-fallback-without-original", "digit-string-status", "status-outside-4xx-5xx"], + ) + def test_raised_failure_status_is_normalized_into_the_error_type(self, raised, expected_error): + async def _gen(): + yield {"type": "response.created"} + raise raised + + chunks = _collect(_gen()) + assert [chunk["type"] for chunk in chunks] == ["message_start", "error"] + assert chunks[1]["error"] == expected_error + + def test_pydantic_response_failed_event_is_mapped_like_a_dict_event(self): + failed = ResponsesAPIResponse( + id="resp_1", + created_at=1, + error={"code": "server_error", "message": "The server had an error while processing your request."}, + status="failed", + output=[], + model="m", + object="response", + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ) + + async def _gen(): + yield {"type": "response.created"} + yield ResponseFailedEvent(type="response.failed", response=failed) + + chunks = _collect(_gen()) + assert [chunk["type"] for chunk in chunks] == ["message_start", "error"] + assert chunks[1]["error"] == { + "type": "api_error", + "message": "The server had an error while processing your request.", + } + + def test_upstream_ending_without_a_terminal_event_is_an_error_not_a_silent_close(self): + async def _gen(): + yield {"type": "response.created"} + yield {"type": "response.output_item.added", "item": {"type": "message", "id": "msg_1"}} + yield {"type": "response.output_text.delta", "item_id": "msg_1", "delta": "Hi"} + + chunks = _collect(_gen()) + assert [chunk["type"] for chunk in chunks] == [ + "message_start", + "content_block_start", + "content_block_delta", + "error", + ] + assert chunks[-1]["error"] == {"type": "api_error", "message": INCOMPLETE_STREAM_ERROR_MESSAGE} + + def test_sync_upstream_ending_before_any_event_is_an_error_not_a_silent_close(self): + chunks = _collect(iter(())) + assert [chunk["type"] for chunk in chunks] == ["message_start", "error"] + assert chunks[1]["error"] == {"type": "api_error", "message": INCOMPLETE_STREAM_ERROR_MESSAGE} diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py b/tests/unit/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py similarity index 100% rename from tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py rename to tests/unit/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/unit/llms/anthropic/test_anthropic_common_utils.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py rename to tests/unit/llms/anthropic/test_anthropic_common_utils.py diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_count_tokens_transformation.py b/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_anthropic_count_tokens_transformation.py rename to tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_files_and_batches.py b/tests/unit/llms/anthropic/test_anthropic_files_and_batches.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_anthropic_files_and_batches.py rename to tests/unit/llms/anthropic/test_anthropic_files_and_batches.py diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_output_format_filter.py b/tests/unit/llms/anthropic/test_anthropic_output_format_filter.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_anthropic_output_format_filter.py rename to tests/unit/llms/anthropic/test_anthropic_output_format_filter.py diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_prompt_cache_prediction.py b/tests/unit/llms/anthropic/test_anthropic_prompt_cache_prediction.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_anthropic_prompt_cache_prediction.py rename to tests/unit/llms/anthropic/test_anthropic_prompt_cache_prediction.py diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_reasoning_effort.py b/tests/unit/llms/anthropic/test_anthropic_reasoning_effort.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_anthropic_reasoning_effort.py rename to tests/unit/llms/anthropic/test_anthropic_reasoning_effort.py diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_schema_filter.py b/tests/unit/llms/anthropic/test_anthropic_schema_filter.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_anthropic_schema_filter.py rename to tests/unit/llms/anthropic/test_anthropic_schema_filter.py diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_structured_output.py b/tests/unit/llms/anthropic/test_anthropic_structured_output.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_anthropic_structured_output.py rename to tests/unit/llms/anthropic/test_anthropic_structured_output.py diff --git a/tests/test_litellm/llms/anthropic/test_azure_ai_cache_pricing.py b/tests/unit/llms/anthropic/test_azure_ai_cache_pricing.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_azure_ai_cache_pricing.py rename to tests/unit/llms/anthropic/test_azure_ai_cache_pricing.py diff --git a/tests/test_litellm/llms/anthropic/test_cost_calculation_dict_safety.py b/tests/unit/llms/anthropic/test_cost_calculation_dict_safety.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_cost_calculation_dict_safety.py rename to tests/unit/llms/anthropic/test_cost_calculation_dict_safety.py diff --git a/tests/test_litellm/llms/anthropic/test_count_tokens_oauth.py b/tests/unit/llms/anthropic/test_count_tokens_oauth.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_count_tokens_oauth.py rename to tests/unit/llms/anthropic/test_count_tokens_oauth.py diff --git a/tests/test_litellm/llms/anthropic/test_message_sanitization.py b/tests/unit/llms/anthropic/test_message_sanitization.py similarity index 100% rename from tests/test_litellm/llms/anthropic/test_message_sanitization.py rename to tests/unit/llms/anthropic/test_message_sanitization.py diff --git a/tests/test_litellm/llms/base_llm/__init__.py b/tests/unit/llms/azure/batches/__init__.py similarity index 100% rename from tests/test_litellm/llms/base_llm/__init__.py rename to tests/unit/llms/azure/batches/__init__.py diff --git a/tests/test_litellm/llms/azure/batches/test_handler.py b/tests/unit/llms/azure/batches/test_handler.py similarity index 100% rename from tests/test_litellm/llms/azure/batches/test_handler.py rename to tests/unit/llms/azure/batches/test_handler.py diff --git a/tests/test_litellm/llms/base_llm/batches/__init__.py b/tests/unit/llms/azure/chat/__init__.py similarity index 100% rename from tests/test_litellm/llms/base_llm/batches/__init__.py rename to tests/unit/llms/azure/chat/__init__.py diff --git a/tests/test_litellm/llms/azure/chat/test_azure_base_model_routing.py b/tests/unit/llms/azure/chat/test_azure_base_model_routing.py similarity index 100% rename from tests/test_litellm/llms/azure/chat/test_azure_base_model_routing.py rename to tests/unit/llms/azure/chat/test_azure_base_model_routing.py diff --git a/tests/test_litellm/llms/azure/chat/test_azure_chat_gpt_transformation.py b/tests/unit/llms/azure/chat/test_azure_chat_gpt_transformation.py similarity index 100% rename from tests/test_litellm/llms/azure/chat/test_azure_chat_gpt_transformation.py rename to tests/unit/llms/azure/chat/test_azure_chat_gpt_transformation.py diff --git a/tests/test_litellm/llms/azure/chat/test_azure_chat_o_series_transformation.py b/tests/unit/llms/azure/chat/test_azure_chat_o_series_transformation.py similarity index 100% rename from tests/test_litellm/llms/azure/chat/test_azure_chat_o_series_transformation.py rename to tests/unit/llms/azure/chat/test_azure_chat_o_series_transformation.py diff --git a/tests/test_litellm/llms/azure/chat/test_azure_gpt5_transformation.py b/tests/unit/llms/azure/chat/test_azure_gpt5_transformation.py similarity index 100% rename from tests/test_litellm/llms/azure/chat/test_azure_gpt5_transformation.py rename to tests/unit/llms/azure/chat/test_azure_gpt5_transformation.py diff --git a/tests/test_litellm/llms/azure/realtime/test_handler.py b/tests/unit/llms/azure/realtime/test_handler.py similarity index 100% rename from tests/test_litellm/llms/azure/realtime/test_handler.py rename to tests/unit/llms/azure/realtime/test_handler.py diff --git a/tests/test_litellm/llms/azure/test_audio_transcriptions.py b/tests/unit/llms/azure/test_audio_transcriptions.py similarity index 100% rename from tests/test_litellm/llms/azure/test_audio_transcriptions.py rename to tests/unit/llms/azure/test_audio_transcriptions.py diff --git a/tests/test_litellm/llms/azure/test_azure.py b/tests/unit/llms/azure/test_azure.py similarity index 100% rename from tests/test_litellm/llms/azure/test_azure.py rename to tests/unit/llms/azure/test_azure.py diff --git a/tests/test_litellm/llms/azure/test_azure_common_utils.py b/tests/unit/llms/azure/test_azure_common_utils.py similarity index 100% rename from tests/test_litellm/llms/azure/test_azure_common_utils.py rename to tests/unit/llms/azure/test_azure_common_utils.py diff --git a/tests/test_litellm/llms/azure/test_azure_cost_calculation.py b/tests/unit/llms/azure/test_azure_cost_calculation.py similarity index 100% rename from tests/test_litellm/llms/azure/test_azure_cost_calculation.py rename to tests/unit/llms/azure/test_azure_cost_calculation.py diff --git a/tests/test_litellm/llms/azure/test_azure_embedding.py b/tests/unit/llms/azure/test_azure_embedding.py similarity index 100% rename from tests/test_litellm/llms/azure/test_azure_embedding.py rename to tests/unit/llms/azure/test_azure_embedding.py diff --git a/tests/test_litellm/llms/azure/test_azure_exception_mapping.py b/tests/unit/llms/azure/test_azure_exception_mapping.py similarity index 100% rename from tests/test_litellm/llms/azure/test_azure_exception_mapping.py rename to tests/unit/llms/azure/test_azure_exception_mapping.py diff --git a/tests/test_litellm/llms/azure/test_azure_fine_tuning_api.py b/tests/unit/llms/azure/test_azure_fine_tuning_api.py similarity index 100% rename from tests/test_litellm/llms/azure/test_azure_fine_tuning_api.py rename to tests/unit/llms/azure/test_azure_fine_tuning_api.py diff --git a/tests/test_litellm/llms/azure/test_azure_speech_audio_transcription.py b/tests/unit/llms/azure/test_azure_speech_audio_transcription.py similarity index 100% rename from tests/test_litellm/llms/azure/test_azure_speech_audio_transcription.py rename to tests/unit/llms/azure/test_azure_speech_audio_transcription.py diff --git a/tests/test_litellm/llms/base_llm/files/__init__.py b/tests/unit/llms/azure/videos/__init__.py similarity index 100% rename from tests/test_litellm/llms/base_llm/files/__init__.py rename to tests/unit/llms/azure/videos/__init__.py diff --git a/tests/test_litellm/llms/azure/videos/test_azure_video_transformation.py b/tests/unit/llms/azure/videos/test_azure_video_transformation.py similarity index 100% rename from tests/test_litellm/llms/azure/videos/test_azure_video_transformation.py rename to tests/unit/llms/azure/videos/test_azure_video_transformation.py diff --git a/tests/test_litellm/llms/base_llm/realtime/__init__.py b/tests/unit/llms/azure_ai/claude/__init__.py similarity index 100% rename from tests/test_litellm/llms/base_llm/realtime/__init__.py rename to tests/unit/llms/azure_ai/claude/__init__.py diff --git a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_count_tokens_transformation.py b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_count_tokens_transformation.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_count_tokens_transformation.py rename to tests/unit/llms/azure_ai/claude/test_azure_anthropic_count_tokens_transformation.py diff --git a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_handler.py b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_handler.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_handler.py rename to tests/unit/llms/azure_ai/claude/test_azure_anthropic_handler.py diff --git a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py rename to tests/unit/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py diff --git a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_provider_routing.py b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_provider_routing.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_provider_routing.py rename to tests/unit/llms/azure_ai/claude/test_azure_anthropic_provider_routing.py diff --git a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_transformation.py b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_transformation.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_transformation.py rename to tests/unit/llms/azure_ai/claude/test_azure_anthropic_transformation.py diff --git a/tests/test_litellm/llms/azure_ai/claude/test_main_azure_anthropic_timeout.py b/tests/unit/llms/azure_ai/claude/test_main_azure_anthropic_timeout.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/claude/test_main_azure_anthropic_timeout.py rename to tests/unit/llms/azure_ai/claude/test_main_azure_anthropic_timeout.py diff --git a/tests/test_litellm/llms/bedrock/__init__.py b/tests/unit/llms/azure_ai/image_generation/__init__.py similarity index 100% rename from tests/test_litellm/llms/bedrock/__init__.py rename to tests/unit/llms/azure_ai/image_generation/__init__.py diff --git a/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py b/tests/unit/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py rename to tests/unit/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py diff --git a/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py b/tests/unit/llms/azure_ai/image_generation/test_mai_image_generation.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py rename to tests/unit/llms/azure_ai/image_generation/test_mai_image_generation.py diff --git a/tests/test_litellm/llms/azure_ai/test_azure_ai_agents_handler.py b/tests/unit/llms/azure_ai/test_azure_ai_agents_handler.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/test_azure_ai_agents_handler.py rename to tests/unit/llms/azure_ai/test_azure_ai_agents_handler.py diff --git a/tests/test_litellm/llms/azure_ai/test_azure_ai_cost_calculator.py b/tests/unit/llms/azure_ai/test_azure_ai_cost_calculator.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/test_azure_ai_cost_calculator.py rename to tests/unit/llms/azure_ai/test_azure_ai_cost_calculator.py diff --git a/tests/test_litellm/llms/azure_ai/test_azure_ai_entra_auth.py b/tests/unit/llms/azure_ai/test_azure_ai_entra_auth.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/test_azure_ai_entra_auth.py rename to tests/unit/llms/azure_ai/test_azure_ai_entra_auth.py diff --git a/tests/test_litellm/llms/azure_ai/test_azure_ai_foundry_catalog_model_metadata.py b/tests/unit/llms/azure_ai/test_azure_ai_foundry_catalog_model_metadata.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/test_azure_ai_foundry_catalog_model_metadata.py rename to tests/unit/llms/azure_ai/test_azure_ai_foundry_catalog_model_metadata.py diff --git a/tests/test_litellm/llms/azure_ai/test_azure_ai_fw_models_metadata.py b/tests/unit/llms/azure_ai/test_azure_ai_fw_models_metadata.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/test_azure_ai_fw_models_metadata.py rename to tests/unit/llms/azure_ai/test_azure_ai_fw_models_metadata.py diff --git a/tests/test_litellm/llms/azure_ai/test_azure_ai_kimi_k26_metadata.py b/tests/unit/llms/azure_ai/test_azure_ai_kimi_k26_metadata.py similarity index 100% rename from tests/test_litellm/llms/azure_ai/test_azure_ai_kimi_k26_metadata.py rename to tests/unit/llms/azure_ai/test_azure_ai_kimi_k26_metadata.py diff --git a/tests/test_litellm/llms/base_llm/batches/base_batches_config_test.py b/tests/unit/llms/base_llm/batches/base_batches_config_test.py similarity index 100% rename from tests/test_litellm/llms/base_llm/batches/base_batches_config_test.py rename to tests/unit/llms/base_llm/batches/base_batches_config_test.py diff --git a/tests/test_litellm/llms/bedrock/batches/__init__.py b/tests/unit/llms/base_llm/files/__init__.py similarity index 100% rename from tests/test_litellm/llms/bedrock/batches/__init__.py rename to tests/unit/llms/base_llm/files/__init__.py diff --git a/tests/test_litellm/llms/base_llm/files/test_azure_blob_storage_backend.py b/tests/unit/llms/base_llm/files/test_azure_blob_storage_backend.py similarity index 100% rename from tests/test_litellm/llms/base_llm/files/test_azure_blob_storage_backend.py rename to tests/unit/llms/base_llm/files/test_azure_blob_storage_backend.py diff --git a/tests/test_litellm/llms/base_llm/files/test_litellm_db_storage_backend.py b/tests/unit/llms/base_llm/files/test_litellm_db_storage_backend.py similarity index 100% rename from tests/test_litellm/llms/base_llm/files/test_litellm_db_storage_backend.py rename to tests/unit/llms/base_llm/files/test_litellm_db_storage_backend.py diff --git a/tests/test_litellm/llms/base_llm/files/test_storage_backend_factory.py b/tests/unit/llms/base_llm/files/test_storage_backend_factory.py similarity index 100% rename from tests/test_litellm/llms/base_llm/files/test_storage_backend_factory.py rename to tests/unit/llms/base_llm/files/test_storage_backend_factory.py diff --git a/tests/test_litellm/llms/bedrock/chat/agentcore/__init__.py b/tests/unit/llms/base_llm/responses/__init__.py similarity index 100% rename from tests/test_litellm/llms/bedrock/chat/agentcore/__init__.py rename to tests/unit/llms/base_llm/responses/__init__.py diff --git a/tests/test_litellm/llms/base_llm/responses/test_codex_compat.py b/tests/unit/llms/base_llm/responses/test_codex_compat.py similarity index 100% rename from tests/test_litellm/llms/base_llm/responses/test_codex_compat.py rename to tests/unit/llms/base_llm/responses/test_codex_compat.py diff --git a/tests/test_litellm/llms/base_llm/responses/test_transformation.py b/tests/unit/llms/base_llm/responses/test_transformation.py similarity index 100% rename from tests/test_litellm/llms/base_llm/responses/test_transformation.py rename to tests/unit/llms/base_llm/responses/test_transformation.py diff --git a/tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/__init__.py b/tests/unit/llms/base_llm/search/__init__.py similarity index 100% rename from tests/test_litellm/llms/bedrock/passthrough/guardrail_translation/__init__.py rename to tests/unit/llms/base_llm/search/__init__.py diff --git a/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py b/tests/unit/llms/base_llm/search/test_base_search_transformation.py similarity index 100% rename from tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py rename to tests/unit/llms/base_llm/search/test_base_search_transformation.py diff --git a/tests/test_litellm/llms/base_llm/test_base_managed_resource.py b/tests/unit/llms/base_llm/test_base_managed_resource.py similarity index 100% rename from tests/test_litellm/llms/base_llm/test_base_managed_resource.py rename to tests/unit/llms/base_llm/test_base_managed_resource.py diff --git a/tests/test_litellm/llms/base_llm/test_base_model_iterator.py b/tests/unit/llms/base_llm/test_base_model_iterator.py similarity index 100% rename from tests/test_litellm/llms/base_llm/test_base_model_iterator.py rename to tests/unit/llms/base_llm/test_base_model_iterator.py diff --git a/tests/test_litellm/llms/base_llm/test_managed_resource_isolation.py b/tests/unit/llms/base_llm/test_managed_resource_isolation.py similarity index 100% rename from tests/test_litellm/llms/base_llm/test_managed_resource_isolation.py rename to tests/unit/llms/base_llm/test_managed_resource_isolation.py diff --git a/tests/test_litellm/llms/base_llm/test_managed_resources_utils.py b/tests/unit/llms/base_llm/test_managed_resources_utils.py similarity index 100% rename from tests/test_litellm/llms/base_llm/test_managed_resources_utils.py rename to tests/unit/llms/base_llm/test_managed_resources_utils.py diff --git a/tests/test_litellm/llms/black_forest_labs/__init__.py b/tests/unit/llms/bedrock/batches/__init__.py similarity index 100% rename from tests/test_litellm/llms/black_forest_labs/__init__.py rename to tests/unit/llms/bedrock/batches/__init__.py diff --git a/tests/test_litellm/llms/bedrock/batches/test_batch_metadata_sanitization.py b/tests/unit/llms/bedrock/batches/test_batch_metadata_sanitization.py similarity index 100% rename from tests/test_litellm/llms/bedrock/batches/test_batch_metadata_sanitization.py rename to tests/unit/llms/bedrock/batches/test_batch_metadata_sanitization.py diff --git a/tests/test_litellm/llms/bedrock/batches/test_handler.py b/tests/unit/llms/bedrock/batches/test_handler.py similarity index 100% rename from tests/test_litellm/llms/bedrock/batches/test_handler.py rename to tests/unit/llms/bedrock/batches/test_handler.py diff --git a/tests/test_litellm/llms/bedrock/batches/test_transformation.py b/tests/unit/llms/bedrock/batches/test_transformation.py similarity index 99% rename from tests/test_litellm/llms/bedrock/batches/test_transformation.py rename to tests/unit/llms/bedrock/batches/test_transformation.py index 347c459a369..5e987239c54 100644 --- a/tests/test_litellm/llms/bedrock/batches/test_transformation.py +++ b/tests/unit/llms/bedrock/batches/test_transformation.py @@ -878,7 +878,7 @@ def test_validate_environment_passes_headers_through(config): # Shared BaseBatchesConfig contract suite. # --------------------------------------------------------------------------- # -from tests.test_litellm.llms.base_llm.batches.base_batches_config_test import ( # noqa: E402 +from tests.unit.llms.base_llm.batches.base_batches_config_test import ( # noqa: E402 BatchesConfigContractTests, ) diff --git a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py b/tests/unit/llms/bedrock/chat/test_bedrock_converse_handler.py similarity index 99% rename from tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py rename to tests/unit/llms/bedrock/chat/test_bedrock_converse_handler.py index 67ffe7570a1..08bcac33a35 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py +++ b/tests/unit/llms/bedrock/chat/test_bedrock_converse_handler.py @@ -20,7 +20,7 @@ from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.rust_bridge import configuration from litellm.types.utils import ModelResponse -from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe +from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe RESOLVED_CREDENTIALS = Credentials( access_key="AKIARESOLVED", diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/unit/llms/bedrock/chat/test_converse_transformation.py similarity index 100% rename from tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py rename to tests/unit/llms/bedrock/chat/test_converse_transformation.py diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation_nova_2.py b/tests/unit/llms/bedrock/chat/test_converse_transformation_nova_2.py similarity index 100% rename from tests/test_litellm/llms/bedrock/chat/test_converse_transformation_nova_2.py rename to tests/unit/llms/bedrock/chat/test_converse_transformation_nova_2.py diff --git a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py b/tests/unit/llms/bedrock/chat/test_invoke_handler.py similarity index 100% rename from tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py rename to tests/unit/llms/bedrock/chat/test_invoke_handler.py diff --git a/tests/test_litellm/llms/bedrock/chat/test_mistral_config.py b/tests/unit/llms/bedrock/chat/test_mistral_config.py similarity index 100% rename from tests/test_litellm/llms/bedrock/chat/test_mistral_config.py rename to tests/unit/llms/bedrock/chat/test_mistral_config.py diff --git a/tests/test_litellm/llms/bedrock/chat/test_service_tier.py b/tests/unit/llms/bedrock/chat/test_service_tier.py similarity index 100% rename from tests/test_litellm/llms/bedrock/chat/test_service_tier.py rename to tests/unit/llms/bedrock/chat/test_service_tier.py diff --git a/tests/test_litellm/llms/bedrock/chat/test_streaming_choice_index.py b/tests/unit/llms/bedrock/chat/test_streaming_choice_index.py similarity index 100% rename from tests/test_litellm/llms/bedrock/chat/test_streaming_choice_index.py rename to tests/unit/llms/bedrock/chat/test_streaming_choice_index.py diff --git a/tests/test_litellm/llms/bedrock/chat/test_writer_palmyra.py b/tests/unit/llms/bedrock/chat/test_writer_palmyra.py similarity index 100% rename from tests/test_litellm/llms/bedrock/chat/test_writer_palmyra.py rename to tests/unit/llms/bedrock/chat/test_writer_palmyra.py diff --git a/tests/unit/llms/bedrock/count_tokens/test_bedrock_count_tokens_handler.py b/tests/unit/llms/bedrock/count_tokens/test_bedrock_count_tokens_handler.py index 3622ce7f212..d67724f261d 100644 --- a/tests/unit/llms/bedrock/count_tokens/test_bedrock_count_tokens_handler.py +++ b/tests/unit/llms/bedrock/count_tokens/test_bedrock_count_tokens_handler.py @@ -7,7 +7,7 @@ from botocore.credentials import RefreshableCredentials from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler -from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe +from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe class _ProbedCountTokensHandler(BedrockCountTokensHandler): diff --git a/tests/test_litellm/llms/black_forest_labs/image_edit/__init__.py b/tests/unit/llms/bedrock/embed/__init__.py similarity index 100% rename from tests/test_litellm/llms/black_forest_labs/image_edit/__init__.py rename to tests/unit/llms/bedrock/embed/__init__.py diff --git a/tests/test_litellm/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py b/tests/unit/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py similarity index 99% rename from tests/test_litellm/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py rename to tests/unit/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py index 18f4b0f6ced..fbcbd0aaea6 100644 --- a/tests/test_litellm/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py +++ b/tests/unit/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py @@ -9,7 +9,7 @@ import respx import litellm from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.types.llms.base import HiddenParams -from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe +from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe # Mock async invoke responses async_invoke_response = { diff --git a/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py b/tests/unit/llms/bedrock/embed/test_bedrock_embedding.py similarity index 99% rename from tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py rename to tests/unit/llms/bedrock/embed/test_bedrock_embedding.py index e5a460e2f1a..ad21cadaa4b 100644 --- a/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py +++ b/tests/unit/llms/bedrock/embed/test_bedrock_embedding.py @@ -11,7 +11,7 @@ import litellm from litellm.llms.bedrock.embed.twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.bedrock.embed.embedding import BedrockEmbedding -from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe +from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe # Mock responses for different embedding models titan_embedding_response = {"embedding": [0.1, 0.2, 0.3], "inputTextTokenCount": 10} diff --git a/tests/test_litellm/llms/bedrock/embed/test_embedding.py b/tests/unit/llms/bedrock/embed/test_embedding.py similarity index 100% rename from tests/test_litellm/llms/bedrock/embed/test_embedding.py rename to tests/unit/llms/bedrock/embed/test_embedding.py diff --git a/tests/test_litellm/llms/bedrock/embed/test_twelvelabs_marengo_3_transformation.py b/tests/unit/llms/bedrock/embed/test_twelvelabs_marengo_3_transformation.py similarity index 100% rename from tests/test_litellm/llms/bedrock/embed/test_twelvelabs_marengo_3_transformation.py rename to tests/unit/llms/bedrock/embed/test_twelvelabs_marengo_3_transformation.py diff --git a/tests/test_litellm/llms/bedrock/event_loop_probe.py b/tests/unit/llms/bedrock/event_loop_probe.py similarity index 100% rename from tests/test_litellm/llms/bedrock/event_loop_probe.py rename to tests/unit/llms/bedrock/event_loop_probe.py diff --git a/tests/test_litellm/llms/black_forest_labs/image_generation/__init__.py b/tests/unit/llms/bedrock/messages/__init__.py similarity index 100% rename from tests/test_litellm/llms/black_forest_labs/image_generation/__init__.py rename to tests/unit/llms/bedrock/messages/__init__.py diff --git a/tests/test_litellm/llms/cerebras/__init__.py b/tests/unit/llms/bedrock/messages/invoke_transformations/__init__.py similarity index 100% rename from tests/test_litellm/llms/cerebras/__init__.py rename to tests/unit/llms/bedrock/messages/invoke_transformations/__init__.py diff --git a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/unit/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py similarity index 100% rename from tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py rename to tests/unit/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py diff --git a/tests/test_litellm/llms/bedrock/rerank/transformation.py b/tests/unit/llms/bedrock/rerank/transformation.py similarity index 100% rename from tests/test_litellm/llms/bedrock/rerank/transformation.py rename to tests/unit/llms/bedrock/rerank/transformation.py diff --git a/tests/test_litellm/llms/chatgpt/__init__.py b/tests/unit/llms/bedrock/responses/__init__.py similarity index 100% rename from tests/test_litellm/llms/chatgpt/__init__.py rename to tests/unit/llms/bedrock/responses/__init__.py diff --git a/tests/test_litellm/llms/bedrock/responses/test_bedrock_openai_responses.py b/tests/unit/llms/bedrock/responses/test_bedrock_openai_responses.py similarity index 100% rename from tests/test_litellm/llms/bedrock/responses/test_bedrock_openai_responses.py rename to tests/unit/llms/bedrock/responses/test_bedrock_openai_responses.py diff --git a/tests/test_litellm/llms/chatgpt/chat/__init__.py b/tests/unit/llms/bedrock/search/__init__.py similarity index 100% rename from tests/test_litellm/llms/chatgpt/chat/__init__.py rename to tests/unit/llms/bedrock/search/__init__.py diff --git a/tests/test_litellm/llms/bedrock/search/test_agentcore_search_transformation.py b/tests/unit/llms/bedrock/search/test_agentcore_search_transformation.py similarity index 100% rename from tests/test_litellm/llms/bedrock/search/test_agentcore_search_transformation.py rename to tests/unit/llms/bedrock/search/test_agentcore_search_transformation.py diff --git a/tests/test_litellm/llms/bedrock/test_anthropic_beta_support.py b/tests/unit/llms/bedrock/test_anthropic_beta_support.py similarity index 100% rename from tests/test_litellm/llms/bedrock/test_anthropic_beta_support.py rename to tests/unit/llms/bedrock/test_anthropic_beta_support.py diff --git a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py b/tests/unit/llms/bedrock/test_base_aws_llm.py similarity index 99% rename from tests/test_litellm/llms/bedrock/test_base_aws_llm.py rename to tests/unit/llms/bedrock/test_base_aws_llm.py index 6b9450afed4..db144ab6d56 100644 --- a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py +++ b/tests/unit/llms/bedrock/test_base_aws_llm.py @@ -28,7 +28,7 @@ from litellm.llms.bedrock.base_aws_llm import ( run_aws_signing, sign_request_off_loop_if_aws, ) -from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe +from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe # Global variable for the base_aws_llm.py file path diff --git a/tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py b/tests/unit/llms/bedrock/test_bedrock_common_utils.py similarity index 100% rename from tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py rename to tests/unit/llms/bedrock/test_bedrock_common_utils.py diff --git a/tests/test_litellm/llms/bedrock/test_bedrock_ssl_verify.py b/tests/unit/llms/bedrock/test_bedrock_ssl_verify.py similarity index 100% rename from tests/test_litellm/llms/bedrock/test_bedrock_ssl_verify.py rename to tests/unit/llms/bedrock/test_bedrock_ssl_verify.py diff --git a/tests/test_litellm/llms/bedrock/test_claude_platform_provider.py b/tests/unit/llms/bedrock/test_claude_platform_provider.py similarity index 100% rename from tests/test_litellm/llms/bedrock/test_claude_platform_provider.py rename to tests/unit/llms/bedrock/test_claude_platform_provider.py diff --git a/tests/test_litellm/llms/bedrock/test_converse_context_management.py b/tests/unit/llms/bedrock/test_converse_context_management.py similarity index 100% rename from tests/test_litellm/llms/bedrock/test_converse_context_management.py rename to tests/unit/llms/bedrock/test_converse_context_management.py diff --git a/tests/test_litellm/llms/bedrock/test_cross_region_inference_profile_mapping.py b/tests/unit/llms/bedrock/test_cross_region_inference_profile_mapping.py similarity index 100% rename from tests/test_litellm/llms/bedrock/test_cross_region_inference_profile_mapping.py rename to tests/unit/llms/bedrock/test_cross_region_inference_profile_mapping.py diff --git a/tests/test_litellm/llms/bedrock/test_mantle.py b/tests/unit/llms/bedrock/test_mantle.py similarity index 100% rename from tests/test_litellm/llms/bedrock/test_mantle.py rename to tests/unit/llms/bedrock/test_mantle.py diff --git a/tests/test_litellm/llms/bedrock/test_nova_imported_models.py b/tests/unit/llms/bedrock/test_nova_imported_models.py similarity index 100% rename from tests/test_litellm/llms/bedrock/test_nova_imported_models.py rename to tests/unit/llms/bedrock/test_nova_imported_models.py diff --git a/tests/test_litellm/llms/bedrock/test_request_metadata.py b/tests/unit/llms/bedrock/test_request_metadata.py similarity index 100% rename from tests/test_litellm/llms/bedrock/test_request_metadata.py rename to tests/unit/llms/bedrock/test_request_metadata.py diff --git a/tests/test_litellm/llms/bedrock/test_web_identity_session_policy.py b/tests/unit/llms/bedrock/test_web_identity_session_policy.py similarity index 100% rename from tests/test_litellm/llms/bedrock/test_web_identity_session_policy.py rename to tests/unit/llms/bedrock/test_web_identity_session_policy.py diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_messages_transformation.py b/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_messages_transformation.py similarity index 100% rename from tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_messages_transformation.py rename to tests/unit/llms/bedrock_mantle/test_bedrock_mantle_messages_transformation.py diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py b/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py similarity index 100% rename from tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py rename to tests/unit/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_transformation.py b/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py similarity index 99% rename from tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_transformation.py rename to tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py index 0cc3963358f..4bf3dd11fa1 100644 --- a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_transformation.py +++ b/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py @@ -19,7 +19,7 @@ import litellm from litellm.llms.bedrock_mantle.chat.transformation import BedrockMantleChatConfig from litellm.llms.bedrock.base_aws_llm import sign_request_off_loop_if_aws from litellm.types.utils import LlmProviders -from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe +from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe @pytest.fixture diff --git a/tests/test_litellm/llms/crusoe/__init__.py b/tests/unit/llms/cometapi/__init__.py similarity index 100% rename from tests/test_litellm/llms/crusoe/__init__.py rename to tests/unit/llms/cometapi/__init__.py diff --git a/tests/test_litellm/llms/databricks/chat/__init__.py b/tests/unit/llms/cometapi/chat/__init__.py similarity index 100% rename from tests/test_litellm/llms/databricks/chat/__init__.py rename to tests/unit/llms/cometapi/chat/__init__.py diff --git a/tests/unit/llms/cometapi/chat/test_cometapi_chat_transformation.py b/tests/unit/llms/cometapi/chat/test_cometapi_chat_transformation.py new file mode 100644 index 00000000000..607648dd6c9 --- /dev/null +++ b/tests/unit/llms/cometapi/chat/test_cometapi_chat_transformation.py @@ -0,0 +1,183 @@ +""" +Unit tests for CometAPI Chat Configuration + +Tests the CometAPIChatConfig class methods using mocks +""" + + +import pytest + + +from litellm.llms.cometapi.chat.transformation import ( + CometAPIChatCompletionStreamingHandler, + CometAPIConfig, +) +from litellm.llms.cometapi.common_utils import CometAPIException + + +class TestCometAPIChatCompletionStreamingHandler: + def test_chunk_parser_successful(self): + handler = CometAPIChatCompletionStreamingHandler( + streaming_response=None, sync_stream=True + ) + + # Test input chunk + chunk = { + "id": "test_id", + "created": 1234567890, + "model": "gpt-3.5-turbo", + "usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}, + "choices": [ + {"delta": {"content": "test content", "reasoning": "test reasoning"}} + ], + } + + # Parse chunk + result = handler.chunk_parser(chunk) + + # Verify response + assert result.id == "test_id" + assert result.object == "chat.completion.chunk" + assert result.created == 1234567890 + assert result.model == "gpt-3.5-turbo" + assert result.usage.prompt_tokens == chunk["usage"]["prompt_tokens"] + assert result.usage.completion_tokens == chunk["usage"]["completion_tokens"] + assert result.usage.total_tokens == chunk["usage"]["total_tokens"] + assert len(result.choices) == 1 + assert result.choices[0]["delta"]["reasoning_content"] == "test reasoning" + + def test_chunk_parser_error_response(self): + handler = CometAPIChatCompletionStreamingHandler( + streaming_response=None, sync_stream=True + ) + + # Test error chunk + error_chunk = { + "error": { + "message": "test error", + "code": 400, + } + } + + # Verify error handling + with pytest.raises(CometAPIException) as exc_info: + handler.chunk_parser(error_chunk) + + assert "CometAPI Error: test error" in str(exc_info.value) + assert exc_info.value.status_code == 400 + + def test_chunk_parser_key_error(self): + handler = CometAPIChatCompletionStreamingHandler( + streaming_response=None, sync_stream=True + ) + + # Test invalid chunk missing required fields + invalid_chunk = {"incomplete": "data"} + + # Verify KeyError handling + with pytest.raises(CometAPIException) as exc_info: + handler.chunk_parser(invalid_chunk) + + assert "KeyError" in str(exc_info.value) + assert exc_info.value.status_code == 400 + + +class TestCometAPIConfig: + def test_transform_request_basic(self): + """Test basic request transformation""" + config = CometAPIConfig() + + transformed_request = config.transform_request( + model="cometapi/gpt-3.5-turbo", + messages=[{"role": "user", "content": "Hello, world!"}], + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert transformed_request["model"] == "cometapi/gpt-3.5-turbo" + assert transformed_request["messages"] == [ + {"role": "user", "content": "Hello, world!"} + ] + + def test_transform_request_with_extra_body(self): + """Test request transformation with extra_body parameters""" + config = CometAPIConfig() + + transformed_request = config.transform_request( + model="cometapi/gpt-4", + messages=[{"role": "user", "content": "Hello, world!"}], + optional_params={"extra_body": {"custom_param": "custom_value"}}, + litellm_params={}, + headers={}, + ) + + # Validate that extra_body parameters are merged into the request + assert transformed_request["custom_param"] == "custom_value" + assert transformed_request["messages"] == [ + {"role": "user", "content": "Hello, world!"} + ] + + def test_cache_control_flag_removal(self): + """Test cache control flag removal from messages""" + config = CometAPIConfig() + + transformed_request = config.transform_request( + model="cometapi/gpt-3.5-turbo", + messages=[ + { + "role": "user", + "content": "Hello, world!", + "cache_control": {"type": "ephemeral"}, + } + ], + optional_params={}, + litellm_params={}, + headers={}, + ) + + # CometAPI should remove cache_control flags by default + assert transformed_request["messages"][0].get("cache_control") is None + + def test_map_openai_params(self): + """Test OpenAI parameter mapping""" + config = CometAPIConfig() + + non_default_params = { + "temperature": 0.7, + "max_tokens": 100, + "top_p": 0.9, + } + + mapped_params = config.map_openai_params( + non_default_params=non_default_params, + optional_params={}, + model="cometapi/gpt-3.5-turbo", + drop_params=False, + ) + + assert mapped_params["temperature"] == 0.7 + assert mapped_params["max_tokens"] == 100 + assert mapped_params["top_p"] == 0.9 + + def test_get_error_class(self): + """Test error class creation""" + config = CometAPIConfig() + + error = config.get_error_class( + error_message="Test error", + status_code=400, + headers={"Content-Type": "application/json"}, + ) + + assert isinstance(error, CometAPIException) + assert error.message == "Test error" + assert error.status_code == 400 + + +# Integration test example (requires real API key) + + +if __name__ == "__main__": + # Quick test runner + pytest.main([__file__, "-v"]) diff --git a/tests/test_litellm/llms/databricks/responses/__init__.py b/tests/unit/llms/compactifai/__init__.py similarity index 100% rename from tests/test_litellm/llms/databricks/responses/__init__.py rename to tests/unit/llms/compactifai/__init__.py diff --git a/tests/test_litellm/llms/compactifai/test_compactifai.py b/tests/unit/llms/compactifai/test_compactifai.py similarity index 84% rename from tests/test_litellm/llms/compactifai/test_compactifai.py rename to tests/unit/llms/compactifai/test_compactifai.py index fd31049731a..1367c703fda 100644 --- a/tests/test_litellm/llms/compactifai/test_compactifai.py +++ b/tests/unit/llms/compactifai/test_compactifai.py @@ -104,56 +104,6 @@ def test_compactifai_completion_streaming(respx_mock): assert chunks[0].choices[0].delta.content == "Hello" -@pytest.mark.respx() -def test_compactifai_models_endpoint(respx_mock): - """Test CompactifAI models listing""" - litellm.disable_aiohttp_transport = True - - mock_response = { - "object": "list", - "data": [ - { - "id": "cai-llama-3-1-8b-slim", - "object": "model", - "created": 1677610602, - "owned_by": "compactifai", - }, - { - "id": "mistral-7b-compressed", - "object": "model", - "created": 1677610602, - "owned_by": "compactifai", - }, - ], - } - - respx_mock.post("https://api.compactif.ai/v1/chat/completions").respond( - json={ - "id": "chatcmpl-123", - "object": "chat.completion", - "created": 1677652288, - "model": "cai-llama-3-1-8b-slim", - "choices": [ - { - "index": 0, - "message": {"role": "assistant", "content": "Test response"}, - "finish_reason": "stop", - } - ], - "usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}, - }, - status_code=200, - ) - - # This would be tested if litellm had a models() function - # For now, we'll test that the provider is properly configured - response = litellm.completion( - model="compactifai/cai-llama-3-1-8b-slim", - messages=[{"role": "user", "content": "test"}], - api_key="test-key", - ) - - @pytest.mark.respx() def test_compactifai_authentication_error(respx_mock): """Test CompactifAI authentication error handling""" diff --git a/tests/test_litellm/llms/deepseek/__init__.py b/tests/unit/llms/custom_httpx/__init__.py similarity index 100% rename from tests/test_litellm/llms/deepseek/__init__.py rename to tests/unit/llms/custom_httpx/__init__.py diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_cleanup_closed.py b/tests/unit/llms/custom_httpx/test_aiohttp_cleanup_closed.py similarity index 100% rename from tests/test_litellm/llms/custom_httpx/test_aiohttp_cleanup_closed.py rename to tests/unit/llms/custom_httpx/test_aiohttp_cleanup_closed.py diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_handler.py b/tests/unit/llms/custom_httpx/test_aiohttp_handler.py similarity index 100% rename from tests/test_litellm/llms/custom_httpx/test_aiohttp_handler.py rename to tests/unit/llms/custom_httpx/test_aiohttp_handler.py diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py b/tests/unit/llms/custom_httpx/test_aiohttp_so_keepalive.py similarity index 100% rename from tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py rename to tests/unit/llms/custom_httpx/test_aiohttp_so_keepalive.py diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py b/tests/unit/llms/custom_httpx/test_aiohttp_transport.py similarity index 100% rename from tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py rename to tests/unit/llms/custom_httpx/test_aiohttp_transport.py diff --git a/tests/test_litellm/llms/custom_httpx/test_asgi_handler.py b/tests/unit/llms/custom_httpx/test_asgi_handler.py similarity index 100% rename from tests/test_litellm/llms/custom_httpx/test_asgi_handler.py rename to tests/unit/llms/custom_httpx/test_asgi_handler.py diff --git a/tests/test_litellm/llms/custom_httpx/test_async_client_cleanup.py b/tests/unit/llms/custom_httpx/test_async_client_cleanup.py similarity index 100% rename from tests/test_litellm/llms/custom_httpx/test_async_client_cleanup.py rename to tests/unit/llms/custom_httpx/test_async_client_cleanup.py diff --git a/tests/test_litellm/llms/custom_httpx/test_container_handler.py b/tests/unit/llms/custom_httpx/test_container_handler.py similarity index 100% rename from tests/test_litellm/llms/custom_httpx/test_container_handler.py rename to tests/unit/llms/custom_httpx/test_container_handler.py diff --git a/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py b/tests/unit/llms/custom_httpx/test_credential_leak_prevention.py similarity index 100% rename from tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py rename to tests/unit/llms/custom_httpx/test_credential_leak_prevention.py diff --git a/tests/test_litellm/llms/custom_httpx/test_gemini_session_leak.py b/tests/unit/llms/custom_httpx/test_gemini_session_leak.py similarity index 100% rename from tests/test_litellm/llms/custom_httpx/test_gemini_session_leak.py rename to tests/unit/llms/custom_httpx/test_gemini_session_leak.py diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/unit/llms/custom_httpx/test_http_handler.py similarity index 100% rename from tests/test_litellm/llms/custom_httpx/test_http_handler.py rename to tests/unit/llms/custom_httpx/test_http_handler.py diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/unit/llms/custom_httpx/test_llm_http_handler.py similarity index 99% rename from tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py rename to tests/unit/llms/custom_httpx/test_llm_http_handler.py index 0350ca74904..399e4dbf206 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_llm_http_handler.py @@ -45,7 +45,7 @@ from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig from litellm.types.llms.openai import HttpxBinaryResponseContent, ResponsesAPIResponse from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ImageObject, ImageResponse, ModelResponse, TranscriptionResponse -from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe +from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe _ACTIVE_KEY = "_code_interpreter_interception_active" _SANDBOX_KEY = "_code_interpreter_interception_sandbox_key" diff --git a/tests/test_litellm/llms/custom_httpx/test_mock_transport.py b/tests/unit/llms/custom_httpx/test_mock_transport.py similarity index 100% rename from tests/test_litellm/llms/custom_httpx/test_mock_transport.py rename to tests/unit/llms/custom_httpx/test_mock_transport.py diff --git a/tests/test_litellm/llms/deepseek/chat/__init__.py b/tests/unit/llms/dashscope/__init__.py similarity index 100% rename from tests/test_litellm/llms/deepseek/chat/__init__.py rename to tests/unit/llms/dashscope/__init__.py diff --git a/tests/test_litellm/llms/dashscope/test_dashscope_chat_transformation.py b/tests/unit/llms/dashscope/test_dashscope_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/dashscope/test_dashscope_chat_transformation.py rename to tests/unit/llms/dashscope/test_dashscope_chat_transformation.py diff --git a/tests/test_litellm/llms/dashscope/test_dashscope_cost_calculator.py b/tests/unit/llms/dashscope/test_dashscope_cost_calculator.py similarity index 100% rename from tests/test_litellm/llms/dashscope/test_dashscope_cost_calculator.py rename to tests/unit/llms/dashscope/test_dashscope_cost_calculator.py diff --git a/tests/test_litellm/llms/dashscope/test_dashscope_embedding_transformation.py b/tests/unit/llms/dashscope/test_dashscope_embedding_transformation.py similarity index 100% rename from tests/test_litellm/llms/dashscope/test_dashscope_embedding_transformation.py rename to tests/unit/llms/dashscope/test_dashscope_embedding_transformation.py diff --git a/tests/test_litellm/llms/dashscope/test_dashscope_rerank_transformation.py b/tests/unit/llms/dashscope/test_dashscope_rerank_transformation.py similarity index 100% rename from tests/test_litellm/llms/dashscope/test_dashscope_rerank_transformation.py rename to tests/unit/llms/dashscope/test_dashscope_rerank_transformation.py diff --git a/tests/test_litellm/llms/dashscope/test_qwen_brand_aliases.py b/tests/unit/llms/dashscope/test_qwen_brand_aliases.py similarity index 100% rename from tests/test_litellm/llms/dashscope/test_qwen_brand_aliases.py rename to tests/unit/llms/dashscope/test_qwen_brand_aliases.py diff --git a/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py b/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py index 52bb89fed5a..9cd17bd3580 100644 --- a/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py +++ b/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py @@ -16,6 +16,9 @@ from litellm.llms.databricks.chat.transformation import ( DatabricksConfig, _sanitize_empty_content, ) +from typing import Final +import httpx +import respx @pytest.fixture() @@ -808,3 +811,75 @@ def test_chunk_parser_surfaces_top_level_reasoning_delta(reasoning_key: str) -> assert parsed.choices[0].delta.reasoning_content == "We need answer" assert parsed.choices[0].delta.content is None + + +def test_completion_merges_leading_system_and_developer_messages_for_chat_template_models( + respx_mock: respx.MockRouter, +): + upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock( + return_value=httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "my-custom-model", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10}, + }, + ) + ) + + response: Final = litellm.completion( + model="databricks/my-custom-model", + messages=[ + {"role": "system", "content": "You are terse."}, + {"role": "developer", "content": "Skills: none."}, + {"role": "user", "content": "Hello"}, + ], + api_base="https://example.databricks.test/serving-endpoints", + api_key="fake-databricks-api-key", + num_retries=0, + ) + + assert upstream.call_count == 1 + request_body: Final = json.loads(upstream.calls[0].request.read()) + assert request_body["messages"] == [ + {"role": "system", "content": "You are terse.\n\nSkills: none."}, + {"role": "user", "content": "Hello"}, + ] + assert response.choices[0].message.content == "Answer" + + +def test_completion_merges_system_messages_when_one_has_empty_content(respx_mock: respx.MockRouter): + upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock( + return_value=httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "my-custom-model", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10}, + }, + ) + ) + + litellm.completion( + model="databricks/my-custom-model", + messages=[ + {"role": "system", "content": "You are terse."}, + {"role": "system", "content": ""}, + {"role": "user", "content": "Hello"}, + ], + api_base="https://example.databricks.test/serving-endpoints", + api_key="fake-databricks-api-key", + num_retries=0, + ) + + request_body: Final = json.loads(upstream.calls[0].request.read()) + assert request_body["messages"] == [ + {"role": "system", "content": "You are terse."}, + {"role": "user", "content": "Hello"}, + ] diff --git a/tests/test_litellm/llms/databricks/test_databricks_common_utils.py b/tests/unit/llms/databricks/test_databricks_common_utils.py similarity index 100% rename from tests/test_litellm/llms/databricks/test_databricks_common_utils.py rename to tests/unit/llms/databricks/test_databricks_common_utils.py diff --git a/tests/test_litellm/llms/databricks/test_databricks_cost_calculator.py b/tests/unit/llms/databricks/test_databricks_cost_calculator.py similarity index 100% rename from tests/test_litellm/llms/databricks/test_databricks_cost_calculator.py rename to tests/unit/llms/databricks/test_databricks_cost_calculator.py diff --git a/tests/test_litellm/llms/databricks/test_databricks_partner_integration.py b/tests/unit/llms/databricks/test_databricks_partner_integration.py similarity index 100% rename from tests/test_litellm/llms/databricks/test_databricks_partner_integration.py rename to tests/unit/llms/databricks/test_databricks_partner_integration.py diff --git a/tests/test_litellm/llms/databricks/test_databricks_streaming_utils.py b/tests/unit/llms/databricks/test_databricks_streaming_utils.py similarity index 100% rename from tests/test_litellm/llms/databricks/test_databricks_streaming_utils.py rename to tests/unit/llms/databricks/test_databricks_streaming_utils.py diff --git a/tests/test_litellm/llms/deepseek/messages/__init__.py b/tests/unit/llms/deepgram/__init__.py similarity index 100% rename from tests/test_litellm/llms/deepseek/messages/__init__.py rename to tests/unit/llms/deepgram/__init__.py diff --git a/tests/test_litellm/llms/gemini/__init__.py b/tests/unit/llms/deepgram/audio_transcription/__init__.py similarity index 100% rename from tests/test_litellm/llms/gemini/__init__.py rename to tests/unit/llms/deepgram/audio_transcription/__init__.py diff --git a/tests/test_litellm/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py b/tests/unit/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py similarity index 100% rename from tests/test_litellm/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py rename to tests/unit/llms/deepgram/audio_transcription/test_deepgram_audio_transcription_transformation.py diff --git a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py b/tests/unit/llms/deepgram/test_deepgram_common_utils.py similarity index 100% rename from tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py rename to tests/unit/llms/deepgram/test_deepgram_common_utils.py diff --git a/tests/test_litellm/llms/deepgram/test_deepgram_mock_transcription.py b/tests/unit/llms/deepgram/test_deepgram_mock_transcription.py similarity index 100% rename from tests/test_litellm/llms/deepgram/test_deepgram_mock_transcription.py rename to tests/unit/llms/deepgram/test_deepgram_mock_transcription.py diff --git a/tests/test_litellm/llms/gemini/audio_transcription/__init__.py b/tests/unit/llms/deepinfra/__init__.py similarity index 100% rename from tests/test_litellm/llms/gemini/audio_transcription/__init__.py rename to tests/unit/llms/deepinfra/__init__.py diff --git a/tests/test_litellm/llms/deepinfra/test_deepinfra_chat_transformation.py b/tests/unit/llms/deepinfra/test_deepinfra_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/deepinfra/test_deepinfra_chat_transformation.py rename to tests/unit/llms/deepinfra/test_deepinfra_chat_transformation.py diff --git a/tests/test_litellm/llms/deepinfra/test_deepinfra_rerank.py b/tests/unit/llms/deepinfra/test_deepinfra_rerank.py similarity index 100% rename from tests/test_litellm/llms/deepinfra/test_deepinfra_rerank.py rename to tests/unit/llms/deepinfra/test_deepinfra_rerank.py diff --git a/tests/unit/llms/deepinfra/test_deepinfra_rerank_integration.py b/tests/unit/llms/deepinfra/test_deepinfra_rerank_integration.py new file mode 100644 index 00000000000..8a2a1d09cb6 --- /dev/null +++ b/tests/unit/llms/deepinfra/test_deepinfra_rerank_integration.py @@ -0,0 +1,159 @@ +""" +Integration tests for DeepInfra rerank functionality. +Tests the full rerank flow following the repository patterns. +""" + +import asyncio +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +import litellm + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post") +@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") +def test_deepinfra_rerank_with_queries_param( + mock_sync_post, mock_async_post, sync_mode +): + """Test DeepInfra rerank with multiple queries parameter.""" + mock_response_data = { + "scores": [0.8, 0.6, 0.2], + "input_tokens": 35, + "request_id": "deepinfra-multi-query-123", + "inference_status": {"status": "success", "runtime_ms": 200}, + } + + def return_val(): + return mock_response_data + + if sync_mode: + mock_response = MagicMock() + mock_response.json = return_val + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.text = json.dumps(mock_response_data) + mock_sync_post.return_value = mock_response + + response = litellm.rerank( + model="deepinfra/Qwen/Qwen3-Reranker-4B", + query="hello", + documents=["hello", "world", "test"], + queries=["hello", "hi there"], # DeepInfra specific param + custom_llm_provider="deepinfra", + api_key="test_key", + api_base="https://api.deepinfra.com", + ) + + mock_sync_post.assert_called_once() + # Verify that queries parameter was passed in request + call_data = json.loads(mock_sync_post.call_args.kwargs["data"]) + assert "queries" in call_data + assert call_data["queries"] == ["hello", "hi there"] + else: + mock_response = AsyncMock() + mock_response.json = return_val + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.text = json.dumps(mock_response_data) + mock_async_post.return_value = mock_response + + response = asyncio.run( + litellm.arerank( + model="deepinfra/Qwen/Qwen3-Reranker-4B", + query="hello", + documents=["hello", "world", "test"], + queries=["hello", "hi there"], + custom_llm_provider="deepinfra", + api_key="test_key", + api_base="https://api.deepinfra.com", + ) + ) + + mock_async_post.assert_called_once() + call_data = json.loads(mock_async_post.call_args.kwargs["data"]) + assert "queries" in call_data + assert call_data["queries"] == ["hello", "hi there"] + + assert response.results is not None + assert len(response.results) == 3 + + +@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") +def test_deepinfra_rerank_with_env_vars(mock_post, monkeypatch): + """Test DeepInfra rerank with environment variable configuration.""" + monkeypatch.setenv("DEEPINFRA_API_KEY", "env_test_key") + monkeypatch.setenv("DEEPINFRA_API_BASE", "https://custom-deepinfra.com") + + mock_response_data = { + "scores": [0.88, 0.22], + "input_tokens": 28, + "request_id": "env-test-123", + } + + def return_val(): + return mock_response_data + + mock_response = MagicMock() + mock_response.json = return_val + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.text = json.dumps(mock_response_data) + mock_post.return_value = mock_response + + response = litellm.rerank( + model="deepinfra/Qwen/Qwen3-Reranker-0.6B", + query="hello", + documents=["hello", "world"], + custom_llm_provider="deepinfra", + ) + + mock_post.assert_called_once() + + # Verify headers contain env API key + headers = mock_post.call_args.kwargs.get("headers", {}) + assert "Bearer env_test_key" in headers.get("Authorization", "") + + assert response.results is not None + + +@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") +def test_deepinfra_rerank_defaults_api_base_when_missing(mock_post, monkeypatch): + """With no api_base anywhere, the call still goes out against DeepInfra's own base.""" + monkeypatch.delenv("DEEPINFRA_API_BASE", raising=False) + + mock_response = MagicMock() + mock_response.json = lambda: {"scores": [0.9, 0.1], "input_tokens": 20} + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_post.return_value = mock_response + + response = litellm.rerank( + model="deepinfra/Qwen/Qwen3-Reranker-0.6B", + query="hello", + documents=["hello", "world"], + custom_llm_provider="deepinfra", + api_key="test_key", + # api_base is intentionally missing + ) + + assert "api.deepinfra.com" in mock_post.call_args.kwargs["url"] + assert [result["relevance_score"] for result in response.results] == [0.9, 0.1] + + +def test_deepinfra_rerank_models(): + """Test that DeepInfra Qwen rerank models are recognized.""" + # These should not raise errors during model validation + models = [ + "deepinfra/Qwen/Qwen3-Reranker-0.6B", + "deepinfra/Qwen/Qwen3-Reranker-4B", + "deepinfra/Qwen/Qwen3-Reranker-8B", + ] + + for model in models: + resolved_model, provider, _, api_base = litellm.get_llm_provider(model=model) + assert provider == "deepinfra" + assert resolved_model == model.removeprefix("deepinfra/") + assert api_base == "https://api.deepinfra.com/v1/openai" diff --git a/tests/test_litellm/llms/deepinfra/test_deepinfra_rerank_transformation.py b/tests/unit/llms/deepinfra/test_deepinfra_rerank_transformation.py similarity index 100% rename from tests/test_litellm/llms/deepinfra/test_deepinfra_rerank_transformation.py rename to tests/unit/llms/deepinfra/test_deepinfra_rerank_transformation.py diff --git a/tests/test_litellm/llms/gemini/google_genai/__init__.py b/tests/unit/llms/edenai/__init__.py similarity index 100% rename from tests/test_litellm/llms/gemini/google_genai/__init__.py rename to tests/unit/llms/edenai/__init__.py diff --git a/tests/test_litellm/llms/gemini/google_genai/guardrail_translation/__init__.py b/tests/unit/llms/edenai/audio_transcription/__init__.py similarity index 100% rename from tests/test_litellm/llms/gemini/google_genai/guardrail_translation/__init__.py rename to tests/unit/llms/edenai/audio_transcription/__init__.py diff --git a/tests/test_litellm/llms/edenai/audio_transcription/test_edenai_audio_transcription_transformation.py b/tests/unit/llms/edenai/audio_transcription/test_edenai_audio_transcription_transformation.py similarity index 100% rename from tests/test_litellm/llms/edenai/audio_transcription/test_edenai_audio_transcription_transformation.py rename to tests/unit/llms/edenai/audio_transcription/test_edenai_audio_transcription_transformation.py diff --git a/tests/test_litellm/llms/gemini/image_edit/__init__.py b/tests/unit/llms/edenai/chat/__init__.py similarity index 100% rename from tests/test_litellm/llms/gemini/image_edit/__init__.py rename to tests/unit/llms/edenai/chat/__init__.py diff --git a/tests/test_litellm/llms/edenai/chat/test_edenai_chat_transformation.py b/tests/unit/llms/edenai/chat/test_edenai_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/edenai/chat/test_edenai_chat_transformation.py rename to tests/unit/llms/edenai/chat/test_edenai_chat_transformation.py diff --git a/tests/test_litellm/llms/edenai/conftest.py b/tests/unit/llms/edenai/conftest.py similarity index 100% rename from tests/test_litellm/llms/edenai/conftest.py rename to tests/unit/llms/edenai/conftest.py diff --git a/tests/test_litellm/llms/gemini/realtime/__init__.py b/tests/unit/llms/edenai/embedding/__init__.py similarity index 100% rename from tests/test_litellm/llms/gemini/realtime/__init__.py rename to tests/unit/llms/edenai/embedding/__init__.py diff --git a/tests/test_litellm/llms/edenai/embedding/test_edenai_embedding_transformation.py b/tests/unit/llms/edenai/embedding/test_edenai_embedding_transformation.py similarity index 100% rename from tests/test_litellm/llms/edenai/embedding/test_edenai_embedding_transformation.py rename to tests/unit/llms/edenai/embedding/test_edenai_embedding_transformation.py diff --git a/tests/test_litellm/llms/gigachat/__init__.py b/tests/unit/llms/edenai/image_generation/__init__.py similarity index 100% rename from tests/test_litellm/llms/gigachat/__init__.py rename to tests/unit/llms/edenai/image_generation/__init__.py diff --git a/tests/test_litellm/llms/edenai/image_generation/test_edenai_image_generation_transformation.py b/tests/unit/llms/edenai/image_generation/test_edenai_image_generation_transformation.py similarity index 100% rename from tests/test_litellm/llms/edenai/image_generation/test_edenai_image_generation_transformation.py rename to tests/unit/llms/edenai/image_generation/test_edenai_image_generation_transformation.py diff --git a/tests/test_litellm/llms/gigachat/embedding/__init__.py b/tests/unit/llms/edenai/messages/__init__.py similarity index 100% rename from tests/test_litellm/llms/gigachat/embedding/__init__.py rename to tests/unit/llms/edenai/messages/__init__.py diff --git a/tests/test_litellm/llms/edenai/messages/test_edenai_anthropic_messages_transformation.py b/tests/unit/llms/edenai/messages/test_edenai_anthropic_messages_transformation.py similarity index 100% rename from tests/test_litellm/llms/edenai/messages/test_edenai_anthropic_messages_transformation.py rename to tests/unit/llms/edenai/messages/test_edenai_anthropic_messages_transformation.py diff --git a/tests/test_litellm/llms/gigachat/passthrough/__init__.py b/tests/unit/llms/edenai/responses/__init__.py similarity index 100% rename from tests/test_litellm/llms/gigachat/passthrough/__init__.py rename to tests/unit/llms/edenai/responses/__init__.py diff --git a/tests/test_litellm/llms/edenai/responses/test_edenai_responses_transformation.py b/tests/unit/llms/edenai/responses/test_edenai_responses_transformation.py similarity index 100% rename from tests/test_litellm/llms/edenai/responses/test_edenai_responses_transformation.py rename to tests/unit/llms/edenai/responses/test_edenai_responses_transformation.py diff --git a/tests/test_litellm/llms/edenai/test_edenai_common_utils.py b/tests/unit/llms/edenai/test_edenai_common_utils.py similarity index 100% rename from tests/test_litellm/llms/edenai/test_edenai_common_utils.py rename to tests/unit/llms/edenai/test_edenai_common_utils.py diff --git a/tests/test_litellm/llms/github_copilot/messages/__init__.py b/tests/unit/llms/edenai/text_to_speech/__init__.py similarity index 100% rename from tests/test_litellm/llms/github_copilot/messages/__init__.py rename to tests/unit/llms/edenai/text_to_speech/__init__.py diff --git a/tests/test_litellm/llms/edenai/text_to_speech/test_edenai_text_to_speech_transformation.py b/tests/unit/llms/edenai/text_to_speech/test_edenai_text_to_speech_transformation.py similarity index 100% rename from tests/test_litellm/llms/edenai/text_to_speech/test_edenai_text_to_speech_transformation.py rename to tests/unit/llms/edenai/text_to_speech/test_edenai_text_to_speech_transformation.py diff --git a/tests/test_litellm/llms/gradient_ai/__init__.py b/tests/unit/llms/edenai/videos/__init__.py similarity index 100% rename from tests/test_litellm/llms/gradient_ai/__init__.py rename to tests/unit/llms/edenai/videos/__init__.py diff --git a/tests/test_litellm/llms/edenai/videos/test_edenai_video_transformation.py b/tests/unit/llms/edenai/videos/test_edenai_video_transformation.py similarity index 100% rename from tests/test_litellm/llms/edenai/videos/test_edenai_video_transformation.py rename to tests/unit/llms/edenai/videos/test_edenai_video_transformation.py diff --git a/tests/test_litellm/llms/gradient_ai/chat/__init__.py b/tests/unit/llms/fal_ai/__init__.py similarity index 100% rename from tests/test_litellm/llms/gradient_ai/chat/__init__.py rename to tests/unit/llms/fal_ai/__init__.py diff --git a/tests/test_litellm/llms/groq/__init__.py b/tests/unit/llms/fal_ai/chat/__init__.py similarity index 100% rename from tests/test_litellm/llms/groq/__init__.py rename to tests/unit/llms/fal_ai/chat/__init__.py diff --git a/tests/test_litellm/llms/fal_ai/chat/test_fal_ai_chat_transformation.py b/tests/unit/llms/fal_ai/chat/test_fal_ai_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/fal_ai/chat/test_fal_ai_chat_transformation.py rename to tests/unit/llms/fal_ai/chat/test_fal_ai_chat_transformation.py diff --git a/tests/test_litellm/llms/groq/chat/__init__.py b/tests/unit/llms/fal_ai/image_edit/__init__.py similarity index 100% rename from tests/test_litellm/llms/groq/chat/__init__.py rename to tests/unit/llms/fal_ai/image_edit/__init__.py diff --git a/tests/test_litellm/llms/fal_ai/image_edit/test_fal_ai_flux_lora_depth_transformation.py b/tests/unit/llms/fal_ai/image_edit/test_fal_ai_flux_lora_depth_transformation.py similarity index 100% rename from tests/test_litellm/llms/fal_ai/image_edit/test_fal_ai_flux_lora_depth_transformation.py rename to tests/unit/llms/fal_ai/image_edit/test_fal_ai_flux_lora_depth_transformation.py diff --git a/tests/test_litellm/llms/fal_ai/image_edit/test_fal_ai_image_edit_transformation.py b/tests/unit/llms/fal_ai/image_edit/test_fal_ai_image_edit_transformation.py similarity index 100% rename from tests/test_litellm/llms/fal_ai/image_edit/test_fal_ai_image_edit_transformation.py rename to tests/unit/llms/fal_ai/image_edit/test_fal_ai_image_edit_transformation.py diff --git a/tests/test_litellm/llms/huggingface/__init__.py b/tests/unit/llms/fal_ai/image_generation/__init__.py similarity index 100% rename from tests/test_litellm/llms/huggingface/__init__.py rename to tests/unit/llms/fal_ai/image_generation/__init__.py diff --git a/tests/test_litellm/llms/fal_ai/image_generation/test_fal_ai_flux_dev_transformation.py b/tests/unit/llms/fal_ai/image_generation/test_fal_ai_flux_dev_transformation.py similarity index 100% rename from tests/test_litellm/llms/fal_ai/image_generation/test_fal_ai_flux_dev_transformation.py rename to tests/unit/llms/fal_ai/image_generation/test_fal_ai_flux_dev_transformation.py diff --git a/tests/test_litellm/llms/fal_ai/image_generation/test_fal_ai_gpt_image_2_transformation.py b/tests/unit/llms/fal_ai/image_generation/test_fal_ai_gpt_image_2_transformation.py similarity index 100% rename from tests/test_litellm/llms/fal_ai/image_generation/test_fal_ai_gpt_image_2_transformation.py rename to tests/unit/llms/fal_ai/image_generation/test_fal_ai_gpt_image_2_transformation.py diff --git a/tests/test_litellm/llms/fal_ai/image_generation/test_fal_ai_nano_banana_transformation.py b/tests/unit/llms/fal_ai/image_generation/test_fal_ai_nano_banana_transformation.py similarity index 100% rename from tests/test_litellm/llms/fal_ai/image_generation/test_fal_ai_nano_banana_transformation.py rename to tests/unit/llms/fal_ai/image_generation/test_fal_ai_nano_banana_transformation.py diff --git a/tests/test_litellm/llms/fal_ai/test_cost_calculator.py b/tests/unit/llms/fal_ai/test_cost_calculator.py similarity index 100% rename from tests/test_litellm/llms/fal_ai/test_cost_calculator.py rename to tests/unit/llms/fal_ai/test_cost_calculator.py diff --git a/tests/test_litellm/llms/inception/__init__.py b/tests/unit/llms/fal_ai/videos/__init__.py similarity index 100% rename from tests/test_litellm/llms/inception/__init__.py rename to tests/unit/llms/fal_ai/videos/__init__.py diff --git a/tests/test_litellm/llms/fal_ai/videos/test_fal_ai_video_transformation.py b/tests/unit/llms/fal_ai/videos/test_fal_ai_video_transformation.py similarity index 100% rename from tests/test_litellm/llms/fal_ai/videos/test_fal_ai_video_transformation.py rename to tests/unit/llms/fal_ai/videos/test_fal_ai_video_transformation.py diff --git a/tests/test_litellm/llms/mistral/batches/__init__.py b/tests/unit/llms/featherless_ai/__init__.py similarity index 100% rename from tests/test_litellm/llms/mistral/batches/__init__.py rename to tests/unit/llms/featherless_ai/__init__.py diff --git a/tests/test_litellm/llms/mistral/files/__init__.py b/tests/unit/llms/featherless_ai/chat/__init__.py similarity index 100% rename from tests/test_litellm/llms/mistral/files/__init__.py rename to tests/unit/llms/featherless_ai/chat/__init__.py diff --git a/tests/test_litellm/llms/featherless_ai/chat/test_featherless_chat_transformation.py b/tests/unit/llms/featherless_ai/chat/test_featherless_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/featherless_ai/chat/test_featherless_chat_transformation.py rename to tests/unit/llms/featherless_ai/chat/test_featherless_chat_transformation.py diff --git a/tests/test_litellm/llms/nvidia_riva/__init__.py b/tests/unit/llms/fireworks_ai/completion/__init__.py similarity index 100% rename from tests/test_litellm/llms/nvidia_riva/__init__.py rename to tests/unit/llms/fireworks_ai/completion/__init__.py diff --git a/tests/test_litellm/llms/fireworks_ai/completion/test_fireworks_ai_completion_transformation.py b/tests/unit/llms/fireworks_ai/completion/test_fireworks_ai_completion_transformation.py similarity index 100% rename from tests/test_litellm/llms/fireworks_ai/completion/test_fireworks_ai_completion_transformation.py rename to tests/unit/llms/fireworks_ai/completion/test_fireworks_ai_completion_transformation.py diff --git a/tests/test_litellm/llms/fireworks_ai/completion/test_fireworks_ai_text_completion_transformation.py b/tests/unit/llms/fireworks_ai/completion/test_fireworks_ai_text_completion_transformation.py similarity index 100% rename from tests/test_litellm/llms/fireworks_ai/completion/test_fireworks_ai_text_completion_transformation.py rename to tests/unit/llms/fireworks_ai/completion/test_fireworks_ai_text_completion_transformation.py diff --git a/tests/test_litellm/llms/oci/rerank/__init__.py b/tests/unit/llms/gdc/__init__.py similarity index 100% rename from tests/test_litellm/llms/oci/rerank/__init__.py rename to tests/unit/llms/gdc/__init__.py diff --git a/tests/test_litellm/llms/ocr/__init__.py b/tests/unit/llms/gdc/chat/__init__.py similarity index 100% rename from tests/test_litellm/llms/ocr/__init__.py rename to tests/unit/llms/gdc/chat/__init__.py diff --git a/tests/test_litellm/llms/gdc/chat/test_gdc_chat_transformation.py b/tests/unit/llms/gdc/chat/test_gdc_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/gdc/chat/test_gdc_chat_transformation.py rename to tests/unit/llms/gdc/chat/test_gdc_chat_transformation.py diff --git a/tests/test_litellm/llms/gemini/test_cost_calculator.py b/tests/unit/llms/gemini/test_cost_calculator.py similarity index 100% rename from tests/test_litellm/llms/gemini/test_cost_calculator.py rename to tests/unit/llms/gemini/test_cost_calculator.py diff --git a/tests/test_litellm/llms/gemini/test_gemini_client_setup.py b/tests/unit/llms/gemini/test_gemini_client_setup.py similarity index 100% rename from tests/test_litellm/llms/gemini/test_gemini_client_setup.py rename to tests/unit/llms/gemini/test_gemini_client_setup.py diff --git a/tests/test_litellm/llms/gemini/test_gemini_common_utils.py b/tests/unit/llms/gemini/test_gemini_common_utils.py similarity index 100% rename from tests/test_litellm/llms/gemini/test_gemini_common_utils.py rename to tests/unit/llms/gemini/test_gemini_common_utils.py diff --git a/tests/test_litellm/llms/gemini/test_gemini_image_generation_transformation.py b/tests/unit/llms/gemini/test_gemini_image_generation_transformation.py similarity index 100% rename from tests/test_litellm/llms/gemini/test_gemini_image_generation_transformation.py rename to tests/unit/llms/gemini/test_gemini_image_generation_transformation.py diff --git a/tests/test_litellm/llms/gemini/test_gemini_tts.py b/tests/unit/llms/gemini/test_gemini_tts.py similarity index 100% rename from tests/test_litellm/llms/gemini/test_gemini_tts.py rename to tests/unit/llms/gemini/test_gemini_tts.py diff --git a/tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py b/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py similarity index 100% rename from tests/test_litellm/llms/github_copilot/test_github_copilot_authenticator.py rename to tests/unit/llms/github_copilot/test_github_copilot_authenticator.py diff --git a/tests/test_litellm/llms/github_copilot/test_github_copilot_transformation.py b/tests/unit/llms/github_copilot/test_github_copilot_transformation.py similarity index 100% rename from tests/test_litellm/llms/github_copilot/test_github_copilot_transformation.py rename to tests/unit/llms/github_copilot/test_github_copilot_transformation.py diff --git a/tests/test_litellm/llms/openai_like/responses/__init__.py b/tests/unit/llms/heroku/__init__.py similarity index 100% rename from tests/test_litellm/llms/openai_like/responses/__init__.py rename to tests/unit/llms/heroku/__init__.py diff --git a/tests/test_litellm/llms/heroku/test_heroku_chat_transformation.py b/tests/unit/llms/heroku/test_heroku_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/heroku/test_heroku_chat_transformation.py rename to tests/unit/llms/heroku/test_heroku_chat_transformation.py diff --git a/tests/test_litellm/llms/parallel_ai/__init__.py b/tests/unit/llms/huggingface/embedding/__init__.py similarity index 100% rename from tests/test_litellm/llms/parallel_ai/__init__.py rename to tests/unit/llms/huggingface/embedding/__init__.py diff --git a/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py b/tests/unit/llms/huggingface/embedding/test_huggingface_embedding_handler.py similarity index 100% rename from tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py rename to tests/unit/llms/huggingface/embedding/test_huggingface_embedding_handler.py diff --git a/tests/test_litellm/llms/langflow/test_langflow_a2a.py b/tests/unit/llms/langflow/test_langflow_a2a.py similarity index 100% rename from tests/test_litellm/llms/langflow/test_langflow_a2a.py rename to tests/unit/llms/langflow/test_langflow_a2a.py diff --git a/tests/test_litellm/llms/pass_through/__init__.py b/tests/unit/llms/lemonade/__init__.py similarity index 100% rename from tests/test_litellm/llms/pass_through/__init__.py rename to tests/unit/llms/lemonade/__init__.py diff --git a/tests/test_litellm/llms/lemonade/test_lemonade.py b/tests/unit/llms/lemonade/test_lemonade.py similarity index 100% rename from tests/test_litellm/llms/lemonade/test_lemonade.py rename to tests/unit/llms/lemonade/test_lemonade.py diff --git a/tests/test_litellm/llms/pass_through/guardrail_translation/__init__.py b/tests/unit/llms/lm_studio/__init__.py similarity index 100% rename from tests/test_litellm/llms/pass_through/guardrail_translation/__init__.py rename to tests/unit/llms/lm_studio/__init__.py diff --git a/tests/test_litellm/llms/lm_studio/test_lm_studio_chat_transformation.py b/tests/unit/llms/lm_studio/test_lm_studio_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/lm_studio/test_lm_studio_chat_transformation.py rename to tests/unit/llms/lm_studio/test_lm_studio_chat_transformation.py diff --git a/tests/test_litellm/llms/perplexity/__init__.py b/tests/unit/llms/mistral/audio_transcription/__init__.py similarity index 100% rename from tests/test_litellm/llms/perplexity/__init__.py rename to tests/unit/llms/mistral/audio_transcription/__init__.py diff --git a/tests/unit/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py b/tests/unit/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py new file mode 100644 index 00000000000..68875ff6d32 --- /dev/null +++ b/tests/unit/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py @@ -0,0 +1,195 @@ +import os +from unittest.mock import MagicMock + +import httpx +import litellm + +from litellm.llms.base_llm.audio_transcription.transformation import ( + BaseAudioTranscriptionConfig, +) +from litellm.llms.mistral.audio_transcription.transformation import ( + MistralAudioTranscriptionConfig, +) +from litellm.types.utils import TranscriptionResponse +from litellm.utils import ProviderConfigManager + + +def test_mistral_audio_transcription_config_installed(): + """Ensure Mistral audio transcription config is registered with ProviderConfigManager.""" + config = ProviderConfigManager.get_provider_audio_transcription_config( + model="mistral/voxtral-mini-latest", + provider=litellm.LlmProviders.MISTRAL, + ) + assert config is not None + assert isinstance(config, BaseAudioTranscriptionConfig) + assert isinstance(config, MistralAudioTranscriptionConfig) + + +def test_mistral_audio_transcription_get_complete_url(): + config = MistralAudioTranscriptionConfig() + url = config.get_complete_url( + api_base=None, + api_key="fake-key", + model="voxtral-mini-latest", + optional_params={}, + litellm_params={}, + ) + assert url == "https://api.mistral.ai/v1/audio/transcriptions" + + +def test_mistral_audio_transcription_get_complete_url_custom_base(): + config = MistralAudioTranscriptionConfig() + url = config.get_complete_url( + api_base="https://custom.api.example.com/v1/", + api_key="fake-key", + model="voxtral-mini-latest", + optional_params={}, + litellm_params={}, + ) + assert url == "https://custom.api.example.com/v1/audio/transcriptions" + + +def test_mistral_audio_transcription_validate_environment(): + config = MistralAudioTranscriptionConfig() + headers = config.validate_environment( + headers={}, + model="voxtral-mini-latest", + messages=[], + optional_params={}, + litellm_params={}, + api_key="test-key-123", + ) + assert headers["Authorization"] == "Bearer test-key-123" + assert headers["accept"] == "application/json" + + +def test_mistral_audio_transcription_supported_params(): + config = MistralAudioTranscriptionConfig() + params = config.get_supported_openai_params("voxtral-mini-latest") + assert "language" in params + assert "temperature" in params + assert "response_format" in params + assert "timestamp_granularities" in params + + +def test_mistral_audio_transcription_request_transform(): + config = MistralAudioTranscriptionConfig() + + wav_path = os.path.join( + os.path.dirname(__file__), + "../../../../..", + "tests", + "llm_translation", + "gettysburg.wav", + ) + audio_file = open(wav_path, "rb") + + result = config.transform_audio_transcription_request( + model="voxtral-mini-latest", + audio_file=audio_file, + optional_params={"language": "en", "temperature": 0.0}, + litellm_params={}, + ) + + audio_file.close() + + assert isinstance(result.data, dict) + assert result.data["model"] == "voxtral-mini-latest" + assert result.data["language"] == "en" + assert result.data["temperature"] == 0.0 + assert result.files is not None + assert "file" in result.files + + +def test_mistral_audio_transcription_request_with_diarize(): + """Test that Mistral-specific params like diarize are passed through.""" + config = MistralAudioTranscriptionConfig() + + wav_path = os.path.join( + os.path.dirname(__file__), + "../../../../..", + "tests", + "llm_translation", + "gettysburg.wav", + ) + audio_file = open(wav_path, "rb") + + result = config.transform_audio_transcription_request( + model="voxtral-mini-latest", + audio_file=audio_file, + optional_params={"diarize": True}, + litellm_params={}, + ) + + audio_file.close() + + assert isinstance(result.data, dict) + assert result.data["diarize"] == "true" + + +def test_mistral_audio_transcription_response_transform(): + config = MistralAudioTranscriptionConfig() + + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = {"text": "Four score and seven years ago..."} + + response = config.transform_audio_transcription_response(mock_response) + + assert isinstance(response, TranscriptionResponse) + assert response.text == "Four score and seven years ago..." + + +def test_mistral_audio_transcription_response_transform_diarized(): + """Test that diarized responses preserve segments and language.""" + config = MistralAudioTranscriptionConfig() + + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = { + "model": "voxtral-mini-latest", + "text": "Hello, how are you? I am fine.", + "language": None, + "segments": [ + { + "text": "Hello, how are you?", + "start": 0.3, + "end": 2.1, + "speaker_id": "speaker_1", + "type": "transcription_segment", + }, + { + "text": "I am fine.", + "start": 2.5, + "end": 3.8, + "speaker_id": "speaker_2", + "type": "transcription_segment", + }, + ], + "usage": { + "prompt_audio_seconds": 4, + "prompt_tokens": 5, + "total_tokens": 50, + "completion_tokens": 20, + }, + } + + response = config.transform_audio_transcription_response(mock_response) + + assert isinstance(response, TranscriptionResponse) + assert response.text == "Hello, how are you? I am fine." + assert response["segments"] is not None + assert len(response["segments"]) == 2 + assert response["segments"][0]["speaker_id"] == "speaker_1" + assert response["segments"][1]["speaker_id"] == "speaker_2" + assert response["language"] is None + + +def test_mistral_audio_transcription_response_transform_empty(): + config = MistralAudioTranscriptionConfig() + + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = {} + + response = config.transform_audio_transcription_response(mock_response) + + assert isinstance(response, TranscriptionResponse) + assert response.text == "" diff --git a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py b/tests/unit/llms/mistral/test_mistral_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py rename to tests/unit/llms/mistral/test_mistral_chat_transformation.py diff --git a/tests/test_litellm/llms/mistral/test_mistral_completion.py b/tests/unit/llms/mistral/test_mistral_completion.py similarity index 100% rename from tests/test_litellm/llms/mistral/test_mistral_completion.py rename to tests/unit/llms/mistral/test_mistral_completion.py diff --git a/tests/test_litellm/llms/perplexity/embedding/__init__.py b/tests/unit/llms/modelscope/chat/__init__.py similarity index 100% rename from tests/test_litellm/llms/perplexity/embedding/__init__.py rename to tests/unit/llms/modelscope/chat/__init__.py diff --git a/tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py b/tests/unit/llms/modelscope/chat/test_modelscope_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py rename to tests/unit/llms/modelscope/chat/test_modelscope_chat_transformation.py diff --git a/tests/test_litellm/llms/stability/__init__.py b/tests/unit/llms/nadir/__init__.py similarity index 100% rename from tests/test_litellm/llms/stability/__init__.py rename to tests/unit/llms/nadir/__init__.py diff --git a/tests/test_litellm/llms/nadir/test_nadir.py b/tests/unit/llms/nadir/test_nadir.py similarity index 100% rename from tests/test_litellm/llms/nadir/test_nadir.py rename to tests/unit/llms/nadir/test_nadir.py diff --git a/tests/test_litellm/llms/stability/image_generation/__init__.py b/tests/unit/llms/nebius/__init__.py similarity index 100% rename from tests/test_litellm/llms/stability/image_generation/__init__.py rename to tests/unit/llms/nebius/__init__.py diff --git a/tests/test_litellm/llms/nebius/test_nebius_chat_transformation.py b/tests/unit/llms/nebius/test_nebius_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/nebius/test_nebius_chat_transformation.py rename to tests/unit/llms/nebius/test_nebius_chat_transformation.py diff --git a/tests/test_litellm/llms/nebius/test_nebius_embedding_transformation.py b/tests/unit/llms/nebius/test_nebius_embedding_transformation.py similarity index 100% rename from tests/test_litellm/llms/nebius/test_nebius_embedding_transformation.py rename to tests/unit/llms/nebius/test_nebius_embedding_transformation.py diff --git a/tests/test_litellm/llms/tencent/__init__.py b/tests/unit/llms/oci/rerank/__init__.py similarity index 100% rename from tests/test_litellm/llms/tencent/__init__.py rename to tests/unit/llms/oci/rerank/__init__.py diff --git a/tests/test_litellm/llms/oci/test_oci_common_utils.py b/tests/unit/llms/oci/test_oci_common_utils.py similarity index 100% rename from tests/test_litellm/llms/oci/test_oci_common_utils.py rename to tests/unit/llms/oci/test_oci_common_utils.py diff --git a/tests/test_litellm/llms/oci/test_oci_coverage_boost.py b/tests/unit/llms/oci/test_oci_coverage_boost.py similarity index 100% rename from tests/test_litellm/llms/oci/test_oci_coverage_boost.py rename to tests/unit/llms/oci/test_oci_coverage_boost.py diff --git a/tests/test_litellm/llms/tencent/chat/__init__.py b/tests/unit/llms/ollama/__init__.py similarity index 100% rename from tests/test_litellm/llms/tencent/chat/__init__.py rename to tests/unit/llms/ollama/__init__.py diff --git a/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py b/tests/unit/llms/ollama/test_ollama_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py rename to tests/unit/llms/ollama/test_ollama_chat_transformation.py diff --git a/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py b/tests/unit/llms/ollama/test_ollama_completion_transformation.py similarity index 100% rename from tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py rename to tests/unit/llms/ollama/test_ollama_completion_transformation.py diff --git a/tests/test_litellm/llms/ollama/test_ollama_embedding.py b/tests/unit/llms/ollama/test_ollama_embedding.py similarity index 100% rename from tests/test_litellm/llms/ollama/test_ollama_embedding.py rename to tests/unit/llms/ollama/test_ollama_embedding.py diff --git a/tests/test_litellm/llms/ollama/test_ollama_model_info.py b/tests/unit/llms/ollama/test_ollama_model_info.py similarity index 100% rename from tests/test_litellm/llms/ollama/test_ollama_model_info.py rename to tests/unit/llms/ollama/test_ollama_model_info.py diff --git a/tests/test_litellm/llms/openai/realtime/README.md b/tests/unit/llms/openai/realtime/README.md similarity index 100% rename from tests/test_litellm/llms/openai/realtime/README.md rename to tests/unit/llms/openai/realtime/README.md diff --git a/tests/test_litellm/llms/tencent/messages/__init__.py b/tests/unit/llms/openai/realtime/__init__.py similarity index 100% rename from tests/test_litellm/llms/tencent/messages/__init__.py rename to tests/unit/llms/openai/realtime/__init__.py diff --git a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py b/tests/unit/llms/openai/realtime/test_openai_realtime_handler.py similarity index 100% rename from tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py rename to tests/unit/llms/openai/realtime/test_openai_realtime_handler.py diff --git a/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py b/tests/unit/llms/openai/realtime/test_transcription_sessions.py similarity index 100% rename from tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py rename to tests/unit/llms/openai/realtime/test_transcription_sessions.py diff --git a/tests/test_litellm/llms/vercel_ai_gateway/embedding/__init__.py b/tests/unit/llms/openai/responses/__init__.py similarity index 100% rename from tests/test_litellm/llms/vercel_ai_gateway/embedding/__init__.py rename to tests/unit/llms/openai/responses/__init__.py diff --git a/tests/test_litellm/llms/openai/responses/test_openai_count_tokens_transformation.py b/tests/unit/llms/openai/responses/test_openai_count_tokens_transformation.py similarity index 100% rename from tests/test_litellm/llms/openai/responses/test_openai_count_tokens_transformation.py rename to tests/unit/llms/openai/responses/test_openai_count_tokens_transformation.py diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_data_residency.py b/tests/unit/llms/openai/responses/test_openai_responses_data_residency.py similarity index 100% rename from tests/test_litellm/llms/openai/responses/test_openai_responses_data_residency.py rename to tests/unit/llms/openai/responses/test_openai_responses_data_residency.py diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py similarity index 100% rename from tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py rename to tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_tool_merge.py b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_tool_merge.py similarity index 100% rename from tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_tool_merge.py rename to tests/unit/llms/openai/responses/test_openai_responses_guardrail_tool_merge.py diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py b/tests/unit/llms/openai/responses/test_openai_responses_transformation.py similarity index 100% rename from tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py rename to tests/unit/llms/openai/responses/test_openai_responses_transformation.py diff --git a/tests/test_litellm/llms/openai/test_cost_calculation.py b/tests/unit/llms/openai/test_cost_calculation.py similarity index 100% rename from tests/test_litellm/llms/openai/test_cost_calculation.py rename to tests/unit/llms/openai/test_cost_calculation.py diff --git a/tests/test_litellm/llms/openai/test_data_residency.py b/tests/unit/llms/openai/test_data_residency.py similarity index 100% rename from tests/test_litellm/llms/openai/test_data_residency.py rename to tests/unit/llms/openai/test_data_residency.py diff --git a/tests/test_litellm/llms/openai/test_gpt5_transformation.py b/tests/unit/llms/openai/test_gpt5_transformation.py similarity index 100% rename from tests/test_litellm/llms/openai/test_gpt5_transformation.py rename to tests/unit/llms/openai/test_gpt5_transformation.py diff --git a/tests/test_litellm/llms/openai/test_is_model_gpt_5_model.py b/tests/unit/llms/openai/test_is_model_gpt_5_model.py similarity index 100% rename from tests/test_litellm/llms/openai/test_is_model_gpt_5_model.py rename to tests/unit/llms/openai/test_is_model_gpt_5_model.py diff --git a/tests/test_litellm/llms/openai/test_o_series_transformation.py b/tests/unit/llms/openai/test_o_series_transformation.py similarity index 100% rename from tests/test_litellm/llms/openai/test_o_series_transformation.py rename to tests/unit/llms/openai/test_o_series_transformation.py diff --git a/tests/test_litellm/llms/openai/test_openai.py b/tests/unit/llms/openai/test_openai.py similarity index 100% rename from tests/test_litellm/llms/openai/test_openai.py rename to tests/unit/llms/openai/test_openai.py diff --git a/tests/test_litellm/llms/openai/test_openai_common_utils.py b/tests/unit/llms/openai/test_openai_common_utils.py similarity index 100% rename from tests/test_litellm/llms/openai/test_openai_common_utils.py rename to tests/unit/llms/openai/test_openai_common_utils.py diff --git a/tests/test_litellm/llms/openai/test_openai_empty_response.py b/tests/unit/llms/openai/test_openai_empty_response.py similarity index 100% rename from tests/test_litellm/llms/openai/test_openai_empty_response.py rename to tests/unit/llms/openai/test_openai_empty_response.py diff --git a/tests/test_litellm/llms/openai/test_openai_file_content_streaming.py b/tests/unit/llms/openai/test_openai_file_content_streaming.py similarity index 100% rename from tests/test_litellm/llms/openai/test_openai_file_content_streaming.py rename to tests/unit/llms/openai/test_openai_file_content_streaming.py diff --git a/tests/test_litellm/llms/openai/test_openai_image_edit_transformation.py b/tests/unit/llms/openai/test_openai_image_edit_transformation.py similarity index 100% rename from tests/test_litellm/llms/openai/test_openai_image_edit_transformation.py rename to tests/unit/llms/openai/test_openai_image_edit_transformation.py diff --git a/tests/test_litellm/llms/openai/test_openai_workload_identity.py b/tests/unit/llms/openai/test_openai_workload_identity.py similarity index 100% rename from tests/test_litellm/llms/openai/test_openai_workload_identity.py rename to tests/unit/llms/openai/test_openai_workload_identity.py diff --git a/tests/test_litellm/llms/openai/test_organization_costs.py b/tests/unit/llms/openai/test_organization_costs.py similarity index 100% rename from tests/test_litellm/llms/openai/test_organization_costs.py rename to tests/unit/llms/openai/test_organization_costs.py diff --git a/tests/test_litellm/llms/openai/test_use_chat_completions_api_no_leak.py b/tests/unit/llms/openai/test_use_chat_completions_api_no_leak.py similarity index 100% rename from tests/test_litellm/llms/openai/test_use_chat_completions_api_no_leak.py rename to tests/unit/llms/openai/test_use_chat_completions_api_no_leak.py diff --git a/tests/test_litellm/llms/openai/transcriptions/test_openai_transcriptions_handler.py b/tests/unit/llms/openai/transcriptions/test_openai_transcriptions_handler.py similarity index 100% rename from tests/test_litellm/llms/openai/transcriptions/test_openai_transcriptions_handler.py rename to tests/unit/llms/openai/transcriptions/test_openai_transcriptions_handler.py diff --git a/tests/test_litellm/llms/vertex_ai/agent_engine/__init__.py b/tests/unit/llms/openai_like/responses/__init__.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/agent_engine/__init__.py rename to tests/unit/llms/openai_like/responses/__init__.py diff --git a/tests/test_litellm/llms/openai_like/responses/test_openai_like_responses.py b/tests/unit/llms/openai_like/responses/test_openai_like_responses.py similarity index 100% rename from tests/test_litellm/llms/openai_like/responses/test_openai_like_responses.py rename to tests/unit/llms/openai_like/responses/test_openai_like_responses.py diff --git a/tests/test_litellm/llms/openai_like/test_abliteration_provider.py b/tests/unit/llms/openai_like/test_abliteration_provider.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_abliteration_provider.py rename to tests/unit/llms/openai_like/test_abliteration_provider.py diff --git a/tests/test_litellm/llms/openai_like/test_assemblyai_provider.py b/tests/unit/llms/openai_like/test_assemblyai_provider.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_assemblyai_provider.py rename to tests/unit/llms/openai_like/test_assemblyai_provider.py diff --git a/tests/test_litellm/llms/openai_like/test_charity_engine.py b/tests/unit/llms/openai_like/test_charity_engine.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_charity_engine.py rename to tests/unit/llms/openai_like/test_charity_engine.py diff --git a/tests/test_litellm/llms/openai_like/test_cognition_provider.py b/tests/unit/llms/openai_like/test_cognition_provider.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_cognition_provider.py rename to tests/unit/llms/openai_like/test_cognition_provider.py diff --git a/tests/test_litellm/llms/openai_like/test_dynamic_config.py b/tests/unit/llms/openai_like/test_dynamic_config.py similarity index 96% rename from tests/test_litellm/llms/openai_like/test_dynamic_config.py rename to tests/unit/llms/openai_like/test_dynamic_config.py index 55e1a1679de..de70f98c3f1 100644 --- a/tests/test_litellm/llms/openai_like/test_dynamic_config.py +++ b/tests/unit/llms/openai_like/test_dynamic_config.py @@ -20,9 +20,6 @@ def _isolate_generated_class_cache(): class TestClassCaching: - def test_same_slug_returns_the_identical_class_object(self): - provider = _provider("cache_same_slug") - assert create_responses_config_class(provider) is create_responses_config_class(provider) def test_cache_is_keyed_on_slug_not_on_the_provider_instance(self): first = create_responses_config_class(_provider("cache_by_slug")) diff --git a/tests/test_litellm/llms/openai_like/test_empiriolabs_provider.py b/tests/unit/llms/openai_like/test_empiriolabs_provider.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_empiriolabs_provider.py rename to tests/unit/llms/openai_like/test_empiriolabs_provider.py diff --git a/tests/unit/llms/openai_like/test_json_providers.py b/tests/unit/llms/openai_like/test_json_providers.py new file mode 100644 index 00000000000..a56108ca9ac --- /dev/null +++ b/tests/unit/llms/openai_like/test_json_providers.py @@ -0,0 +1,317 @@ +""" +Tests for JSON-based provider configuration system. +""" + +import os +import sys +from unittest.mock import patch + +try: + import pytest +except ImportError: + # pytest not available, will run as standalone script + pytest = None + +# Add workspace to path +workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../..")) +sys.path.insert(0, workspace_path) + + + +class TestJSONProviderLoader: + """Test JSON provider loading and configuration""" + + def test_load_json_providers(self): + """Test that JSON providers load correctly""" + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + # Verify publicai is loaded + assert JSONProviderRegistry.exists("publicai") + + # Get publicai config + publicai = JSONProviderRegistry.get("publicai") + assert publicai is not None + assert publicai.base_url == "https://api.publicai.co/v1" + assert publicai.api_key_env == "PUBLICAI_API_KEY" + assert publicai.api_base_env == "PUBLICAI_API_BASE" + assert publicai.param_mappings.get("max_completion_tokens") == "max_tokens" + + def test_dynamic_config_generation(self): + """Test dynamic config class creation""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + provider = JSONProviderRegistry.get("publicai") + config_class = create_config_class(provider) + config = config_class() + + # Test API info resolution + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + assert api_base == "https://api.publicai.co/v1" + + # Test with custom base + api_base, api_key = config._get_openai_compatible_provider_info( + "https://custom.api.com", "test-key" + ) + assert api_base == "https://custom.api.com" + assert api_key == "test-key" + + def test_parameter_mapping(self): + """Test parameter mapping works""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + provider = JSONProviderRegistry.get("publicai") + config_class = create_config_class(provider) + config = config_class() + + # Test parameter mapping + optional_params = {} + non_default_params = {"max_completion_tokens": 100, "temperature": 0.7} + result = config.map_openai_params( + non_default_params, optional_params, "gpt-4", False + ) + + # max_completion_tokens should be mapped to max_tokens + assert "max_tokens" in result + assert result["max_tokens"] == 100 + assert "max_completion_tokens" not in result + + # temperature should be passed through + assert result["temperature"] == 0.7 + + def test_supported_params(self): + """Test that config returns supported params""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + provider = JSONProviderRegistry.get("publicai") + config_class = create_config_class(provider) + config = config_class() + + # Get supported params + supported = config.get_supported_openai_params("gpt-4") + + # Should have standard OpenAI params + assert isinstance(supported, list) + assert len(supported) > 0 + + def test_tool_params_excluded_when_function_calling_not_supported(self): + """Test that tool-related params are excluded for models that don't support + function calling. Regression test for https://github.com/BerriAI/litellm/issues/21125 + """ + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + provider = JSONProviderRegistry.get("publicai") + config_class = create_config_class(provider) + config = config_class() + + # Mock supports_function_calling to return False + with patch("litellm.utils.supports_function_calling", return_value=False): + supported = config.get_supported_openai_params("some-model-without-fc") + + tool_params = [ + "tools", + "tool_choice", + "function_call", + "functions", + "parallel_tool_calls", + ] + for param in tool_params: + assert ( + param not in supported + ), f"'{param}' should not be in supported params when function calling is not supported" + + # Non-tool params should still be present + assert "temperature" in supported + assert "max_tokens" in supported + assert "stop" in supported + + def test_tool_params_included_when_function_calling_supported(self): + """Test that tool-related params are included for models that support function calling.""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + provider = JSONProviderRegistry.get("publicai") + config_class = create_config_class(provider) + config = config_class() + + # Mock supports_function_calling to return True + with patch("litellm.utils.supports_function_calling", return_value=True): + supported = config.get_supported_openai_params("some-model-with-fc") + + assert "tools" in supported + assert "tool_choice" in supported + + def test_provider_resolution(self): + """Test that provider resolution finds JSON providers""" + from litellm.litellm_core_utils.get_llm_provider_logic import ( + get_llm_provider, + ) + + model, provider, api_key, api_base = get_llm_provider( + model="publicai/gpt-4", + custom_llm_provider=None, + api_base=None, + api_key=None, + ) + + assert model == "gpt-4" + assert provider == "publicai" + assert api_base == "https://api.publicai.co/v1" + + def test_provider_config_manager(self): + """Test that ProviderConfigManager returns JSON-based configs""" + from litellm import LlmProviders + from litellm.utils import ProviderConfigManager + + config = ProviderConfigManager.get_provider_chat_config( + model="gpt-4", provider=LlmProviders.PUBLICAI + ) + + assert config is not None + assert config.custom_llm_provider == "publicai" + + +class TestPinstripes: + """Tests for Pinstripes JSON-configured provider""" + + def test_pinstripes_json_config_exists(self): + """Test that pinstripes is configured in providers.json""" + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + assert JSONProviderRegistry.exists("pinstripes") + + pinstripes = JSONProviderRegistry.get("pinstripes") + assert pinstripes is not None + assert pinstripes.base_url == "https://pinstripes.io/v1" + assert pinstripes.api_key_env == "PINSTRIPES_API_KEY" + assert pinstripes.param_mappings.get("max_completion_tokens") == "max_tokens" + + def test_pinstripes_provider_resolution(self): + """Test that provider resolution finds pinstripes and returns the default base URL""" + from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + model, provider, api_key, api_base = get_llm_provider( + model="pinstripes/ps/glm-4.5-air", + custom_llm_provider=None, + api_base=None, + api_key=None, + ) + + assert model == "ps/glm-4.5-air" + assert provider == "pinstripes" + assert api_base == "https://pinstripes.io/v1" + + def test_pinstripes_dynamic_config(self): + """Test dynamic config class creation for pinstripes""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + provider = JSONProviderRegistry.get("pinstripes") + config_class = create_config_class(provider) + config = config_class() + + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + assert api_base == "https://pinstripes.io/v1" + + api_base, api_key = config._get_openai_compatible_provider_info( + "https://custom.pinstripes.io/v1", "test-key" + ) + assert api_base == "https://custom.pinstripes.io/v1" + assert api_key == "test-key" + + def test_pinstripes_parameter_mapping(self): + """Test that max_completion_tokens is mapped to max_tokens for pinstripes""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + provider = JSONProviderRegistry.get("pinstripes") + config_class = create_config_class(provider) + config = config_class() + + optional_params = {} + non_default_params = {"max_completion_tokens": 100, "temperature": 0.7} + result = config.map_openai_params( + non_default_params, optional_params, "ps/glm-4.5-air", False + ) + + assert "max_tokens" in result + assert result["max_tokens"] == 100 + assert "max_completion_tokens" not in result + assert result["temperature"] == 0.7 + + +class TestDarkbloom: + def test_darkbloom_json_config_exists(self): + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + darkbloom = JSONProviderRegistry.get("darkbloom") + assert darkbloom is not None + assert darkbloom.base_url == "https://api.darkbloom.dev/v1" + assert darkbloom.api_key_env == "DARKBLOOM_API_KEY" + assert darkbloom.api_base_env == "DARKBLOOM_API_BASE" + assert darkbloom.param_mappings.get("max_completion_tokens") == "max_tokens" + + def test_darkbloom_provider_resolution(self): + from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + model, provider, api_key, api_base = get_llm_provider( + model="darkbloom/gemma-4-26b", + custom_llm_provider=None, + api_base=None, + api_key=None, + ) + + assert model == "gemma-4-26b" + assert provider == "darkbloom" + assert api_key is None + assert api_base == "https://api.darkbloom.dev/v1" + + def test_darkbloom_dynamic_config(self): + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + provider = JSONProviderRegistry.get("darkbloom") + config_class = create_config_class(provider) + config = config_class() + + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + assert api_base == "https://api.darkbloom.dev/v1" + + api_base, api_key = config._get_openai_compatible_provider_info( + "https://custom.darkbloom.dev/v1", "test-key" + ) + assert api_base == "https://custom.darkbloom.dev/v1" + assert api_key == "test-key" + + def test_darkbloom_complete_url_appends_endpoint(self): + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + provider = JSONProviderRegistry.get("darkbloom") + config_class = create_config_class(provider) + config = config_class() + + url = config.get_complete_url( + api_base="https://api.darkbloom.dev/v1", + api_key="test-key", + model="darkbloom/gemma-4-26b", + optional_params={}, + litellm_params={}, + stream=True, + ) + + assert url == "https://api.darkbloom.dev/v1/chat/completions" + + def test_darkbloom_provider_config_manager(self): + from litellm import LlmProviders + from litellm.utils import ProviderConfigManager + + config = ProviderConfigManager.get_provider_chat_config( + model="gemma-4-26b", provider=LlmProviders.DARKBLOOM + ) + + assert config is not None + assert config.custom_llm_provider == "darkbloom" diff --git a/tests/test_litellm/llms/openai_like/test_libertai_provider.py b/tests/unit/llms/openai_like/test_libertai_provider.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_libertai_provider.py rename to tests/unit/llms/openai_like/test_libertai_provider.py diff --git a/tests/test_litellm/llms/openai_like/test_meta_provider.py b/tests/unit/llms/openai_like/test_meta_provider.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_meta_provider.py rename to tests/unit/llms/openai_like/test_meta_provider.py diff --git a/tests/test_litellm/llms/openai_like/test_model_info.py b/tests/unit/llms/openai_like/test_model_info.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_model_info.py rename to tests/unit/llms/openai_like/test_model_info.py diff --git a/tests/test_litellm/llms/openai_like/test_pinstripes_provider.py b/tests/unit/llms/openai_like/test_pinstripes_provider.py similarity index 68% rename from tests/test_litellm/llms/openai_like/test_pinstripes_provider.py rename to tests/unit/llms/openai_like/test_pinstripes_provider.py index 70bb786b2e6..e7a2dfb92dc 100644 --- a/tests/test_litellm/llms/openai_like/test_pinstripes_provider.py +++ b/tests/unit/llms/openai_like/test_pinstripes_provider.py @@ -16,17 +16,6 @@ class TestPinstripeProviderConfig: assert LlmProviders.PINSTRIPES.value == "pinstripes" assert "pinstripes" in litellm.provider_list - def test_pinstripes_json_config_exists(self): - """Test that pinstripes is configured in providers.json""" - from litellm.llms.openai_like.json_loader import JSONProviderRegistry - - assert JSONProviderRegistry.exists("pinstripes") - - pinstripes = JSONProviderRegistry.get("pinstripes") - assert pinstripes is not None - assert pinstripes.base_url == "https://pinstripes.io/v1" - assert pinstripes.api_key_env == "PINSTRIPES_API_KEY" - assert pinstripes.param_mappings.get("max_completion_tokens") == "max_tokens" def test_pinstripes_in_openai_compatible_providers(self): """Test that pinstripes is in the openai_compatible_providers list""" @@ -34,20 +23,6 @@ class TestPinstripeProviderConfig: assert "pinstripes" in openai_compatible_providers - def test_pinstripes_provider_resolution(self): - """Test that provider resolution finds pinstripes and returns the default base URL""" - from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider - - model, provider, api_key, api_base = get_llm_provider( - model="pinstripes/ps/glm-4.5-air", - custom_llm_provider=None, - api_base=None, - api_key=None, - ) - - assert model == "ps/glm-4.5-air" - assert provider == "pinstripes" - assert api_base == "https://pinstripes.io/v1" def test_pinstripes_api_base_override(self): """Test that an explicit api_base / api_key overrides the default""" diff --git a/tests/test_litellm/llms/openai_like/test_provider_affinity_forwarding.py b/tests/unit/llms/openai_like/test_provider_affinity_forwarding.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_provider_affinity_forwarding.py rename to tests/unit/llms/openai_like/test_provider_affinity_forwarding.py diff --git a/tests/test_litellm/llms/openai_like/test_scx_ai_provider.py b/tests/unit/llms/openai_like/test_scx_ai_provider.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_scx_ai_provider.py rename to tests/unit/llms/openai_like/test_scx_ai_provider.py diff --git a/tests/test_litellm/llms/openai_like/test_tensormesh_provider.py b/tests/unit/llms/openai_like/test_tensormesh_provider.py similarity index 100% rename from tests/test_litellm/llms/openai_like/test_tensormesh_provider.py rename to tests/unit/llms/openai_like/test_tensormesh_provider.py diff --git a/tests/unit/llms/openai_like/test_xiaomi_mimo.py b/tests/unit/llms/openai_like/test_xiaomi_mimo.py new file mode 100644 index 00000000000..a642cc91f90 --- /dev/null +++ b/tests/unit/llms/openai_like/test_xiaomi_mimo.py @@ -0,0 +1,84 @@ +""" +Tests for Xiaomi MiMo provider configuration and integration. +Related to issue #18794 +""" + +import os +import sys +from unittest.mock import MagicMock, patch + +try: + import pytest +except ImportError: + pytest = None + +# Add workspace to path +workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../..")) +sys.path.insert(0, workspace_path) + +import litellm + + +class TestXiaomiMiMoProviderConfig: + """Test Xiaomi MiMo provider configuration""" + + def test_xiaomi_mimo_in_provider_list(self): + """Test that xiaomi_mimo is in the provider list (fixes #18794)""" + from litellm import LlmProviders + + # Verify xiaomi_mimo is in the enum + assert hasattr(LlmProviders, "XIAOMI_MIMO") + assert LlmProviders.XIAOMI_MIMO.value == "xiaomi_mimo" + + # Verify it's in the provider list + assert "xiaomi_mimo" in litellm.provider_list + + def test_xiaomi_mimo_json_config_exists(self): + """Test that xiaomi_mimo is configured in providers.json""" + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + # Verify xiaomi_mimo is loaded + assert JSONProviderRegistry.exists("xiaomi_mimo") + + # Get xiaomi_mimo config + xiaomi_mimo = JSONProviderRegistry.get("xiaomi_mimo") + assert xiaomi_mimo is not None + assert xiaomi_mimo.base_url == "https://api.xiaomimimo.com/v1" + assert xiaomi_mimo.api_key_env == "XIAOMI_MIMO_API_KEY" + assert xiaomi_mimo.param_mappings.get("max_completion_tokens") == "max_tokens" + + def test_xiaomi_mimo_provider_resolution(self): + """Test that provider resolution finds xiaomi_mimo""" + from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + model, provider, api_key, api_base = get_llm_provider( + model="xiaomi_mimo/mimo-v2-flash", + custom_llm_provider=None, + api_base=None, + api_key=None, + ) + + assert model == "mimo-v2-flash" + assert provider == "xiaomi_mimo" + assert api_base == "https://api.xiaomimimo.com/v1" + + def test_xiaomi_mimo_router_config(self): + """Test that xiaomi_mimo can be used in Router configuration (fixes #18794)""" + from litellm import Router + + # This should not raise "Unsupported provider - xiaomi_mimo" + router = Router( + model_list=[ + { + "model_name": "mimo-v2-flash", + "litellm_params": { + "model": "xiaomi_mimo/mimo-v2-flash", + "api_key": "test-key", + }, + } + ] + ) + + # Verify the deployment was created successfully + assert len(router.model_list) == 1 + assert router.model_list[0]["model_name"] == "mimo-v2-flash" diff --git a/tests/test_litellm/llms/vertex_ai/audio_transcription/__init__.py b/tests/unit/llms/ovhcloud/__init__.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/audio_transcription/__init__.py rename to tests/unit/llms/ovhcloud/__init__.py diff --git a/tests/unit/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py b/tests/unit/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py new file mode 100644 index 00000000000..87e54dfba9b --- /dev/null +++ b/tests/unit/llms/ovhcloud/test_ovhcloud_audio_transcription_transformation.py @@ -0,0 +1,58 @@ + + + + +class TestOVHCloudDurationFieldMigration: + """Tests for OVHCloud duration -> seconds field migration.""" + + def test_seconds_field_mapped_to_duration(self): + """New `seconds` field should be normalized to `duration`.""" + from litellm.llms.ovhcloud.audio_transcription.transformation import ( + OVHCloudAudioTranscriptionConfig, + ) + from unittest.mock import MagicMock + + config = OVHCloudAudioTranscriptionConfig() + mock_response = MagicMock() + mock_response.json.return_value = { + "text": "Hello world", + "seconds": 3.14, + } + + result = config.transform_audio_transcription_response(mock_response) + + assert result.text == "Hello world" + assert result._hidden_params["duration"] == 3.14 + + def test_legacy_duration_field_still_works(self): + """Legacy `duration` field should still be accepted.""" + from litellm.llms.ovhcloud.audio_transcription.transformation import ( + OVHCloudAudioTranscriptionConfig, + ) + from unittest.mock import MagicMock + + config = OVHCloudAudioTranscriptionConfig() + mock_response = MagicMock() + mock_response.json.return_value = { + "text": "Hello world", + "duration": 2.71, + } + + result = config.transform_audio_transcription_response(mock_response) + + assert result.text == "Hello world" + assert result._hidden_params["duration"] == 2.71 + + + def test_seconds_zero_mapped_to_duration(self): + """seconds=0.0 must not be treated as falsy and lost.""" + from litellm.llms.ovhcloud.audio_transcription.transformation import ( + OVHCloudAudioTranscriptionConfig, + ) + from unittest.mock import MagicMock + + config = OVHCloudAudioTranscriptionConfig() + mock_response = MagicMock() + 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 diff --git a/tests/unit/llms/ovhcloud/test_ovhcloud_chat_transformation.py b/tests/unit/llms/ovhcloud/test_ovhcloud_chat_transformation.py new file mode 100644 index 00000000000..c2bc4ee4a4c --- /dev/null +++ b/tests/unit/llms/ovhcloud/test_ovhcloud_chat_transformation.py @@ -0,0 +1,250 @@ +""" +Unit tests for OVHCloud AI Endpoints chat integration. +""" + + +import pytest + +from litellm.llms.ovhcloud.utils import OVHCloudException +from litellm.utils import get_optional_params + + +from litellm.llms.ovhcloud.chat.transformation import ( + OVHCloudChatCompletionStreamingHandler, + OVHCloudChatConfig, +) + +config = OVHCloudChatConfig() +model = "ovhcloud/Mistral-7B-Instruct-v0.3" + + +class TestOvhCloudChatCompletionStreamingHandler: + def test_chunk_parser_successful(self): + handler = OVHCloudChatCompletionStreamingHandler( + streaming_response=None, sync_stream=True + ) + + chunk = { + "id": "test_id", + "created": 1234567890, + "model": "gpt-oss-20b", + "usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}, + "choices": [ + {"delta": {"content": "test content", "reasoning": "test reasoning"}} + ], + } + + result = handler.chunk_parser(chunk) + + assert result.id == "test_id" + assert result.object == "chat.completion.chunk" + assert result.created == 1234567890 + assert result.model == "gpt-oss-20b" + assert result.usage.prompt_tokens == chunk["usage"]["prompt_tokens"] + assert result.usage.completion_tokens == chunk["usage"]["completion_tokens"] + assert result.usage.total_tokens == chunk["usage"]["total_tokens"] + assert len(result.choices) == 1 + assert result.choices[0]["delta"]["reasoning_content"] == "test reasoning" + + def test_chunk_parser_error_response(self): + handler = OVHCloudChatCompletionStreamingHandler( + streaming_response=None, sync_stream=True + ) + + error_chunk = { + "error": { + "message": "test error", + "code": 400, + } + } + + with pytest.raises(OVHCloudException) as exc_info: + handler.chunk_parser(error_chunk) + + assert "OVHCloud Error: test error" in str(exc_info.value) + assert exc_info.value.status_code == 400 + + def test_chunk_parser_key_error(self): + handler = OVHCloudChatCompletionStreamingHandler( + streaming_response=None, sync_stream=True + ) + + invalid_chunk = {"incomplete": "data"} + + with pytest.raises(OVHCloudException) as exc_info: + handler.chunk_parser(invalid_chunk) + + assert "KeyError" in str(exc_info.value) + assert exc_info.value.status_code == 400 + + +class TestOVHCloudConfig: + def test_transform_request_basic(self): + """Test basic request transformation""" + transformed_request = config.transform_request( + model, + messages=[{"role": "user", "content": "Hello, world!"}], + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert transformed_request["model"] == model + assert transformed_request["messages"] == [ + {"role": "user", "content": "Hello, world!"} + ] + + def test_transform_request_with_extra_body(self): + """Test request transformation with extra_body parameters""" + transformed_request = config.transform_request( + model, + messages=[{"role": "user", "content": "Hello, world!"}], + optional_params={"extra_body": {"custom_param": "custom_value"}}, + litellm_params={}, + headers={}, + ) + + assert transformed_request["custom_param"] == "custom_value" + assert transformed_request["messages"] == [ + {"role": "user", "content": "Hello, world!"} + ] + + def test_map_openai_params(self): + """Test OpenAI parameter mapping""" + non_default_params = { + "temperature": 0.7, + "max_tokens": 100, + "top_p": 0.9, + } + + mapped_params = config.map_openai_params( + non_default_params=non_default_params, + optional_params={}, + model=model, + drop_params=False, + ) + + assert mapped_params["temperature"] == 0.7 + assert mapped_params["max_tokens"] == 100 + assert mapped_params["top_p"] == 0.9 + + def test_get_error_class(self): + """Test error class creation""" + error = config.get_error_class( + error_message="Test error", + status_code=400, + headers={"Content-Type": "application/json"}, + ) + + assert isinstance(error, OVHCloudException) + assert error.message == "Test error" + assert error.status_code == 400 + + @pytest.mark.parametrize( + "model", + [ + "Meta-Llama-3_3-70B-Instruct", + "Meta-Llama-3_1-70B-Instruct", + "Mixtral-8x7B-Instruct-v0.1", + "gpt-oss-120b", + "some-model-not-in-the-cost-map", + ], + ) + def test_tools_not_filtered_by_static_model_map(self, model): + """ + OVHCloud AI Endpoints are OpenAI-compatible; tools/tool_choice must pass + through for any model. The server is responsible for rejecting unsupported + tool calls — LiteLLM must not strip them based on a stale static catalog. + """ + + params = get_optional_params( + model=model, + custom_llm_provider="ovhcloud", + tools=[ + { + "type": "function", + "function": {"name": "x", "parameters": {}}, + } + ], + tool_choice="auto", + ) + + assert "tools" in params + assert "tool_choice" in params + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) + + +class TestOVHCloudReasoningFieldMigration: + """Tests for OVHCloud reasoning_content -> reasoning field migration.""" + + def test_streaming_new_reasoning_field(self): + """New `reasoning` field should be mapped to `reasoning_content`.""" + handler = OVHCloudChatCompletionStreamingHandler( + streaming_response=iter([]), + sync_stream=True, + ) + chunk = { + "id": "test-id", + "created": 1234567890, + "model": "test-model", + "choices": [ + { + "delta": { + "role": "assistant", + "reasoning": "Let me think...", + }, + "index": 0, + } + ], + } + result = handler.chunk_parser(chunk) + assert result.choices[0]["delta"]["reasoning_content"] == "Let me think..." + + def test_streaming_legacy_reasoning_content_unchanged(self): + """Legacy `reasoning_content` field should pass through untouched.""" + handler = OVHCloudChatCompletionStreamingHandler( + streaming_response=iter([]), + sync_stream=True, + ) + chunk = { + "id": "test-id", + "created": 1234567890, + "model": "test-model", + "choices": [ + { + "delta": { + "role": "assistant", + "reasoning_content": "Already correct field.", + }, + "index": 0, + } + ], + } + result = handler.chunk_parser(chunk) + assert result.choices[0]["delta"]["reasoning_content"] == "Already correct field." + + def test_streaming_both_fields_legacy_wins(self): + """When both fields present, existing `reasoning_content` is not overwritten.""" + handler = OVHCloudChatCompletionStreamingHandler( + streaming_response=iter([]), + sync_stream=True, + ) + chunk = { + "id": "test-id", + "created": 1234567890, + "model": "test-model", + "choices": [ + { + "delta": { + "reasoning": "new field", + "reasoning_content": "legacy field", + }, + "index": 0, + } + ], + } + result = handler.chunk_parser(chunk) + assert result.choices[0]["delta"]["reasoning_content"] == "legacy field" diff --git a/tests/test_litellm/llms/ovhcloud/test_ovhcloud_embeddings_transformation.py b/tests/unit/llms/ovhcloud/test_ovhcloud_embeddings_transformation.py similarity index 100% rename from tests/test_litellm/llms/ovhcloud/test_ovhcloud_embeddings_transformation.py rename to tests/unit/llms/ovhcloud/test_ovhcloud_embeddings_transformation.py diff --git a/tests/test_litellm/llms/vertex_ai/batches/__init__.py b/tests/unit/llms/pass_through/__init__.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/batches/__init__.py rename to tests/unit/llms/pass_through/__init__.py diff --git a/tests/test_litellm/llms/vertex_ai/files/__init__.py b/tests/unit/llms/pass_through/guardrail_translation/__init__.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/files/__init__.py rename to tests/unit/llms/pass_through/guardrail_translation/__init__.py diff --git a/tests/test_litellm/llms/perplexity/test_perplexity.py b/tests/unit/llms/perplexity/test_perplexity.py similarity index 100% rename from tests/test_litellm/llms/perplexity/test_perplexity.py rename to tests/unit/llms/perplexity/test_perplexity.py diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py b/tests/unit/llms/perplexity/test_perplexity_cost_calculator.py similarity index 100% rename from tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py rename to tests/unit/llms/perplexity/test_perplexity_cost_calculator.py diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_integration.py b/tests/unit/llms/perplexity/test_perplexity_integration.py similarity index 100% rename from tests/test_litellm/llms/perplexity/test_perplexity_integration.py rename to tests/unit/llms/perplexity/test_perplexity_integration.py diff --git a/tests/test_litellm/llms/vertex_ai/gemini_embeddings/__init__.py b/tests/unit/llms/pg_vector/__init__.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/gemini_embeddings/__init__.py rename to tests/unit/llms/pg_vector/__init__.py diff --git a/tests/test_litellm/llms/vertex_ai/text_to_speech/__init__.py b/tests/unit/llms/pg_vector/vector_stores/__init__.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/text_to_speech/__init__.py rename to tests/unit/llms/pg_vector/vector_stores/__init__.py diff --git a/tests/test_litellm/llms/pg_vector/vector_stores/test_pg_vector_transformation.py b/tests/unit/llms/pg_vector/vector_stores/test_pg_vector_transformation.py similarity index 100% rename from tests/test_litellm/llms/pg_vector/vector_stores/test_pg_vector_transformation.py rename to tests/unit/llms/pg_vector/vector_stores/test_pg_vector_transformation.py diff --git a/tests/test_litellm/llms/azure/realtime/__init__.py b/tests/unit/llms/reducto/__init__.py similarity index 100% rename from tests/test_litellm/llms/azure/realtime/__init__.py rename to tests/unit/llms/reducto/__init__.py diff --git a/tests/test_litellm/llms/reducto/conftest.py b/tests/unit/llms/reducto/conftest.py similarity index 100% rename from tests/test_litellm/llms/reducto/conftest.py rename to tests/unit/llms/reducto/conftest.py diff --git a/tests/test_litellm/llms/reducto/test_cost.py b/tests/unit/llms/reducto/test_cost.py similarity index 100% rename from tests/test_litellm/llms/reducto/test_cost.py rename to tests/unit/llms/reducto/test_cost.py diff --git a/tests/test_litellm/llms/reducto/test_model_info.py b/tests/unit/llms/reducto/test_model_info.py similarity index 100% rename from tests/test_litellm/llms/reducto/test_model_info.py rename to tests/unit/llms/reducto/test_model_info.py diff --git a/tests/test_litellm/llms/reducto/test_parse_legacy.py b/tests/unit/llms/reducto/test_parse_legacy.py similarity index 100% rename from tests/test_litellm/llms/reducto/test_parse_legacy.py rename to tests/unit/llms/reducto/test_parse_legacy.py diff --git a/tests/test_litellm/llms/reducto/test_parse_v3.py b/tests/unit/llms/reducto/test_parse_v3.py similarity index 100% rename from tests/test_litellm/llms/reducto/test_parse_v3.py rename to tests/unit/llms/reducto/test_parse_v3.py diff --git a/tests/test_litellm/llms/reducto/test_upload.py b/tests/unit/llms/reducto/test_upload.py similarity index 100% rename from tests/test_litellm/llms/reducto/test_upload.py rename to tests/unit/llms/reducto/test_upload.py diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/__init__.py b/tests/unit/llms/sagemaker/__init__.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/__init__.py rename to tests/unit/llms/sagemaker/__init__.py diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_chat_handler.py b/tests/unit/llms/sagemaker/test_sagemaker_chat_handler.py similarity index 100% rename from tests/test_litellm/llms/sagemaker/test_sagemaker_chat_handler.py rename to tests/unit/llms/sagemaker/test_sagemaker_chat_handler.py diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_chat_transformation.py b/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/sagemaker/test_sagemaker_chat_transformation.py rename to tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py b/tests/unit/llms/sagemaker/test_sagemaker_common_utils.py similarity index 100% rename from tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py rename to tests/unit/llms/sagemaker/test_sagemaker_common_utils.py diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_completion_handler.py b/tests/unit/llms/sagemaker/test_sagemaker_completion_handler.py similarity index 100% rename from tests/test_litellm/llms/sagemaker/test_sagemaker_completion_handler.py rename to tests/unit/llms/sagemaker/test_sagemaker_completion_handler.py diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_embedding_role_assumption.py b/tests/unit/llms/sagemaker/test_sagemaker_embedding_role_assumption.py similarity index 100% rename from tests/test_litellm/llms/sagemaker/test_sagemaker_embedding_role_assumption.py rename to tests/unit/llms/sagemaker/test_sagemaker_embedding_role_assumption.py diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_embedding_voyage.py b/tests/unit/llms/sagemaker/test_sagemaker_embedding_voyage.py similarity index 100% rename from tests/test_litellm/llms/sagemaker/test_sagemaker_embedding_voyage.py rename to tests/unit/llms/sagemaker/test_sagemaker_embedding_voyage.py diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_nova_transformation.py b/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py similarity index 100% rename from tests/test_litellm/llms/sagemaker/test_sagemaker_nova_transformation.py rename to tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/__init__.py b/tests/unit/llms/sambanova/__init__.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/__init__.py rename to tests/unit/llms/sambanova/__init__.py diff --git a/tests/test_litellm/llms/sambanova/tests_sambanova_embedding_transformation.py b/tests/unit/llms/sambanova/tests_sambanova_embedding_transformation.py similarity index 100% rename from tests/test_litellm/llms/sambanova/tests_sambanova_embedding_transformation.py rename to tests/unit/llms/sambanova/tests_sambanova_embedding_transformation.py diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/qwen/__init__.py b/tests/unit/llms/sap/chat/__init__.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/qwen/__init__.py rename to tests/unit/llms/sap/chat/__init__.py diff --git a/tests/test_litellm/llms/sap/chat/test_sap_chat_calls.py b/tests/unit/llms/sap/chat/test_sap_chat_calls.py similarity index 100% rename from tests/test_litellm/llms/sap/chat/test_sap_chat_calls.py rename to tests/unit/llms/sap/chat/test_sap_chat_calls.py diff --git a/tests/test_litellm/llms/sap/chat/test_sap_langchain_strict_param.py b/tests/unit/llms/sap/chat/test_sap_langchain_strict_param.py similarity index 100% rename from tests/test_litellm/llms/sap/chat/test_sap_langchain_strict_param.py rename to tests/unit/llms/sap/chat/test_sap_langchain_strict_param.py diff --git a/tests/test_litellm/llms/sap/chat/test_sap_response_format.py b/tests/unit/llms/sap/chat/test_sap_response_format.py similarity index 100% rename from tests/test_litellm/llms/sap/chat/test_sap_response_format.py rename to tests/unit/llms/sap/chat/test_sap_response_format.py diff --git a/tests/test_litellm/llms/sap/chat/test_sap_tool_parameters.py b/tests/unit/llms/sap/chat/test_sap_tool_parameters.py similarity index 100% rename from tests/test_litellm/llms/sap/chat/test_sap_tool_parameters.py rename to tests/unit/llms/sap/chat/test_sap_tool_parameters.py diff --git a/tests/test_litellm/llms/sap/chat/test_sap_transformation.py b/tests/unit/llms/sap/chat/test_sap_transformation.py similarity index 100% rename from tests/test_litellm/llms/sap/chat/test_sap_transformation.py rename to tests/unit/llms/sap/chat/test_sap_transformation.py diff --git a/tests/test_litellm/llms/voyage/rerank/__init__.py b/tests/unit/llms/sap/embed/__init__.py similarity index 100% rename from tests/test_litellm/llms/voyage/rerank/__init__.py rename to tests/unit/llms/sap/embed/__init__.py diff --git a/tests/test_litellm/llms/sap/embed/test_sap_embed_transformation.py b/tests/unit/llms/sap/embed/test_sap_embed_transformation.py similarity index 100% rename from tests/test_litellm/llms/sap/embed/test_sap_embed_transformation.py rename to tests/unit/llms/sap/embed/test_sap_embed_transformation.py diff --git a/tests/test_litellm/llms/sap/embed/test_sap_embedding.py b/tests/unit/llms/sap/embed/test_sap_embedding.py similarity index 100% rename from tests/test_litellm/llms/sap/embed/test_sap_embedding.py rename to tests/unit/llms/sap/embed/test_sap_embedding.py diff --git a/tests/test_litellm/llms/watsonx/__init__.py b/tests/unit/llms/snowflake/chat/__init__.py similarity index 100% rename from tests/test_litellm/llms/watsonx/__init__.py rename to tests/unit/llms/snowflake/chat/__init__.py diff --git a/tests/test_litellm/llms/snowflake/chat/test_snowflake_chat_transformation.py b/tests/unit/llms/snowflake/chat/test_snowflake_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/snowflake/chat/test_snowflake_chat_transformation.py rename to tests/unit/llms/snowflake/chat/test_snowflake_chat_transformation.py diff --git a/tests/test_litellm/llms/watsonx/audio_transcription/__init__.py b/tests/unit/llms/snowflake/embedding/__init__.py similarity index 100% rename from tests/test_litellm/llms/watsonx/audio_transcription/__init__.py rename to tests/unit/llms/snowflake/embedding/__init__.py diff --git a/tests/test_litellm/llms/snowflake/embedding/test_snowflake_embedding.py b/tests/unit/llms/snowflake/embedding/test_snowflake_embedding.py similarity index 100% rename from tests/test_litellm/llms/snowflake/embedding/test_snowflake_embedding.py rename to tests/unit/llms/snowflake/embedding/test_snowflake_embedding.py diff --git a/tests/unit/llms/snowflake/test_snowflake_native_endpoints.py b/tests/unit/llms/snowflake/test_snowflake_native_endpoints.py index 7970f7771fc..344b8e5573d 100644 --- a/tests/unit/llms/snowflake/test_snowflake_native_endpoints.py +++ b/tests/unit/llms/snowflake/test_snowflake_native_endpoints.py @@ -7,7 +7,7 @@ Covers: - Claude models → /messages (Anthropic format) Run: - pytest tests/test_litellm/llms/snowflake/test_snowflake_native_endpoints.py -v + pytest tests/unit/llms/snowflake/test_snowflake_native_endpoints.py -v """ import json diff --git a/tests/test_litellm/llms/soniox/audio_transcription/__init__.py b/tests/unit/llms/soniox/audio_transcription/__init__.py similarity index 100% rename from tests/test_litellm/llms/soniox/audio_transcription/__init__.py rename to tests/unit/llms/soniox/audio_transcription/__init__.py diff --git a/tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_handler.py b/tests/unit/llms/soniox/audio_transcription/test_soniox_audio_transcription_handler.py similarity index 100% rename from tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_handler.py rename to tests/unit/llms/soniox/audio_transcription/test_soniox_audio_transcription_handler.py diff --git a/tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py b/tests/unit/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py similarity index 100% rename from tests/test_litellm/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py rename to tests/unit/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py diff --git a/tests/test_litellm/llms/test_cache_control_and_reasoning.py b/tests/unit/llms/test_cache_control_and_reasoning.py similarity index 100% rename from tests/test_litellm/llms/test_cache_control_and_reasoning.py rename to tests/unit/llms/test_cache_control_and_reasoning.py diff --git a/tests/test_litellm/llms/test_file_content_block.py b/tests/unit/llms/test_file_content_block.py similarity index 100% rename from tests/test_litellm/llms/test_file_content_block.py rename to tests/unit/llms/test_file_content_block.py diff --git a/tests/test_litellm/llms/test_file_search_responses.py b/tests/unit/llms/test_file_search_responses.py similarity index 100% rename from tests/test_litellm/llms/test_file_search_responses.py rename to tests/unit/llms/test_file_search_responses.py diff --git a/tests/test_litellm/llms/test_lifecycle_fix.py b/tests/unit/llms/test_lifecycle_fix.py similarity index 100% rename from tests/test_litellm/llms/test_lifecycle_fix.py rename to tests/unit/llms/test_lifecycle_fix.py diff --git a/tests/test_litellm/llms/test_polling_url_origin_match.py b/tests/unit/llms/test_polling_url_origin_match.py similarity index 100% rename from tests/test_litellm/llms/test_polling_url_origin_match.py rename to tests/unit/llms/test_polling_url_origin_match.py diff --git a/tests/test_litellm/llms/test_predibase_transformation.py b/tests/unit/llms/test_predibase_transformation.py similarity index 100% rename from tests/test_litellm/llms/test_predibase_transformation.py rename to tests/unit/llms/test_predibase_transformation.py diff --git a/tests/test_litellm/llms/watsonx/rerank/__init__.py b/tests/unit/llms/tinyfish/__init__.py similarity index 100% rename from tests/test_litellm/llms/watsonx/rerank/__init__.py rename to tests/unit/llms/tinyfish/__init__.py diff --git a/tests/test_litellm/llms/tinyfish/test_tinyfish_search.py b/tests/unit/llms/tinyfish/test_tinyfish_search.py similarity index 100% rename from tests/test_litellm/llms/tinyfish/test_tinyfish_search.py rename to tests/unit/llms/tinyfish/test_tinyfish_search.py diff --git a/tests/test_litellm/llms/vercel_ai_gateway/test_vercel_ai_gateway.py b/tests/unit/llms/vercel_ai_gateway/test_vercel_ai_gateway.py similarity index 100% rename from tests/test_litellm/llms/vercel_ai_gateway/test_vercel_ai_gateway.py rename to tests/unit/llms/vercel_ai_gateway/test_vercel_ai_gateway.py diff --git a/tests/test_litellm/llms/you_com/__init__.py b/tests/unit/llms/vertex_ai/audio_transcription/__init__.py similarity index 100% rename from tests/test_litellm/llms/you_com/__init__.py rename to tests/unit/llms/vertex_ai/audio_transcription/__init__.py diff --git a/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_audio_transcription_transformation.py b/tests/unit/llms/vertex_ai/audio_transcription/test_vertex_ai_audio_transcription_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_audio_transcription_transformation.py rename to tests/unit/llms/vertex_ai/audio_transcription/test_vertex_ai_audio_transcription_transformation.py diff --git a/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_gemini_transcribe_transformation.py b/tests/unit/llms/vertex_ai/audio_transcription/test_vertex_ai_gemini_transcribe_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_gemini_transcribe_transformation.py rename to tests/unit/llms/vertex_ai/audio_transcription/test_vertex_ai_gemini_transcribe_transformation.py diff --git a/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py b/tests/unit/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py rename to tests/unit/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py diff --git a/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_transformation.py b/tests/unit/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_transformation.py rename to tests/unit/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_transformation.py diff --git a/tests/test_litellm/messages/__init__.py b/tests/unit/llms/vertex_ai/batches/__init__.py similarity index 100% rename from tests/test_litellm/messages/__init__.py rename to tests/unit/llms/vertex_ai/batches/__init__.py diff --git a/tests/test_litellm/llms/vertex_ai/batches/test_handler.py b/tests/unit/llms/vertex_ai/batches/test_handler.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/batches/test_handler.py rename to tests/unit/llms/vertex_ai/batches/test_handler.py diff --git a/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py b/tests/unit/llms/vertex_ai/batches/test_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/batches/test_transformation.py rename to tests/unit/llms/vertex_ai/batches/test_transformation.py diff --git a/tests/test_litellm/llms/vertex_ai/files/test_transformation.py b/tests/unit/llms/vertex_ai/files/test_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/files/test_transformation.py rename to tests/unit/llms/vertex_ai/files/test_transformation.py diff --git a/tests/test_litellm/rag/__init__.py b/tests/unit/llms/vertex_ai/gemini/__init__.py similarity index 100% rename from tests/test_litellm/rag/__init__.py rename to tests/unit/llms/vertex_ai/gemini/__init__.py diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_context_circulation.py b/tests/unit/llms/vertex_ai/gemini/test_context_circulation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/gemini/test_context_circulation.py rename to tests/unit/llms/vertex_ai/gemini/test_context_circulation.py diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_function_call_args_serialization.py b/tests/unit/llms/vertex_ai/gemini/test_function_call_args_serialization.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/gemini/test_function_call_args_serialization.py rename to tests/unit/llms/vertex_ai/gemini/test_function_call_args_serialization.py diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py b/tests/unit/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py rename to tests/unit/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py b/tests/unit/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py rename to tests/unit/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_grounding_requests.py b/tests/unit/llms/vertex_ai/gemini/test_grounding_requests.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/gemini/test_grounding_requests.py rename to tests/unit/llms/vertex_ai/gemini/test_grounding_requests.py diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py b/tests/unit/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py rename to tests/unit/llms/vertex_ai/gemini/test_thought_signature_in_tool_call_id.py diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_tool_call_followed_by_text_assistant.py b/tests/unit/llms/vertex_ai/gemini/test_tool_call_followed_by_text_assistant.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/gemini/test_tool_call_followed_by_text_assistant.py rename to tests/unit/llms/vertex_ai/gemini/test_tool_call_followed_by_text_assistant.py diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py b/tests/unit/llms/vertex_ai/gemini/test_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py rename to tests/unit/llms/vertex_ai/gemini/test_transformation.py diff --git a/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py new file mode 100644 index 00000000000..4f23ac1773a --- /dev/null +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -0,0 +1,2729 @@ +import base64 + +import pytest + +from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_result, +) +from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + _transform_request_body, + check_if_part_exists_in_parts, + _get_highest_media_resolution, + _extract_max_media_resolution_from_messages, +) +from litellm.types.llms.vertex_ai import BlobType +from litellm.types.utils import Message + + +def test_check_if_part_exists_in_parts(): + parts = [ + {"text": "Hello", "thought": True}, + {"text": "World", "thought": False}, + ] + part = {"text": "Hello", "thought": True} + new_part = {"text": "Hello World", "thought": True} + assert check_if_part_exists_in_parts(parts, part) + assert not check_if_part_exists_in_parts(parts, new_part, ["thought"]) + assert check_if_part_exists_in_parts(parts, new_part, ["text"]) + + +def test_check_if_part_exists_in_parts_camel_case_snake_case(): + """Test that function handles both camelCase and snake_case key variations""" + # Test snake_case to camelCase matching + parts_with_snake_case = [ + { + "function_call": { + "name": "get_current_weather", + "args": {"location": "San Francisco, CA"}, + } + }, + {"text": "Some other content"}, + ] + + part_with_camel_case = { + "functionCall": { + "name": "get_current_weather", + "args": {"location": "San Francisco, CA"}, + } + } + + # Should find match between function_call and functionCall + assert check_if_part_exists_in_parts(parts_with_snake_case, part_with_camel_case) + + # Test camelCase to snake_case matching + parts_with_camel_case = [ + {"functionCall": {"name": "calculate_sum", "args": {"a": 1, "b": 2}}} + ] + + part_with_snake_case = { + "function_call": {"name": "calculate_sum", "args": {"a": 1, "b": 2}} + } + + # Should find match between functionCall and function_call + assert check_if_part_exists_in_parts(parts_with_camel_case, part_with_snake_case) + + # Test no match when values differ + part_with_different_values = { + "function_call": {"name": "different_function", "args": {"x": 5}} + } + + assert not check_if_part_exists_in_parts( + parts_with_snake_case, part_with_different_values + ) + + # Test multiple keys with mixed casing + parts_mixed = [ + { + "function_call": {"name": "test"}, + "thoughtSignature": "reasoning", + "text": "content", + } + ] + + part_mixed_casing = { + "functionCall": {"name": "test"}, + "thought_signature": "reasoning", + "text": "content", + } + + assert check_if_part_exists_in_parts(parts_mixed, part_mixed_casing) + + +def test_cached_content_respects_modify_params_for_cache_incompatible_fields(): + """Regression: cachedContent drops system/tools/toolConfig only when modify_params=True.""" + import litellm + + cache_name = "projects/p/locations/us-central1/cachedContents/abc123" + messages = [ + {"role": "system", "content": "You are helpful"}, + {"role": "user", "content": "hi"}, + ] + optional_params = { + "tools": [ + { + "functionDeclarations": [ + {"name": "get_weather", "description": "Get weather"}, + ] + } + ], + "tool_choice": {"functionCallingConfig": {"mode": "AUTO"}}, + } + + original_modify_params = litellm.modify_params + try: + # With modify_params=False (default), keep fields even with cachedContent. + litellm.modify_params = False + result = _transform_request_body( + messages=list(messages), + model="gemini-2.5-pro", + optional_params=dict(optional_params), + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=cache_name, + ) + assert result.get("cachedContent") == cache_name + assert "system_instruction" in result + assert "tools" in result + assert "toolConfig" in result + assert "contents" in result + + # With modify_params=True, drop cache-incompatible fields. + litellm.modify_params = True + result_modify_true = _transform_request_body( + messages=list(messages), + model="gemini-2.5-pro", + optional_params=dict(optional_params), + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=cache_name, + ) + assert result_modify_true.get("cachedContent") == cache_name + assert "system_instruction" not in result_modify_true + assert "tools" not in result_modify_true + assert "toolConfig" not in result_modify_true + assert "contents" in result_modify_true + + # Without cache, fields are always included. + result_no_cache = _transform_request_body( + messages=list(messages), + model="gemini-2.5-pro", + optional_params=dict(optional_params), + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=None, + ) + assert "system_instruction" in result_no_cache + assert "tools" in result_no_cache + assert "toolConfig" in result_no_cache + finally: + litellm.modify_params = original_modify_params + + +# Tests for issue #14556: Labels field provider-aware filtering +def test_google_genai_excludes_labels(): + """Test that Google GenAI/AI Studio endpoints exclude labels when custom_llm_provider='gemini'""" + messages = [{"role": "user", "content": "test"}] + optional_params = {"labels": {"project": "test", "team": "ai"}} + litellm_params = {} + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params, + custom_llm_provider="gemini", + litellm_params=litellm_params, + cached_content=None, + ) + + # Google GenAI/AI Studio should NOT include labels + assert "labels" not in result + assert "contents" in result + + +def test_vertex_ai_includes_labels(): + """Test that Vertex AI endpoints include labels when custom_llm_provider='vertex_ai'""" + messages = [{"role": "user", "content": "test"}] + optional_params = {"labels": {"project": "test", "team": "ai"}} + litellm_params = {} + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params, + custom_llm_provider="vertex_ai", + litellm_params=litellm_params, + cached_content=None, + ) + + # Vertex AI SHOULD include labels + assert "labels" in result + assert result["labels"] == {"project": "test", "team": "ai"} + + +def test_service_tier_forwarded_to_vertex_ai(): + """Test that service_tier in optional_params is mapped to serviceTier in request body.""" + messages = [{"role": "user", "content": "test"}] + optional_params = {"service_tier": "flex"} + litellm_params = {} + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params, + custom_llm_provider="vertex_ai", + litellm_params=litellm_params, + cached_content=None, + ) + + assert "serviceTier" in result + assert result["serviceTier"] == "flex" + + +def test_extra_body_cache_not_forwarded_to_vertex_ai(): + """ + 'cache' inside extra_body is a LiteLLM-internal proxy caching control. + It must NOT be forwarded to the Vertex AI request body. + + Regression test for: "Invalid JSON payload received. Unknown name \"cache\": Cannot find field." + Vertex AI enforces a strict JSON schema and rejects any unknown field. + """ + messages = [{"role": "user", "content": "test"}] + optional_params = { + "extra_body": { + "cache": {"use-cache": True, "ttl": 86400}, # LiteLLM-internal + "some_vertex_param": "value", # legitimate provider extra + }, + } + litellm_params = {} + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params, + custom_llm_provider="vertex_ai", + litellm_params=litellm_params, + cached_content=None, + ) + + # 'cache' must be stripped — Vertex AI has no such field + assert "cache" not in result, ( + "extra_body.cache must not be forwarded to Vertex AI. " + 'Vertex AI rejects it with 400: Unknown name "cache": Cannot find field.' + ) + + # Other legitimate extra_body keys should still pass through + assert "some_vertex_param" in result + assert result["some_vertex_param"] == "value" + + # Core request fields must be present + assert "contents" in result + + +def test_extra_body_tags_not_forwarded_to_vertex_ai(): + """ + 'tags' inside extra_body is a LiteLLM-internal param for logging/tracking. + It must NOT be forwarded to the Vertex AI request body. + Documented in litellm_proxy.md: "Send tags by including them in the extra_body parameter" + """ + messages = [{"role": "user", "content": "test"}] + optional_params = { + "extra_body": { + "tags": ["user:alice", "env:prod"], + "custom_param": "allowed", + }, + } + litellm_params = {} + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params, + custom_llm_provider="vertex_ai", + litellm_params=litellm_params, + cached_content=None, + ) + + assert "tags" not in result + assert "custom_param" in result + assert result["custom_param"] == "allowed" + + +def test_extra_body_google_maps_rewrites_json_response_format(): + messages = [{"role": "user", "content": "test"}] + optional_params = { + "response_mime_type": "application/json", + "response_schema": { + "type": "object", + "properties": {"answer": {"type": "string"}}, + }, + "extra_body": { + "tools": [{"googleMaps": {}}], + }, + } + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params, + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=None, + ) + + generation_config = result["generationConfig"] + assert "response_mime_type" not in generation_config + assert generation_config["responseFormat"] == { + "text": { + "mimeType": "APPLICATION_JSON", + "schema": { + "type": "object", + "properties": {"answer": {"type": "string"}}, + }, + } + } + + +def test_extra_body_generation_config_cannot_restore_google_maps_json_mime_type(): + messages = [{"role": "user", "content": "test"}] + optional_params = { + "tools": [{"googleMaps": {}}], + "response_mime_type": "application/json", + "extra_body": { + "generationConfig": { + "response_mime_type": "application/json", + "response_json_schema": { + "type": "object", + "properties": {"answer": {"type": "string"}}, + }, + }, + }, + } + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params, + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=None, + ) + + generation_config = result["generationConfig"] + assert "response_mime_type" not in generation_config + assert "response_json_schema" not in generation_config + assert generation_config["responseFormat"] == { + "text": { + "mimeType": "APPLICATION_JSON", + "schema": { + "type": "object", + "properties": {"answer": {"type": "string"}}, + }, + } + } + + +def test_metadata_to_labels_vertex_only(): + """Test that metadata->labels conversion only happens for Vertex AI""" + messages = [{"role": "user", "content": "test"}] + optional_params = {} + litellm_params = { + "metadata": { + "requester_metadata": {"user": "john_doe", "project": "test-project"} + } + } + + # Google GenAI/AI Studio should not include labels from metadata + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params.copy(), + custom_llm_provider="gemini", + litellm_params=litellm_params.copy(), + cached_content=None, + ) + assert "labels" not in result + + # Vertex AI should include labels from metadata + result = _transform_request_body( + messages=messages, + model="gemini-2.5-pro", + optional_params=optional_params.copy(), + custom_llm_provider="vertex_ai", + litellm_params=litellm_params.copy(), + cached_content=None, + ) + assert "labels" in result + assert result["labels"] == {"user": "john_doe", "project": "test-project"} + + +def test_empty_content_handling(): + """Test that empty content strings are properly handled in Gemini message transformation""" + # Test with empty content in user message + messages = [{"content": "", "role": "user"}] + + contents = _gemini_convert_messages_with_history(messages=messages) + + # Verify that the content was properly transformed + assert len(contents) == 1 + assert contents[0]["role"] == "user" + assert len(contents[0]["parts"]) == 1 + assert "text" in contents[0]["parts"][0] + assert contents[0]["parts"][0]["text"] == "" + + +def test_thought_signature_extraction_from_response(): + """Test that thought signatures are extracted from Gemini response parts and stored in provider_specific_fields""" + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + from litellm.types.llms.vertex_ai import HttpxPartType + + # Test case: Single function call with thought signature + test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" + + parts_with_signature = [ + HttpxPartType( + functionCall={ + "name": "get_current_temperature", + "args": {"location": "Paris"}, + }, + thoughtSignature=test_signature, + ) + ] + + function, tools, _ = VertexGeminiConfig._transform_parts( + parts=parts_with_signature, + cumulative_tool_call_idx=0, + is_function_call=False, + ) + + # Verify thought signature is stored in provider_specific_fields + assert tools is not None + assert len(tools) == 1 + assert "provider_specific_fields" in tools[0] + assert tools[0]["provider_specific_fields"]["thought_signature"] == test_signature + + +def test_thought_signature_parallel_function_calls(): + """Test that only the first function call in parallel calls has thought signature""" + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + from litellm.types.llms.vertex_ai import HttpxPartType + + test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" + + # Parallel function calls - only first has signature + parts_parallel = [ + HttpxPartType( + functionCall={ + "name": "get_current_temperature", + "args": {"location": "Paris"}, + }, + thoughtSignature=test_signature, # First FC has signature + ), + HttpxPartType( + functionCall={ + "name": "get_current_temperature", + "args": {"location": "London"}, + }, + # Second FC has no signature (parallel call) + ), + ] + + function, tools, _ = VertexGeminiConfig._transform_parts( + parts=parts_parallel, + cumulative_tool_call_idx=0, + is_function_call=False, + ) + + # Verify only first tool call has thought signature + assert tools is not None + assert len(tools) == 2 + assert "provider_specific_fields" in tools[0] + assert tools[0]["provider_specific_fields"]["thought_signature"] == test_signature + # Second tool call should not have thought signature + assert "provider_specific_fields" not in tools[ + 1 + ] or "thought_signature" not in tools[1].get("provider_specific_fields", {}) + + +def test_thought_signature_preservation_in_conversion(): + """Test that thought signatures are preserved when converting assistant messages back to Gemini format""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" + + # Assistant message with tool calls containing thought signatures + assistant_message = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": { + "name": "get_current_temperature", + "arguments": '{"location": "Paris"}', + }, + "index": 0, + "provider_specific_fields": { + "thought_signature": test_signature, + }, + }, + { + "id": "call_def456", + "type": "function", + "function": { + "name": "get_current_temperature", + "arguments": '{"location": "London"}', + }, + "index": 1, + # No thought signature for parallel call + }, + ], + } + + gemini_parts = convert_to_gemini_tool_call_invoke(assistant_message) + + # Verify thought signature is preserved in first function call part + assert len(gemini_parts) == 2 + assert "function_call" in gemini_parts[0] + assert "thoughtSignature" in gemini_parts[0] + assert gemini_parts[0]["thoughtSignature"] == test_signature + + # Verify second function call part does not have thought signature + assert "function_call" in gemini_parts[1] + assert "thoughtSignature" not in gemini_parts[1] + + +def test_thought_signature_sequential_function_calls(): + """Test that each sequential function call preserves its own thought signature""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + signature_1 = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" + signature_2 = "DifferentSignatureForSecondCall1234567890ABCDEFGHIJKLMNOPQRSTUVWXYZ" + + # Sequential function calls - each has its own signature + # This simulates a multi-step conversation where each step has a signature + assistant_message_step1 = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_step1", + "type": "function", + "function": { + "name": "check_flight", + "arguments": '{"flight": "AA100"}', + }, + "index": 0, + "provider_specific_fields": { + "thought_signature": signature_1, + }, + }, + ], + } + + assistant_message_step2 = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_step2", + "type": "function", + "function": { + "name": "book_taxi", + "arguments": '{"destination": "airport"}', + }, + "index": 0, + "provider_specific_fields": { + "thought_signature": signature_2, + }, + }, + ], + } + + gemini_parts_step1 = convert_to_gemini_tool_call_invoke(assistant_message_step1) + gemini_parts_step2 = convert_to_gemini_tool_call_invoke(assistant_message_step2) + + # Verify each step preserves its own signature + assert len(gemini_parts_step1) == 1 + assert gemini_parts_step1[0]["thoughtSignature"] == signature_1 + + assert len(gemini_parts_step2) == 1 + assert gemini_parts_step2[0]["thoughtSignature"] == signature_2 + + +def test_thought_signature_with_function_call_mode(): + """Test thought signature extraction in function_call mode (is_function_call=True)""" + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + from litellm.types.llms.vertex_ai import HttpxPartType + + test_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" + + parts_with_signature = [ + HttpxPartType( + functionCall={ + "name": "get_current_weather", + "args": {"location": "Tokyo"}, + }, + thoughtSignature=test_signature, + ) + ] + + function, tools, _ = VertexGeminiConfig._transform_parts( + parts=parts_with_signature, + cumulative_tool_call_idx=0, + is_function_call=True, + ) + + # Verify thought signature is stored in function's provider_specific_fields + assert function is not None + # Function should be dict-like (TypedDict or dict) + assert hasattr(function, "__getitem__") or isinstance(function, dict) + assert "provider_specific_fields" in function + assert function["provider_specific_fields"]["thought_signature"] == test_signature + assert tools is None + + +def test_dummy_signature_added_for_gemini_3_conversation_history(): + """Test that dummy signatures are added when transferring conversation history from older models (like gemini-2.5-flash) to gemini-3.""" + import base64 + + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + # Simulate conversation history from gemini-2.5-flash (no thought signature) + assistant_message_from_older_model = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": { + "name": "get_current_temperature", + "arguments": '{"location": "Paris"}', + }, + "index": 0, + # No provider_specific_fields - older model doesn't provide signatures + }, + ], + } + + # Convert to Gemini format for gemini-3-pro-preview (should add dummy signature) + gemini_parts = convert_to_gemini_tool_call_invoke( + assistant_message_from_older_model, model="gemini-3-pro-preview" + ) + + # Verify dummy signature is added + assert len(gemini_parts) == 1 + assert "function_call" in gemini_parts[0] + assert "thoughtSignature" in gemini_parts[0] + + # Verify it's the expected dummy signature (base64 encoded "skip_thought_signature_validator") + expected_dummy = base64.b64encode(b"skip_thought_signature_validator").decode( + "utf-8" + ) + assert gemini_parts[0]["thoughtSignature"] == expected_dummy + + +def test_dummy_signature_not_added_for_gemini_2_5(): + """Test that dummy signatures are NOT added when target model is not gemini-3.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + # Simulate conversation history from gemini-2.5-flash (no thought signature) + assistant_message = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": { + "name": "get_current_temperature", + "arguments": '{"location": "Paris"}', + }, + "index": 0, + # No provider_specific_fields + }, + ], + } + + # Convert to Gemini format for gemini-2.5-flash (should NOT add dummy signature) + gemini_parts = convert_to_gemini_tool_call_invoke( + assistant_message, model="gemini-2.5-flash" + ) + + # Verify no dummy signature is added for non-gemini-3 models + assert len(gemini_parts) == 1 + assert "function_call" in gemini_parts[0] + assert "thoughtSignature" not in gemini_parts[0] + + +def test_dummy_signature_not_added_when_signature_exists(): + """Test that dummy signatures are NOT added when a real signature already exists.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + real_signature = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n/4ZMmksdTtfQcJMoT76S1DGwhnAiLwTgWCNXs3lEb4M19EVYoWFxhrH5Lr9YMIquoU9U4paydGwvZyIyigamIg4B6WnxrRsf0KZV12gJed0DZuKczvOFtHz3zUnmZRlOiTzd5gBVyQM+5jv1VI8m4WUKd6cN/5a5ZvaA0ggiO6kdVhlpIVs7GczSEVJD8KH4u02X7VSnb7CvykqDntZzV0y8rZFBEFGKrChmeHlWXP4D1IB3F9KQyhuLgWImMzg4BajKVxxMU737JGnNISy5" + + # Assistant message with existing thought signature + assistant_message_with_signature = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": { + "name": "get_current_temperature", + "arguments": '{"location": "Paris"}', + "provider_specific_fields": { + "thought_signature": real_signature, + }, + }, + "index": 0, + }, + ], + } + + # Convert to Gemini format for gemini-3-pro-preview + gemini_parts = convert_to_gemini_tool_call_invoke( + assistant_message_with_signature, model="gemini-3-pro-preview" + ) + + # Verify real signature is preserved, not replaced with dummy + assert len(gemini_parts) == 1 + assert "function_call" in gemini_parts[0] + assert "thoughtSignature" in gemini_parts[0] + assert gemini_parts[0]["thoughtSignature"] == real_signature + + +def test_dummy_signature_with_function_call_mode(): + """Test that dummy signatures are added for function_call mode when converting to gemini-3.""" + import base64 + + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + # Assistant message with function_call (not tool_calls) and no signature + assistant_message_function_call = { + "role": "assistant", + "content": None, + "function_call": { + "name": "get_current_temperature", + "arguments": '{"location": "Paris"}', + # No provider_specific_fields + }, + } + + # Convert to Gemini format for gemini-3-pro-preview + gemini_parts = convert_to_gemini_tool_call_invoke( + assistant_message_function_call, model="gemini-3-pro-preview" + ) + + # Verify dummy signature is added + assert len(gemini_parts) == 1 + assert "function_call" in gemini_parts[0] + assert "thoughtSignature" in gemini_parts[0] + + # Verify it's the expected dummy signature + expected_dummy = base64.b64encode(b"skip_thought_signature_validator").decode( + "utf-8" + ) + assert gemini_parts[0]["thoughtSignature"] == expected_dummy + + +def _parallel_tool_calls(*signatures): + return [ + { + "id": f"call_{idx}", + "type": "function", + "function": { + "name": f"tool_{idx}", + "arguments": '{"location": "Paris"}', + **( + {"provider_specific_fields": {"thought_signature": signature}} + if signature is not None + else {} + ), + }, + "index": idx, + } + for idx, signature in enumerate(signatures) + ] + + +def _parallel_tool_calls_signed_via_id(*signatures): + """Parallel tool calls in the shape LiteLLM actually hands back to clients. + + The signature rides in the tool call id behind __thought__, which is what an + OpenAI-format client echoes back on the next turn. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + _encode_tool_call_id_with_signature, + ) + + return [ + { + "id": _encode_tool_call_id_with_signature(f"call_{idx}", signature), + "type": "function", + "function": {"name": f"tool_{idx}", "arguments": '{"location": "Paris"}'}, + "index": idx, + } + for idx, signature in enumerate(signatures) + ] + + +REAL_THOUGHT_SIGNATURE = "Co4CAdHtim/rWgXbz2Ghp4tShzLeMASrPw6JJyYIC3cbVyZnKzU3uv8/wVzyS2sKRPL2m8QQHHXbNQhEEz500G7n" +PLACEHOLDER_SIGNATURE = base64.b64encode(b"skip_thought_signature_validator").decode( + "utf-8" +) + + +def test_dummy_signature_only_on_first_parallel_tool_call(): + """Google documents the placeholder as a last resort that degrades quality, so an unsigned + parallel turn replayed to gemini-3 gets a budget of exactly one.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + gemini_parts = convert_to_gemini_tool_call_invoke( + { + "role": "assistant", + "content": None, + "tool_calls": _parallel_tool_calls(None, None, None), + }, + model="gemini-3-pro-preview", + ) + + assert len(gemini_parts) == 3 + assert gemini_parts[0]["thoughtSignature"] == PLACEHOLDER_SIGNATURE + assert "thoughtSignature" not in gemini_parts[1] + assert "thoughtSignature" not in gemini_parts[2] + + +def test_real_signature_on_first_parallel_tool_call_leaves_siblings_empty(): + """Gemini signs only the first of N parallel function calls, so a faithful replay has + nothing to attach to the siblings.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + gemini_parts = convert_to_gemini_tool_call_invoke( + { + "role": "assistant", + "content": None, + "tool_calls": _parallel_tool_calls(REAL_THOUGHT_SIGNATURE, None, None), + }, + model="gemini-3-pro-preview", + ) + + assert len(gemini_parts) == 3 + assert gemini_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE + assert "thoughtSignature" not in gemini_parts[1] + assert "thoughtSignature" not in gemini_parts[2] + + +def test_real_signature_on_later_parallel_tool_call_is_preserved(): + """Clients may reorder or drop calls, so a signature that lands on a non-first call is + still the model's own and must survive the round trip.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + gemini_parts = convert_to_gemini_tool_call_invoke( + { + "role": "assistant", + "content": None, + "tool_calls": _parallel_tool_calls(None, REAL_THOUGHT_SIGNATURE), + }, + model="gemini-3-pro-preview", + ) + + assert len(gemini_parts) == 2 + assert gemini_parts[0]["thoughtSignature"] == PLACEHOLDER_SIGNATURE + assert gemini_parts[1]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE + + +def test_no_signatures_on_parallel_tool_calls_for_gemini_2_5(): + """Non-gemini-3 models never get a placeholder signature, on any call.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + gemini_parts = convert_to_gemini_tool_call_invoke( + { + "role": "assistant", + "content": None, + "tool_calls": _parallel_tool_calls(None, None), + }, + model="gemini-2.5-flash", + ) + + assert len(gemini_parts) == 2 + assert all("thoughtSignature" not in part for part in gemini_parts) + + +def test_signature_embedded_in_tool_call_id_only_on_first_parallel_call(): + """The production shape: the signature arrives inside the first call's id, siblings have bare ids.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + gemini_parts = convert_to_gemini_tool_call_invoke( + { + "role": "assistant", + "content": None, + "tool_calls": _parallel_tool_calls_signed_via_id( + REAL_THOUGHT_SIGNATURE, None, None + ), + }, + model="gemini-3-pro-preview", + ) + + assert len(gemini_parts) == 3 + assert gemini_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE + assert "thoughtSignature" not in gemini_parts[1] + assert "thoughtSignature" not in gemini_parts[2] + + +def test_tool_level_provider_specific_fields_signature_leaves_siblings_empty(): + """A signature on the tool call itself, rather than on its function, behaves the same way.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + tool_calls = _parallel_tool_calls(None, None) + tool_calls[0]["provider_specific_fields"] = { + "thought_signature": REAL_THOUGHT_SIGNATURE + } + + gemini_parts = convert_to_gemini_tool_call_invoke( + {"role": "assistant", "content": None, "tool_calls": tool_calls}, + model="gemini-3-pro-preview", + ) + + assert len(gemini_parts) == 2 + assert gemini_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE + assert "thoughtSignature" not in gemini_parts[1] + + +def test_placeholder_lands_on_first_emitted_part_not_first_tool_call_entry(): + """A non-function entry (e.g. an OpenAI custom tool call) emits no part, so it must not + consume the one placeholder slot and leave the real first function call bare.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + tool_calls = [ + {"id": "call_custom", "type": "custom", "custom": {"name": "noop", "input": ""}} + ] + _parallel_tool_calls(None, None) + + gemini_parts = convert_to_gemini_tool_call_invoke( + {"role": "assistant", "content": None, "tool_calls": tool_calls}, + model="gemini-3-pro-preview", + ) + + assert len(gemini_parts) == 2 + assert gemini_parts[0]["thoughtSignature"] == PLACEHOLDER_SIGNATURE + assert "thoughtSignature" not in gemini_parts[1] + + +def test_no_placeholder_when_model_is_unknown(): + """Without a model there is nothing to prove the target needs a placeholder, so none is added.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + gemini_parts = convert_to_gemini_tool_call_invoke( + { + "role": "assistant", + "content": None, + "tool_calls": _parallel_tool_calls(None, None), + }, + ) + + assert len(gemini_parts) == 2 + assert all("thoughtSignature" not in part for part in gemini_parts) + + +def test_real_signature_forwarded_to_gemini_2_5_without_placeholder_siblings(): + """Older models still receive a real signature that a client replays, and still get no placeholder.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + gemini_parts = convert_to_gemini_tool_call_invoke( + { + "role": "assistant", + "content": None, + "tool_calls": _parallel_tool_calls(REAL_THOUGHT_SIGNATURE, None), + }, + model="gemini-2.5-flash", + ) + + assert len(gemini_parts) == 2 + assert gemini_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE + assert "thoughtSignature" not in gemini_parts[1] + + +def test_parallel_tool_call_history_replayed_through_full_message_conversion(): + """End to end through the message-history converter, the path a real /chat/completions replay takes.""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + messages = [ + {"role": "user", "content": "Weather in Paris, London and Tokyo?"}, + { + "role": "assistant", + "content": None, + "tool_calls": _parallel_tool_calls_signed_via_id( + REAL_THOUGHT_SIGNATURE, None, None + ), + }, + ] + + contents = _gemini_convert_messages_with_history( + messages=messages, model="gemini-3-pro-preview" + ) + + model_parts = contents[1]["parts"] + assert len(model_parts) == 3 + assert model_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE + assert "thoughtSignature" not in model_parts[1] + assert "thoughtSignature" not in model_parts[2] + + +@pytest.mark.parametrize( + "model", + ["gemini-3.5-flash", "vertex_ai/gemini-3.5-flash", "gemini/gemini-3.5-flash"], +) +def test_natively_signed_parallel_turn_never_carries_a_placeholder(model): + """A native gemini-3.5 parallel turn replays with zero skip_thought_signature_validator parts. + + Fabricating the placeholder alongside a real signature is what produced empty text responses + on gemini-3.5 parallel function calling, so the whole payload has to stay placeholder-free. + """ + import json + + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + messages = [ + {"role": "user", "content": "Weather in Paris, London and Tokyo?"}, + { + "role": "assistant", + "content": None, + "tool_calls": _parallel_tool_calls_signed_via_id( + REAL_THOUGHT_SIGNATURE, None, None + ), + }, + ] + + contents = _gemini_convert_messages_with_history(messages=messages, model=model) + + model_parts = contents[1]["parts"] + assert len(model_parts) == 3 + assert model_parts[0]["thoughtSignature"] == REAL_THOUGHT_SIGNATURE + assert "thoughtSignature" not in model_parts[1] + assert "thoughtSignature" not in model_parts[2] + assert PLACEHOLDER_SIGNATURE not in json.dumps(contents) + + +@pytest.mark.parametrize( + "model", + [ + "gemini-3-pro-preview", + "gemini-3-flash-preview", + "gemini-3.1-pro-preview", + "gemini-3.5-flash", + "gemini-3.6-flash", + "gemini-3.7-flash", + "gemini-3.8-flash", + "vertex_ai/gemini-3.5-flash", + "vertex_ai/gemini-3.7-flash", + "vertex_ai/gemini-3.8-flash", + "gemini/gemini-3.5-flash", + "gemini/gemini-3.7-flash", + "gemini/gemini-3.8-flash", + ], +) +def test_placeholder_scoped_to_first_call_across_gemini_3_variants(model): + """The gemini-3 gate is a substring match, so every family member and prefix form has to + land on the same one-placeholder budget rather than only the versions we happened to try.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + convert_to_gemini_tool_call_invoke, + ) + + gemini_parts = convert_to_gemini_tool_call_invoke( + { + "role": "assistant", + "content": None, + "tool_calls": _parallel_tool_calls(None, None, None), + }, + model=model, + ) + + assert len(gemini_parts) == 3 + assert gemini_parts[0]["thoughtSignature"] == PLACEHOLDER_SIGNATURE + assert "thoughtSignature" not in gemini_parts[1] + assert "thoughtSignature" not in gemini_parts[2] + + +def test_signed_text_part_survives_alongside_unsigned_parallel_tool_calls(): + """Text-part and function-call signatures are collected by separate code paths, so scoping the + placeholder must not disturb a real signature that arrived on the text part.""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + msg = { + "role": "assistant", + "content": "Checking all three cities.", + "provider_specific_fields": {"thought_signatures": ["real_25_signature"]}, + "tool_calls": _parallel_tool_calls(None, None, None), + } + + parts = _gemini_convert_messages_with_history( + messages=[msg], model="gemini-3-pro-preview" + )[0]["parts"] + + assert parts[0]["text"] == "Checking all three cities." + assert parts[0]["thoughtSignature"] == "real_25_signature" + assert parts[1]["thoughtSignature"] == PLACEHOLDER_SIGNATURE + assert "thoughtSignature" not in parts[2] + assert "thoughtSignature" not in parts[3] + + +# Tests for media_resolution (detail parameter) handling - Issue #17084 +class TestMediaResolution: + """Tests for media_resolution handling in Gemini 2.x models""" + + def test_get_highest_media_resolution_high_wins(self): + """Test that 'high' resolution takes precedence over 'low'""" + assert _get_highest_media_resolution("low", "high") == "high" + assert _get_highest_media_resolution("high", "low") == "high" + assert _get_highest_media_resolution(None, "high") == "high" + assert _get_highest_media_resolution("high", None) == "high" + + def test_get_highest_media_resolution_low_over_none(self): + """Test that 'low' resolution takes precedence over None""" + assert _get_highest_media_resolution(None, "low") == "low" + assert _get_highest_media_resolution("low", None) == "low" + + def test_get_highest_media_resolution_same_values(self): + """Test handling of same resolution values""" + assert _get_highest_media_resolution("high", "high") == "high" + assert _get_highest_media_resolution("low", "low") == "low" + assert _get_highest_media_resolution(None, None) is None + + def test_get_highest_media_resolution_medium(self): + """Test that 'medium' resolution is correctly ranked between 'low' and 'high'""" + assert _get_highest_media_resolution("low", "medium") == "medium" + assert _get_highest_media_resolution("medium", "low") == "medium" + assert _get_highest_media_resolution("medium", "high") == "high" + assert _get_highest_media_resolution("high", "medium") == "high" + assert _get_highest_media_resolution(None, "medium") == "medium" + assert _get_highest_media_resolution("medium", None) == "medium" + + def test_get_highest_media_resolution_ultra_high(self): + """Test that 'ultra_high' resolution takes precedence over all others""" + assert _get_highest_media_resolution("high", "ultra_high") == "ultra_high" + assert _get_highest_media_resolution("ultra_high", "high") == "ultra_high" + assert _get_highest_media_resolution("medium", "ultra_high") == "ultra_high" + assert _get_highest_media_resolution("low", "ultra_high") == "ultra_high" + assert _get_highest_media_resolution(None, "ultra_high") == "ultra_high" + assert _get_highest_media_resolution("ultra_high", None) == "ultra_high" + + def test_extract_max_media_resolution_single_image_high(self): + """Test extraction of media resolution from single image with detail=high""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is this?"}, + { + "type": "image_url", + "image_url": { + "url": "data:image/png;base64,abc123", + "detail": "high", + }, + }, + ], + } + ] + assert _extract_max_media_resolution_from_messages(messages) == "high" + + def test_extract_max_media_resolution_single_image_low(self): + """Test extraction of media resolution from single image with detail=low""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is this?"}, + { + "type": "image_url", + "image_url": { + "url": "data:image/png;base64,abc123", + "detail": "low", + }, + }, + ], + } + ] + assert _extract_max_media_resolution_from_messages(messages) == "low" + + def test_extract_max_media_resolution_no_detail(self): + """Test extraction when no detail parameter is provided""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is this?"}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,abc123"}, + }, + ], + } + ] + assert _extract_max_media_resolution_from_messages(messages) is None + + def test_extract_max_media_resolution_multiple_images_mixed(self): + """Test that highest resolution is returned when multiple images have different details""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Compare these images"}, + { + "type": "image_url", + "image_url": { + "url": "data:image/png;base64,abc123", + "detail": "low", + }, + }, + { + "type": "image_url", + "image_url": { + "url": "data:image/png;base64,def456", + "detail": "high", + }, + }, + ], + } + ] + assert _extract_max_media_resolution_from_messages(messages) == "high" + + def test_extract_max_media_resolution_text_only(self): + """Test extraction from messages with no images""" + messages = [ + {"role": "user", "content": "Hello, how are you?"}, + {"role": "assistant", "content": "I'm doing well!"}, + ] + assert _extract_max_media_resolution_from_messages(messages) is None + + def test_transform_request_body_gemini_2x_adds_media_resolution(self): + """Test that media_resolution is added to generationConfig for Gemini 2.x models""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is this?"}, + { + "type": "image_url", + "image_url": { + "url": "data:image/png;base64,iVBORw0KGgo=", + "detail": "high", + }, + }, + ], + } + ] + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-flash", + optional_params={}, + custom_llm_provider="gemini", + litellm_params={}, + cached_content=None, + ) + + assert "generationConfig" in result + assert "mediaResolution" in result["generationConfig"] + assert result["generationConfig"]["mediaResolution"] == "MEDIA_RESOLUTION_HIGH" + + def test_transform_request_body_gemini_2x_low_resolution(self): + """Test that low media_resolution is correctly added for Gemini 2.x""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is this?"}, + { + "type": "image_url", + "image_url": { + "url": "data:image/png;base64,iVBORw0KGgo=", + "detail": "low", + }, + }, + ], + } + ] + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-flash", + optional_params={}, + custom_llm_provider="gemini", + litellm_params={}, + cached_content=None, + ) + + assert "generationConfig" in result + assert "mediaResolution" in result["generationConfig"] + assert result["generationConfig"]["mediaResolution"] == "MEDIA_RESOLUTION_LOW" + + def test_transform_request_body_gemini_3_no_global_media_resolution(self): + """Test that Gemini 3 models don't add media_resolution to generationConfig (they use per-part)""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is this?"}, + { + "type": "image_url", + "image_url": { + "url": "data:image/png;base64,iVBORw0KGgo=", + "detail": "high", + }, + }, + ], + } + ] + + result = _transform_request_body( + messages=messages, + model="gemini-3-pro-preview", + optional_params={}, + custom_llm_provider="gemini", + litellm_params={}, + cached_content=None, + ) + + # Gemini 3 should NOT have mediaResolution in generationConfig + # (it's handled per-part in the content transformation) + if "generationConfig" in result: + assert "mediaResolution" not in result["generationConfig"] + + def test_transform_request_body_no_detail_no_media_resolution(self): + """Test that no mediaResolution is added when detail is not specified""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is this?"}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, + }, + ], + } + ] + + result = _transform_request_body( + messages=messages, + model="gemini-2.5-flash", + optional_params={}, + custom_llm_provider="gemini", + litellm_params={}, + cached_content=None, + ) + + # When no detail is specified, mediaResolution should not be in generationConfig + if "generationConfig" in result: + assert "mediaResolution" not in result["generationConfig"] + + def test_extract_max_media_resolution_file_type_with_detail(self): + """Test that detail is extracted from file content type, not just image_url""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is in this file?"}, + { + "type": "file", + "file": { + "url": "data:image/png;base64,abc123", + "detail": "high", + }, + }, + ], + } + ] + assert _extract_max_media_resolution_from_messages(messages) == "high" + + def test_extract_max_media_resolution_mixed_image_and_file(self): + """Test that highest detail is returned across both image_url and file types""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Compare these"}, + { + "type": "image_url", + "image_url": { + "url": "data:image/png;base64,abc123", + "detail": "low", + }, + }, + { + "type": "file", + "file": { + "url": "data:image/png;base64,def456", + "detail": "high", + }, + }, + ], + } + ] + assert _extract_max_media_resolution_from_messages(messages) == "high" + + def test_transform_request_body_gemini_1x_no_media_resolution(self): + """Test that Gemini 1.x models don't get mediaResolution in generationConfig""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is this?"}, + { + "type": "image_url", + "image_url": { + "url": "data:image/png;base64,iVBORw0KGgo=", + "detail": "high", + }, + }, + ], + } + ] + + result = _transform_request_body( + messages=messages, + model="gemini-1.5-pro", + optional_params={}, + custom_llm_provider="gemini", + litellm_params={}, + cached_content=None, + ) + + # Gemini 1.x should NOT have mediaResolution (not supported) + if "generationConfig" in result: + assert "mediaResolution" not in result["generationConfig"] + + +# Tests for VideoMetadata support across all Gemini models (Issue #25474) +class TestVideoMetadataAllGeminiModels: + """Tests that video_metadata (fps, start_offset, end_offset) works for all Gemini models""" + + def _make_video_messages(self, video_metadata: dict) -> list: + return [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Analyze this video"}, + { + "type": "file", + "file": { + "file_id": "gs://bucket/video.mp4", + "format": "video/mp4", + "video_metadata": video_metadata, + }, + }, + ], + } + ] + + def _get_file_part(self, contents: list) -> dict: + for part in contents[0]["parts"]: + if "file_data" in part: + return part + raise AssertionError("No file part found in contents") + + def test_video_metadata_fps_gemini_2_5_flash(self): + """Gemini 2.5 Flash: fps in video_metadata should be forwarded (Issue #25474)""" + messages = self._make_video_messages({"fps": 5}) + contents = _gemini_convert_messages_with_history( + messages=messages, model="gemini-2.5-flash" + ) + file_part = self._get_file_part(contents) + assert "video_metadata" in file_part + assert file_part["video_metadata"]["fps"] == 5 + + def test_video_metadata_fps_gemini_2_5_pro(self): + """Gemini 2.5 Pro: fps in video_metadata should be forwarded (Issue #25474)""" + messages = self._make_video_messages({"fps": 10}) + contents = _gemini_convert_messages_with_history( + messages=messages, model="gemini-2.5-pro" + ) + file_part = self._get_file_part(contents) + assert "video_metadata" in file_part + assert file_part["video_metadata"]["fps"] == 10 + + def test_video_metadata_offsets_gemini_2_5_flash(self): + """Gemini 2.5 Flash: start_offset/end_offset converted to camelCase (Issue #25474)""" + messages = self._make_video_messages( + {"start_offset": "5s", "end_offset": "30s"} + ) + contents = _gemini_convert_messages_with_history( + messages=messages, model="gemini-2.5-flash" + ) + file_part = self._get_file_part(contents) + assert "video_metadata" in file_part + vm = file_part["video_metadata"] + assert vm["startOffset"] == "5s" + assert vm["endOffset"] == "30s" + + def test_video_metadata_all_fields_gemini_2_5_flash(self): + """Gemini 2.5 Flash: all video_metadata fields forwarded correctly (Issue #25474)""" + messages = self._make_video_messages( + {"fps": 5, "start_offset": "10s", "end_offset": "60s"} + ) + contents = _gemini_convert_messages_with_history( + messages=messages, model="gemini-2.5-flash" + ) + file_part = self._get_file_part(contents) + assert "video_metadata" in file_part + vm = file_part["video_metadata"] + assert vm["fps"] == 5 + assert vm["startOffset"] == "10s" + assert vm["endOffset"] == "60s" + + def test_video_metadata_gemini_1_5_pro(self): + """Gemini 1.5 Pro: video_metadata should also be forwarded (Issue #25474)""" + messages = self._make_video_messages({"fps": 2}) + contents = _gemini_convert_messages_with_history( + messages=messages, model="gemini-1.5-pro" + ) + file_part = self._get_file_part(contents) + assert "video_metadata" in file_part + assert file_part["video_metadata"]["fps"] == 2 + + +def test_convert_tool_response_with_base64_image(): + """Test tool response with base64 data URI image.""" + # Create a small test image (1x1 red pixel PNG) + test_image_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" + image_data_uri = f"data:image/png;base64,{test_image_base64}" + + # Create tool message with image + tool_message = { + "role": "tool", + "tool_call_id": "call_test123", + "content": [ + { + "type": "text", + "text": '{"url": "https://example.com", "status": "success"}', + }, + {"type": "input_image", "image_url": image_data_uri}, + ], + } + + # Mock last message with tool calls + last_message_with_tool_calls = { + "tool_calls": [ + { + "id": "call_test123", + "function": {"name": "click_at", "arguments": '{"x": 100, "y": 200}'}, + } + ] + } + + # Convert tool response with nested multimodal functionResponse.parts. + result = convert_to_gemini_tool_call_result( + tool_message, last_message_with_tool_calls + ) + + assert isinstance(result, list), "Should return a parts list when media is present" + assert len(result) == 1, "Should return one function_response part" + result_part = result[0] + assert "function_response" in result_part + assert "inline_data" not in result_part + function_response = result_part["function_response"] + assert function_response["name"] == "click_at" + assert "response" in function_response + # Verify JSON response is parsed correctly + assert "url" in function_response["response"] + assert function_response["response"]["url"] == "https://example.com" + + # Check inline_data is nested under functionResponse.parts. + assert "parts" in function_response + assert len(function_response["parts"]) == 1 + inline_data: BlobType = function_response["parts"][0]["inline_data"] + assert "data" in inline_data + assert "mime_type" in inline_data + assert inline_data["mime_type"] == "image/png" + assert inline_data["data"] == test_image_base64 + + +def test_gemini_history_nests_multimodal_tool_response_parts(): + """Full history conversion should not emit sibling inline_data tool result parts.""" + test_image_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" + messages = [ + {"role": "user", "content": "Get me an image"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_get_image", + "type": "function", + "function": {"name": "get_image", "arguments": "{}"}, + } + ], + }, + { + "role": "tool", + "tool_call_id": "call_get_image", + "content": [ + {"type": "text", "text": '{"image_ref": "inline"}'}, + { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": test_image_base64, + }, + }, + ], + }, + ] + + contents = _gemini_convert_messages_with_history(messages=messages) + + tool_response_parts = contents[-1]["parts"] + assert len(tool_response_parts) == 1 + assert "inline_data" not in tool_response_parts[0] + function_response = tool_response_parts[0]["function_response"] + assert function_response["parts"] == [ + { + "inline_data": { + "data": test_image_base64, + "mime_type": "image/png", + } + } + ] + + +def test_convert_tool_response_text_only(): + """Test tool response with only text (no image).""" + tool_message = { + "role": "tool", + "tool_call_id": "call_test789", + "content": [ + {"type": "text", "text": '{"status": "completed", "result": "success"}'} + ], + } + + last_message_with_tool_calls = { + "tool_calls": [ + { + "id": "call_test789", + "function": {"name": "wait_5_seconds", "arguments": "{}"}, + } + ] + } + + result = convert_to_gemini_tool_call_result( + tool_message, last_message_with_tool_calls + ) + + # Should be a single part (no list) when no image + assert not isinstance(result, list), "Should return single part when no image" + + # Check function_response exists + assert "function_response" in result + function_response = result["function_response"] + assert function_response["name"] == "wait_5_seconds" + # Verify JSON response is parsed correctly + assert "status" in function_response["response"] + assert function_response["response"]["status"] == "completed" + + # Check inline_data does NOT exist (no image provided) + assert "inline_data" not in result + + +def test_file_data_field_order(): + """ + Test that file_data fields are in the correct order (mime_type before file_uri). + + The Gemini API is sensitive to field order in the file_data object. + This test verifies that mime_type comes before file_uri in both: + 1. Dictionary key order + 2. JSON serialization + + Related issue: Gemini API returns 400 INVALID_ARGUMENT when fields are in wrong order. + """ + import json + + from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_media + + # Test with HTTPS URL and explicit format (audio file) + file_url = "https://generativelanguage.googleapis.com/v1beta/files/test123" + format = "audio/mpeg" + + result = _process_gemini_media(image_url=file_url, format=format) + + # Verify the result has file_data + assert "file_data" in result + file_data = result["file_data"] + + # Verify both fields are present + assert "mime_type" in file_data + assert "file_uri" in file_data + assert file_data["mime_type"] == "audio/mpeg" + assert file_data["file_uri"] == file_url + + # Verify field order by checking dictionary keys + # In Python 3.7+, dict maintains insertion order + file_data_keys = list(file_data.keys()) + assert file_data_keys.index("mime_type") < file_data_keys.index( + "file_uri" + ), "mime_type must come before file_uri in the file_data dict" + + # Also verify by serializing to JSON string + json_str = json.dumps(file_data) + mime_type_pos = json_str.find('"mime_type"') + file_uri_pos = json_str.find('"file_uri"') + assert ( + mime_type_pos < file_uri_pos + ), "mime_type must appear before file_uri in JSON serialization" + + +def test_file_data_field_order_gcs_urls(): + """Test that GCS URLs also maintain correct field order.""" + import json + + from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_media + + # Test with GCS URL + gcs_url = "gs://bucket/audio.mp3" + + result = _process_gemini_media(image_url=gcs_url) + + # Verify the result has file_data + assert "file_data" in result + file_data = result["file_data"] + + # Verify both fields are present + assert "mime_type" in file_data + assert "file_uri" in file_data + + # Verify field order + file_data_keys = list(file_data.keys()) + assert file_data_keys.index("mime_type") < file_data_keys.index( + "file_uri" + ), "mime_type must come before file_uri in the file_data dict" + + +def test_gemini_files_api_uri_without_format(): + """ + Test that Gemini Files API URIs work WITHOUT an explicit format/mime_type. + + When a user uploads a file via the Gemini Files API and then references it + by URI (https://generativelanguage.googleapis.com/v1beta/files/...), + the file is already on Google's servers. These URLs return 403 when + fetched directly, so _process_gemini_media must NOT try to resolve the + MIME type via HTTP. Instead it should pass the URI through as file_data + and let the Gemini API resolve the type from its stored metadata. + + Related issue: https://github.com/BerriAI/litellm/issues/24907 + """ + from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_media + + file_url = "https://generativelanguage.googleapis.com/v1beta/files/37eh7rsw1vfe" + + # Should NOT raise — previously this hit the generic https:// handler + # which called _get_image_mime_type_from_url() and got a 403. + result = _process_gemini_media(image_url=file_url) + + assert "file_data" in result + file_data = result["file_data"] + assert file_data["file_uri"] == file_url + # When no format is provided, mime_type should be absent so the + # Gemini API infers it from the stored file metadata. + assert "mime_type" not in file_data + + +def test_gemini_files_api_uri_with_format(): + """ + Test that Gemini Files API URIs correctly forward an explicit format. + + Related issue: https://github.com/BerriAI/litellm/issues/24907 + """ + from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_media + + file_url = "https://generativelanguage.googleapis.com/v1beta/files/n1vhxa28lyaw" + + result = _process_gemini_media(image_url=file_url, format="text/plain") + + assert "file_data" in result + file_data = result["file_data"] + assert file_data["file_uri"] == file_url + assert file_data["mime_type"] == "text/plain" + + +def test_extract_file_data_with_path_object(): + """ + Test that filename is correctly extracted from Path objects for MIME type detection. + + When uploading files using Path objects (e.g., Path("speech.mp3")), the filename + must be extracted to enable proper MIME type detection. Without this, files get + uploaded with 'application/octet-stream' instead of the correct MIME type. + + Related issue: Files uploaded with wrong MIME type cause Gemini API to reject + requests where the specified format doesn't match the uploaded file's MIME type. + """ + import os + import tempfile + from pathlib import Path + + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + extract_file_data, + ) + + # Create a temporary MP3 file + with tempfile.NamedTemporaryFile(suffix=".mp3", delete=False) as tmp: + tmp.write(b"fake mp3 content") + tmp_path = tmp.name + + try: + # Test with Path object + path_obj = Path(tmp_path) + extracted = extract_file_data(path_obj) + + # Verify filename was extracted + assert extracted["filename"] is not None + assert extracted["filename"].endswith(".mp3") + + # Verify MIME type was correctly detected + assert ( + extracted["content_type"] == "audio/mpeg" + ), f"Expected 'audio/mpeg' but got '{extracted['content_type']}'" + + # Verify content was read + assert extracted["content"] == b"fake mp3 content" + + finally: + # Clean up temporary file + os.unlink(tmp_path) + + +def test_extract_file_data_with_pathlib_path(): + """Test that filename is correctly extracted from pathlib.Path inputs. + Bare str paths are rejected — when this runs in a proxy request handler + the value is attacker-controlled and opening it as a path is an LFI.""" + import os + import tempfile + from pathlib import Path + + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + extract_file_data, + ) + + with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp: + tmp.write(b"fake wav content") + tmp_path = Path(tmp.name) + + try: + extracted = extract_file_data(tmp_path) + + assert extracted["filename"] is not None + assert extracted["filename"].endswith(".wav") + assert extracted["content_type"] in [ + "audio/wav", + "audio/x-wav", + ], f"Expected 'audio/wav' or 'audio/x-wav' but got '{extracted['content_type']}'" + assert extracted["content"] == b"fake wav content" + finally: + os.unlink(str(tmp_path)) + + +def test_extract_file_data_with_tuple_format(): + """Test that tuple format (with explicit content_type) still works correctly.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + extract_file_data, + ) + + # Test with tuple format: (filename, content, content_type) + filename = "test_audio.mp3" + content = b"test audio content" + content_type = "audio/mpeg" + + extracted = extract_file_data((filename, content, content_type)) + + # Verify all fields are correct + assert extracted["filename"] == filename + assert extracted["content"] == content + assert extracted["content_type"] == content_type + + +def test_extract_file_data_fallback_to_octet_stream(): + """Unknown file types fall back to application/octet-stream.""" + import os + import tempfile + from pathlib import Path + + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + extract_file_data, + ) + + with tempfile.NamedTemporaryFile(suffix=".xyz123", delete=False) as tmp: + tmp.write(b"unknown content") + tmp_path = Path(tmp.name) + + try: + extracted = extract_file_data(tmp_path) + + assert extracted["filename"] is not None + assert extracted["filename"].endswith(".xyz123") + assert ( + extracted["content_type"] == "application/octet-stream" + ), f"Expected 'application/octet-stream' for unknown type, got '{extracted['content_type']}'" + finally: + os.unlink(str(tmp_path)) + + +def test_convert_tool_response_with_pdf_file(): + """Test tool response with PDF file content using file_data field.""" + # Create a minimal test PDF (base64 encoded) + test_pdf_base64 = "JVBERi0xLjQKJeLjz9MKMSAwIG9iago8PC9UeXBlL0NhdGFsb2cvUGFnZXMgMiAwIFI+PgplbmRvYmoKdHJhaWxlcgo8PC9TaXplIDQvUm9vdCAxIDAgUj4+CnN0YXJ0eHJlZgoyMTYKJSVFT0Y=" + file_data_uri = f"data:application/pdf;base64,{test_pdf_base64}" + + # Create tool message with file + tool_message = { + "role": "tool", + "tool_call_id": "call_pdf_test", + "content": [ + {"type": "text", "text": '{"status": "success", "pages": 1}'}, + {"type": "file", "file_data": file_data_uri}, + ], + } + + # Mock last message with tool calls + last_message_with_tool_calls = { + "tool_calls": [ + { + "id": "call_pdf_test", + "function": { + "name": "analyze_document", + "arguments": '{"path": "/tmp/doc.pdf"}', + }, + } + ] + } + + # Convert tool response with nested multimodal functionResponse.parts. + result = convert_to_gemini_tool_call_result( + tool_message, last_message_with_tool_calls + ) + + assert isinstance(result, list), "Should return a parts list when media is present" + assert len(result) == 1, "Should return one function_response part" + result_part = result[0] + assert "function_response" in result_part + assert "inline_data" not in result_part + function_response = result_part["function_response"] + assert function_response["name"] == "analyze_document" + assert "response" in function_response + # Verify JSON response is parsed correctly + assert "status" in function_response["response"] + assert function_response["response"]["status"] == "success" + + # Check inline_data is nested under functionResponse.parts. + assert "parts" in function_response + assert len(function_response["parts"]) == 1 + inline_data: BlobType = function_response["parts"][0]["inline_data"] + assert "data" in inline_data + assert "mime_type" in inline_data + assert inline_data["mime_type"] == "application/pdf" + assert inline_data["data"] == test_pdf_base64 + + +def test_convert_tool_response_with_input_file_type(): + """Test tool response with input_file content type (Responses API format).""" + # Create a minimal test PDF (base64 encoded) + test_pdf_base64 = "JVBERi0xLjQKJeLjz9MKMSAwIG9iago8PC9UeXBlL0NhdGFsb2cvUGFnZXMgMiAwIFI+PgplbmRvYmoKdHJhaWxlcgo8PC9TaXplIDQvUm9vdCAxIDAgUj4+CnN0YXJ0eHJlZgoyMTYKJSVFT0Y=" + file_data_uri = f"data:application/pdf;base64,{test_pdf_base64}" + + # Create tool message with input_file type + tool_message = { + "role": "tool", + "tool_call_id": "call_input_file_test", + "content": [{"type": "input_file", "file_data": file_data_uri}], + } + + # Mock last message with tool calls + last_message_with_tool_calls = { + "tool_calls": [ + { + "id": "call_input_file_test", + "function": {"name": "read_file", "arguments": "{}"}, + } + ] + } + + # Convert tool response + result = convert_to_gemini_tool_call_result( + tool_message, last_message_with_tool_calls + ) + + # Check inline_data is nested under functionResponse.parts. + assert isinstance(result, list), "Should return a parts list when media is present" + assert len(result) == 1, "Should return one function_response part" + function_response = result[0]["function_response"] + assert ( + function_response["parts"][0]["inline_data"]["mime_type"] == "application/pdf" + ) + + +def test_convert_tool_response_with_nested_file_object(): + """Test tool response with file content using nested file object format.""" + # Create a minimal test PDF (base64 encoded) + test_pdf_base64 = "JVBERi0xLjQKJeLjz9MKMSAwIG9iago8PC9UeXBlL0NhdGFsb2cvUGFnZXMgMiAwIFI+PgplbmRvYmoKdHJhaWxlcgo8PC9TaXplIDQvUm9vdCAxIDAgUj4+CnN0YXJ0eHJlZgoyMTYKJSVFT0Y=" + file_data_uri = f"data:application/pdf;base64,{test_pdf_base64}" + + # Create tool message with nested file object (OpenAI Agents SDK format) + tool_message = { + "role": "tool", + "tool_call_id": "call_nested_test", + "content": [{"type": "file", "file": {"file_data": file_data_uri}}], + } + + # Mock last message with tool calls + last_message_with_tool_calls = { + "tool_calls": [ + { + "id": "call_nested_test", + "function": {"name": "process_document", "arguments": "{}"}, + } + ] + } + + # Convert tool response + result = convert_to_gemini_tool_call_result( + tool_message, last_message_with_tool_calls + ) + + # Check inline_data is nested under functionResponse.parts. + assert isinstance(result, list), "Should return a parts list when media is present" + assert len(result) == 1, "Should return one function_response part" + function_response = result[0]["function_response"] + inline_data: BlobType = function_response["parts"][0]["inline_data"] + assert "data" in inline_data + assert "mime_type" in inline_data + assert inline_data["mime_type"] == "application/pdf" + assert inline_data["data"] == test_pdf_base64 + + +def test_assistant_message_with_images_field(): + """ + Test that assistant messages with images field are properly converted to Gemini format. + + This handles the case where an assistant message contains generated images in the + `images` field (e.g., from image generation models like gemini-2.5-flash-image). + The images should be converted to inline_data parts in the Gemini format. + """ + # Create a small test image (1x1 red pixel PNG) + test_image_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" + image_data_uri = f"data:image/png;base64,{test_image_base64}" + + # Create messages with assistant message containing images field + messages = [ + { + "role": "user", + "content": "Generate an image of a banana wearing a costume that says LiteLLM", + }, + { + "role": "assistant", + "content": "Here's your banana in a LiteLLM costume!", + "images": [ + { + "image_url": {"url": image_data_uri, "detail": "auto"}, + "index": 0, + "type": "image_url", + } + ], + }, + ] + + # Convert messages to Gemini format + contents = _gemini_convert_messages_with_history(messages=messages) + + # Verify structure + assert len(contents) == 2, f"Expected 2 content blocks, got {len(contents)}" + + # Verify user message + assert contents[0]["role"] == "user" + assert len(contents[0]["parts"]) == 1 + assert ( + contents[0]["parts"][0]["text"] + == "Generate an image of a banana wearing a costume that says LiteLLM" + ) + + # Verify assistant message + assert contents[1]["role"] == "model" + assert ( + len(contents[1]["parts"]) == 2 + ), f"Expected 2 parts (text + image), got {len(contents[1]['parts'])}" + + # Find text part and inline_data part + text_part = None + inline_data_part = None + for part in contents[1]["parts"]: + if "text" in part: + text_part = part + elif "inline_data" in part: + inline_data_part = part + + # Verify text part + assert text_part is not None, "Missing text part in assistant message" + assert text_part["text"] == "Here's your banana in a LiteLLM costume!" + + # Verify inline_data part (image) + assert inline_data_part is not None, "Missing inline_data part in assistant message" + inline_data: BlobType = inline_data_part["inline_data"] + assert "data" in inline_data + assert "mime_type" in inline_data + assert inline_data["mime_type"] == "image/png" + assert inline_data["data"] == test_image_base64 + + +def test_assistant_message_with_multiple_images(): + """Test that assistant messages with multiple images are properly converted.""" + # Create two test images + test_image1_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" + test_image2_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8DwHwAFBQIAX8jx0gAAAABJRU5ErkJggg==" + image1_data_uri = f"data:image/png;base64,{test_image1_base64}" + image2_data_uri = f"data:image/jpeg;base64,{test_image2_base64}" + + messages = [ + {"role": "user", "content": "Generate two images"}, + { + "role": "assistant", + "content": "Here are your images:", + "images": [ + { + "image_url": {"url": image1_data_uri, "detail": "auto"}, + "index": 0, + "type": "image_url", + }, + { + "image_url": {"url": image2_data_uri, "detail": "high"}, + "index": 1, + "type": "image_url", + }, + ], + }, + ] + + # Convert messages to Gemini format + contents = _gemini_convert_messages_with_history(messages=messages) + + # Verify assistant message has 3 parts (1 text + 2 images) + assert contents[1]["role"] == "model" + assert ( + len(contents[1]["parts"]) == 3 + ), f"Expected 3 parts (text + 2 images), got {len(contents[1]['parts'])}" + + # Count inline_data parts + inline_data_parts = [part for part in contents[1]["parts"] if "inline_data" in part] + assert ( + len(inline_data_parts) == 2 + ), f"Expected 2 inline_data parts, got {len(inline_data_parts)}" + + # Verify first image + assert inline_data_parts[0]["inline_data"]["mime_type"] == "image/png" + assert inline_data_parts[0]["inline_data"]["data"] == test_image1_base64 + + # Verify second image + assert inline_data_parts[1]["inline_data"]["mime_type"] == "image/jpeg" + assert inline_data_parts[1]["inline_data"]["data"] == test_image2_base64 + + +def test_assistant_message_with_images_using_message_object(): + """Test that Message objects with images field are properly converted.""" + # Create a small test image + test_image_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" + image_data_uri = f"data:image/png;base64,{test_image_base64}" + + # Create messages using Message object (as returned by LiteLLM) + user_message = {"role": "user", "content": "Generate an image"} + + assistant_message = Message( + content="Here's your image!", + role="assistant", + tool_calls=None, + function_call=None, + images=[ + { + "image_url": {"url": image_data_uri, "detail": "auto"}, + "index": 0, + "type": "image_url", + } + ], + ) + + messages = [user_message, assistant_message] + + # Convert messages to Gemini format + contents = _gemini_convert_messages_with_history(messages=messages) + + # Verify assistant message has both text and image + assert contents[1]["role"] == "model" + assert len(contents[1]["parts"]) == 2 + + # Verify image was converted + inline_data_parts = [part for part in contents[1]["parts"] if "inline_data" in part] + assert len(inline_data_parts) == 1 + assert inline_data_parts[0]["inline_data"]["mime_type"] == "image/png" + assert inline_data_parts[0]["inline_data"]["data"] == test_image_base64 + + +def test_assistant_message_with_images_in_conversation_history(): + """ + Test multi-turn conversation where assistant message with images is in history. + + This simulates the real use case where: + 1. User asks for image generation + 2. Assistant generates image (with images field) + 3. User asks follow-up question about the image + """ + test_image_base64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" + image_data_uri = f"data:image/png;base64,{test_image_base64}" + + messages = [ + {"role": "user", "content": "Generate an image of a cat"}, + { + "role": "assistant", + "content": "Here's a cat image:", + "images": [ + { + "image_url": {"url": image_data_uri, "detail": "auto"}, + "index": 0, + "type": "image_url", + } + ], + }, + {"role": "user", "content": "Can you make it more colorful?"}, + ] + + # Convert messages to Gemini format + contents = _gemini_convert_messages_with_history(messages=messages) + + # Verify structure: user -> model (with image) -> user + assert len(contents) == 3 + assert contents[0]["role"] == "user" + assert contents[1]["role"] == "model" + assert contents[2]["role"] == "user" + + # Verify assistant message has image in history + inline_data_parts = [part for part in contents[1]["parts"] if "inline_data" in part] + assert len(inline_data_parts) == 1 + assert inline_data_parts[0]["inline_data"]["mime_type"] == "image/png" + + +def test_function_response_has_user_role(): + """ + Test that function response ContentType blocks include role="user". + + Gemini API only accepts two roles: "user" and "model". Function responses + must be sent with role="user". Previously, LiteLLM omitted the role field + entirely, causing 400 errors from the Gemini API. + + Fixes: https://github.com/BerriAI/litellm/issues/22003 + Fixes: https://github.com/BerriAI/litellm/issues/20690 + """ + messages = [ + {"role": "user", "content": "What is the weather in Berlin?"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"city": "Berlin"}', + }, + } + ], + }, + { + "role": "tool", + "tool_call_id": "call_abc123", + "content": '{"temperature": "15°C", "condition": "Cloudy"}', + }, + ] + + contents = _gemini_convert_messages_with_history(messages=messages) + + # Expect: user -> model (functionCall) -> user (functionResponse) + assert len(contents) == 3 + + assert contents[0]["role"] == "user" + assert contents[1]["role"] == "model" + assert "function_call" in contents[1]["parts"][0] + + # The critical assertion: function response must have role="user" + assert contents[2]["role"] == "user" + assert "function_response" in contents[2]["parts"][0] + + +def test_multi_turn_function_calling_roles(): + """ + Test a full multi-turn function calling conversation produces correct roles. + + Simulates: user asks → model calls tool → tool responds → model answers → user asks again. + Every content block must have an explicit role of "user" or "model". + + Fixes: https://github.com/BerriAI/litellm/issues/22003 + """ + messages = [ + {"role": "user", "content": "What is the weather in Berlin?"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_001", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"city": "Berlin"}', + }, + } + ], + }, + { + "role": "tool", + "tool_call_id": "call_001", + "content": '{"temperature": "15°C"}', + }, + { + "role": "assistant", + "content": "The weather in Berlin is 15°C.", + }, + {"role": "user", "content": "And in Paris?"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_002", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"city": "Paris"}', + }, + } + ], + }, + { + "role": "tool", + "tool_call_id": "call_002", + "content": '{"temperature": "18°C"}', + }, + ] + + contents = _gemini_convert_messages_with_history(messages=messages) + + # Every content block must have a valid role + for i, content in enumerate(contents): + assert "role" in content, f"Content block {i} missing 'role' field" + assert content["role"] in ( + "user", + "model", + ), f"Content block {i} has invalid role: {content.get('role')}" + + # Verify the function response blocks specifically have role="user" + for i, content in enumerate(contents): + for part in content["parts"]: + if "function_response" in part: + assert ( + content["role"] == "user" + ), f"Content block {i} with function_response has role='{content['role']}', expected 'user'" + + +def test_gemini_thought_signature_preservation_real_response(): + """Test that thought signatures are preserved on the text part if originally there, without dropping or duplicating (real response case).""" + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + real_candidate = { + "content": { + "parts": [ + { + "text": "I will explain and then list files.", + "thoughtSignature": "mock_signature_from_text_part", + }, + { + "functionCall": { + "name": "list_files", + "args": {}, + } + }, + ] + } + } + + parts = real_candidate["content"]["parts"] + + content, reasoning_content = ( + VertexGeminiConfig().get_assistant_content_message(parts=parts) + ) + thought_signatures = ( + VertexGeminiConfig()._extract_thought_signatures_from_parts( + parts=parts + ) + ) + functions, tools, _ = VertexGeminiConfig._transform_parts( + parts=parts, + cumulative_tool_call_idx=0, + is_function_call=False, + ) + + msg: dict = {"role": "assistant"} + if content is not None: + msg["content"] = content + if tools: + msg["tool_calls"] = tools + if functions is not None: + msg["function_call"] = functions + if thought_signatures is not None: + msg["provider_specific_fields"] = { + "thought_signatures": thought_signatures + } + + converted_real = _gemini_convert_messages_with_history( + messages=[msg], + model="gemini-2.5-pro", + ) + + assert len(converted_real) == 1 + assert "parts" in converted_real[0] + parts_out = converted_real[0]["parts"] + assert len(parts_out) == 2 + assert "text" in parts_out[0] + assert ( + parts_out[0]["thoughtSignature"] == "mock_signature_from_text_part" + ) + assert "function_call" in parts_out[1] + assert "thoughtSignature" not in parts_out[1] + + +def test_gemini_thought_signature_deduplication_assumed_response(): + """Test that thought signatures are deduplicated and not attached to the text part if already present in the tool call (assumed response case).""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + pr_assumed_msg = { + "role": "assistant", + "content": "I will list the directory.", + "provider_specific_fields": { + "thought_signatures": ["mock_signature_63k"] + }, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "list_files", "arguments": "{}"}, + "provider_specific_fields": { + "thought_signature": "mock_signature_63k" + }, + } + ], + } + + converted_pr = _gemini_convert_messages_with_history( + messages=[pr_assumed_msg], + model="gemini-2.5-pro", + ) + + assert len(converted_pr) == 1 + assert "parts" in converted_pr[0] + parts_out = converted_pr[0]["parts"] + assert len(parts_out) == 2 + assert "text" in parts_out[0] + assert "thoughtSignature" not in parts_out[0] + assert "function_call" in parts_out[1] + assert parts_out[1]["thoughtSignature"] == "mock_signature_63k" + + +def test_gemini_thought_signature_pure_text(): + """Test that thought signatures are preserved on the text part for responses with no tool calls.""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + msg = { + "role": "assistant", + "content": "Hello, I am a model.", + "provider_specific_fields": { + "thought_signatures": ["pure_text_signature"] + }, + } + + converted = _gemini_convert_messages_with_history( + messages=[msg], + model="gemini-2.5-pro", + ) + + assert len(converted) == 1 + assert "parts" in converted[0] + parts_out = converted[0]["parts"] + assert len(parts_out) == 1 + assert "text" in parts_out[0] + assert parts_out[0]["thoughtSignature"] == "pure_text_signature" + + +def test_gemini_thought_signature_pure_tool_call(): + """Test that thought signatures are preserved on the tool call for responses with no intermediate text.""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + msg = { + "role": "assistant", + "content": None, + "provider_specific_fields": { + "thought_signatures": ["pure_tool_signature"] + }, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "list_files", "arguments": "{}"}, + "provider_specific_fields": { + "thought_signature": "pure_tool_signature" + }, + } + ], + } + + converted = _gemini_convert_messages_with_history( + messages=[msg], + model="gemini-2.5-pro", + ) + + assert len(converted) == 1 + assert "parts" in converted[0] + parts_out = converted[0]["parts"] + assert len(parts_out) == 1 + assert "function_call" in parts_out[0] + assert parts_out[0]["thoughtSignature"] == "pure_tool_signature" + + +def test_gemini_distinct_text_and_tool_signatures_are_both_preserved(): + """A text-part signature that differs from the tool-call signature must stay on the text part.""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + msg = { + "role": "assistant", + "content": "Some analysis.", + "provider_specific_fields": { + "thought_signatures": ["text_signature", "tool_signature"] + }, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "list_files", "arguments": "{}"}, + "provider_specific_fields": {"thought_signature": "tool_signature"}, + } + ], + } + + parts = _gemini_convert_messages_with_history( + messages=[msg], model="gemini-2.5-pro" + )[0]["parts"] + + assert parts[0]["text"] == "Some analysis." + assert parts[0]["thoughtSignature"] == "text_signature" + assert "function_call" in parts[1] + assert parts[1]["thoughtSignature"] == "tool_signature" + + +def test_gemini_25_text_signature_survives_replay_to_gemini_3(): + """gemini-2.5 history (signed text, unsigned tool call) replayed to gemini-3 keeps the real + text signature; the dummy signature synthesized for the unsigned tool call must not suppress it.""" + from litellm.litellm_core_utils.prompt_templates.factory import ( + _get_dummy_thought_signature, + ) + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + msg = { + "role": "assistant", + "content": "I will list the directory.", + "provider_specific_fields": {"thought_signatures": ["real_25_signature"]}, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "list_files", "arguments": "{}"}, + } + ], + } + + parts = _gemini_convert_messages_with_history(messages=[msg], model="gemini-3-pro")[ + 0 + ]["parts"] + + assert parts[0]["text"] == "I will list the directory." + assert parts[0]["thoughtSignature"] == "real_25_signature" + assert "function_call" in parts[1] + assert parts[1]["thoughtSignature"] == _get_dummy_thought_signature() + + +def test_gemini_function_call_signature_round_trip_no_duplicate(): + """End to end: a gemini-3-style response (unsigned text + signed functionCall) parsed and + re-serialized sends the signature exactly once, on the function-call part.""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + + response_parts = [ + {"text": "I will calculate the result for you."}, + { + "functionCall": {"name": "add_numbers", "args": {"a": 17, "b": 25}}, + "thoughtSignature": "signature_from_function_call", + }, + ] + + config = VertexGeminiConfig() + content, _ = config.get_assistant_content_message(parts=response_parts) + thought_signatures = config._extract_thought_signatures_from_parts( + parts=response_parts + ) + _, tools, _ = VertexGeminiConfig._transform_parts( + parts=response_parts, cumulative_tool_call_idx=0, is_function_call=False + ) + + msg = { + "role": "assistant", + "content": content, + "tool_calls": tools, + "provider_specific_fields": {"thought_signatures": thought_signatures}, + } + + parts = _gemini_convert_messages_with_history(messages=[msg], model="gemini-3-pro")[ + 0 + ]["parts"] + + signatures = [p["thoughtSignature"] for p in parts if "thoughtSignature" in p] + assert signatures == ["signature_from_function_call"] + assert "thoughtSignature" not in parts[0] + assert "function_call" in parts[1] + + +def test_gemini_server_side_tool_signature_not_duplicated_on_text(): + """A signature already re-injected on a server-side toolCall part is not attached to the text part again.""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + msg = { + "role": "assistant", + "content": "The weather in Buenos Aires is sunny.", + "provider_specific_fields": { + "thought_signatures": ["server_side_signature"], + "server_side_tool_invocations": [ + { + "tool_type": "GOOGLE_SEARCH_WEB", + "id": "abc123", + "args": {"queries": ["weather Buenos Aires"]}, + "response": {"weather": "Sunny"}, + "thought_signature": "server_side_signature", + } + ], + }, + } + + parts = _gemini_convert_messages_with_history( + messages=[msg], model="gemini-2.5-pro" + )[0]["parts"] + + text_part = next(p for p in parts if "text" in p) + assert "thoughtSignature" not in text_part + tool_call_part = next(p for p in parts if "toolCall" in p) + assert tool_call_part["thoughtSignature"] == "server_side_signature" diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py similarity index 99% rename from tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py rename to tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 88ba7fc37d9..739744336a1 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -1894,42 +1894,6 @@ def test_vertex_ai_tool_call_id_format(): ), f"All 10 IDs should be unique, got {len(ids_generated)} unique IDs" -def test_vertex_ai_code_line_length(): - """ - Test that the specific code line generating tool call IDs is within character limit. - - This is a meta-test to ensure the code change meets the 40-character requirement. - """ - import inspect - - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - - # Get the source code of the _transform_parts method - source_lines = inspect.getsource(VertexGeminiConfig._transform_parts).split("\n") - - # Find the line that generates the ID - id_line = None - for line in source_lines: - if '"id": f"call_' in line and "uuid.uuid4().hex[:28]" in line: - id_line = line.strip() # Remove indentation for length check - break - - assert id_line is not None, "Could not find the ID generation line in source code" - - # Check that the line is 40 characters or less (excluding indentation) - line_length = len(id_line) - assert ( - line_length <= 40 - ), f"ID generation line is {line_length} characters, should be ≤40: {id_line}" - - # Verify it contains the expected UUID format - assert ( - "uuid.uuid4().hex[:28]" in id_line - ), f"Line should contain shortened UUID format: {id_line}" - - def test_vertex_ai_map_google_maps_tool_simple(): """ Test googleMaps tool transformation without location data. @@ -2530,8 +2494,6 @@ def test_fine_tuned_endpoint_and_gemma_get_no_gemini_3_default_temperature(model assert "temperature" not in mapped - - def _tool_call_messages(tool_call_id: str): return [ {"role": "user", "content": "hi"}, diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_gemini_unbound_local_error.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_gemini_unbound_local_error.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/gemini/test_vertex_gemini_unbound_local_error.py rename to tests/unit/llms/vertex_ai/gemini/test_vertex_gemini_unbound_local_error.py diff --git a/tests/test_litellm/rag/ingestion/__init__.py b/tests/unit/llms/vertex_ai/image_generation/__init__.py similarity index 100% rename from tests/test_litellm/rag/ingestion/__init__.py rename to tests/unit/llms/vertex_ai/image_generation/__init__.py diff --git a/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_cost_calculator.py b/tests/unit/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_cost_calculator.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_cost_calculator.py rename to tests/unit/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_cost_calculator.py 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 new file mode 100644 index 00000000000..a72a570c2a2 --- /dev/null +++ b/tests/unit/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py @@ -0,0 +1,637 @@ +from unittest.mock import MagicMock, patch + +import httpx + + +from litellm.llms.vertex_ai.image_generation import ( + get_vertex_ai_image_generation_config, +) +from litellm.llms.vertex_ai.image_generation.vertex_gemini_transformation import ( + VertexAIGeminiImageGenerationConfig, +) +from litellm.llms.vertex_ai.image_generation.vertex_imagen_transformation import ( + VertexAIImagenImageGenerationConfig, +) + + +class TestVertexAIGeminiImageGenerationConfig: + def setup_method(self): + """Set up test fixtures""" + self.config = VertexAIGeminiImageGenerationConfig() + + def test_get_supported_openai_params(self): + """Test get_supported_openai_params returns correct params""" + supported = self.config.get_supported_openai_params("gemini-2.5-flash-image") + assert "n" in supported + assert "size" in supported + + def test_map_openai_params_n(self): + """Test mapping n parameter to candidate_count""" + non_default_params = {"n": 3} + optional_params = {} + result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False) + assert result.get("candidate_count") == 3 + + def test_map_openai_params_size(self): + """Test mapping size parameter to aspectRatio""" + non_default_params = {"size": "1024x1024"} + optional_params = {} + result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False) + assert result.get("aspectRatio") == "1:1" + + def test_map_openai_params_size_16_9(self): + """Test mapping 16:9 size""" + non_default_params = {"size": "1792x1024"} + optional_params = {} + result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False) + assert result.get("aspectRatio") == "16:9" + + def test_map_size_to_aspect_ratio(self): + """Test size to aspect ratio mapping""" + assert self.config._map_size_to_aspect_ratio("1024x1024") == "1:1" + assert self.config._map_size_to_aspect_ratio("1792x1024") == "16:9" + assert self.config._map_size_to_aspect_ratio("1024x1792") == "9:16" + assert self.config._map_size_to_aspect_ratio("1280x896") == "4:3" + assert self.config._map_size_to_aspect_ratio("896x1280") == "3:4" + assert self.config._map_size_to_aspect_ratio("unknown") == "1:1" # default + + def test_get_supported_openai_params_includes_native_gemini_params(self): + """Test that native Gemini imageConfig params are supported""" + supported = self.config.get_supported_openai_params("gemini-3-pro-image-preview") + assert "aspectRatio" in supported + assert "aspect_ratio" in supported + assert "imageSize" in supported + assert "image_size" in supported + assert "imageConfig" in supported + + def test_map_openai_params_aspect_ratio_camel_case(self): + """Test mapping native aspectRatio parameter""" + result = self.config.map_openai_params({"aspectRatio": "9:16"}, {}, "gemini-3-pro-image-preview", False) + assert result["aspectRatio"] == "9:16" + + def test_map_openai_params_aspect_ratio_snake_case(self): + """Test mapping native aspect_ratio parameter""" + result = self.config.map_openai_params({"aspect_ratio": "16:9"}, {}, "gemini-3-pro-image-preview", False) + assert result["aspectRatio"] == "16:9" + + def test_map_openai_params_image_size_camel_case(self): + """Test mapping native imageSize parameter""" + result = self.config.map_openai_params({"imageSize": "4K"}, {}, "gemini-3-pro-image-preview", False) + assert result["imageSize"] == "4K" + + def test_map_openai_params_image_size_snake_case(self): + """Test mapping native image_size parameter""" + result = self.config.map_openai_params({"image_size": "2K"}, {}, "gemini-3-pro-image-preview", False) + assert result["imageSize"] == "2K" + + def test_map_openai_params_image_config_dict_stored_whole(self): + """imageConfig dict is stored as-is so all fields survive""" + result = self.config.map_openai_params( + {"imageConfig": {"aspectRatio": "16:9", "imageSize": "2K"}}, + {}, + "gemini-3.1-flash-image", + False, + ) + assert result["imageConfig"] == {"aspectRatio": "16:9", "imageSize": "2K"} + + def test_map_openai_params_image_config_all_fields(self): + """All ImageConfig fields (personGeneration, imageOutputOptions) pass through""" + payload = { + "imageConfig": { + "aspectRatio": "9:16", + "imageSize": "4K", + "personGeneration": "DONT_ALLOW", + "imageOutputOptions": { + "mimeType": "image/jpeg", + "compressionQuality": 80, + }, + } + } + result = self.config.map_openai_params(payload, {}, "gemini-3.1-flash-image", False) + assert result["imageConfig"] == payload["imageConfig"] + + def test_map_openai_params_image_config_non_dict_warns_and_drops(self): + """Non-dict imageConfig is dropped with a warning, not silently discarded""" + with patch("litellm.llms.vertex_ai.image_generation.vertex_gemini_transformation.verbose_logger") as mock_log: + result = self.config.map_openai_params( + {"imageConfig": "bad-string-value"}, {}, "gemini-3.1-flash-image", False + ) + assert "imageConfig" not in result + mock_log.warning.assert_called_once() + + def test_transform_image_generation_request_from_image_config(self): + """Full imageConfig dict is forwarded verbatim into generationConfig""" + full_config = { + "aspectRatio": "16:9", + "imageSize": "2K", + "personGeneration": "DONT_ALLOW", + "imageOutputOptions": {"mimeType": "image/jpeg", "compressionQuality": 85}, + } + mapped = self.config.map_openai_params( + {"imageConfig": full_config}, + {}, + "gemini-3.1-flash-image", + False, + ) + request = self.config.transform_image_generation_request( + model="gemini-3.1-flash-image", + prompt="A nano banana on a desk", + optional_params=mapped, + litellm_params={}, + headers={}, + ) + assert request["generationConfig"]["imageConfig"] == full_config + + def test_transform_image_generation_flat_params_override_image_config(self): + """Explicit flat params win over the same key inside imageConfig""" + request = self.config.transform_image_generation_request( + model="gemini-3.1-flash-image", + prompt="A nano banana", + optional_params={ + "imageConfig": {"aspectRatio": "1:1", "personGeneration": "DONT_ALLOW"}, + "aspectRatio": "16:9", # should win + }, + litellm_params={}, + headers={}, + ) + assert request["generationConfig"]["imageConfig"]["aspectRatio"] == "16:9" + assert request["generationConfig"]["imageConfig"]["personGeneration"] == "DONT_ALLOW" + + def test_transform_image_generation_request_basic(self): + """Test basic request transformation""" + request = self.config.transform_image_generation_request( + model="gemini-2.5-flash-image", + prompt="A nano banana", + optional_params={}, + litellm_params={}, + headers={}, + ) + assert "contents" in request + assert "generationConfig" in request + assert request["generationConfig"]["responseModalities"] == ["IMAGE"] + assert request["contents"][0]["parts"][0]["text"] == "A nano banana" + + def test_transform_image_generation_request_with_aspect_ratio(self): + """Test request transformation with aspectRatio""" + request = self.config.transform_image_generation_request( + model="gemini-2.5-flash-image", + prompt="A nano banana", + optional_params={"aspectRatio": "16:9"}, + litellm_params={}, + headers={}, + ) + assert request["generationConfig"]["imageConfig"]["aspectRatio"] == "16:9" + + def test_transform_image_generation_request_with_image_size(self): + """Test request transformation with imageSize (Gemini 3 Pro)""" + request = self.config.transform_image_generation_request( + model="gemini-3-pro-image-preview", + prompt="A nano banana", + optional_params={"imageSize": "4K"}, + litellm_params={}, + headers={}, + ) + assert request["generationConfig"]["imageConfig"]["imageSize"] == "4K" + + def test_map_openai_params_web_search_options(self): + """Test web_search_options maps to googleSearch tool""" + result = self.config.map_openai_params({"web_search_options": {}}, {}, "gemini-3.1-flash-image-preview", False) + assert result["tools"] == [{"googleSearch": {}}] + + def test_transform_image_generation_request_with_web_search_tools(self): + """Test request transformation includes googleSearch tools""" + request = self.config.transform_image_generation_request( + model="gemini-3.1-flash-image-preview", + prompt="Generate an image of the latest iPhone", + optional_params={"tools": [{"googleSearch": {}}]}, + litellm_params={}, + headers={}, + ) + assert request["tools"] == [{"googleSearch": {}}] + + def test_transform_image_generation_request_forwards_tool_config(self): + """Test request transformation forwards toolConfig side-effects from tool mapping""" + mapped = self.config.map_openai_params( + {"tools": [{"googleMaps": {"latitude": 37.7, "longitude": -122.4}}]}, + {}, + "gemini-3.1-flash-image-preview", + False, + ) + request = self.config.transform_image_generation_request( + model="gemini-3.1-flash-image-preview", + prompt="Generate an image of a coffee shop nearby", + optional_params=mapped, + litellm_params={}, + headers={}, + ) + assert request["tools"] == [{"googleMaps": {}}] + assert request["toolConfig"] == {"retrievalConfig": {"latLng": {"latitude": 37.7, "longitude": -122.4}}} + + def test_transform_image_generation_request_with_candidate_count(self): + """Test request transformation with candidate_count""" + request = self.config.transform_image_generation_request( + model="gemini-2.5-flash-image", + prompt="A nano banana", + optional_params={"candidate_count": 2}, + litellm_params={}, + headers={}, + ) + assert request["generationConfig"]["candidateCount"] == 2 + + def test_transform_image_generation_request_with_n(self): + """Test request transformation with n parameter""" + request = self.config.transform_image_generation_request( + model="gemini-2.5-flash-image", + prompt="A nano banana", + optional_params={"n": 2}, + litellm_params={}, + headers={}, + ) + assert request["generationConfig"]["candidateCount"] == 2 + + def test_transform_image_generation_response(self): + """Test response transformation""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.json.return_value = { + "candidates": [ + { + "content": { + "parts": [ + { + "inlineData": { + "mimeType": "image/png", + "data": "base64_encoded_image_data", + } + } + ] + } + } + ], + "usageMetadata": { + "promptTokenCount": 93, + "promptTokensDetails": [ + { + "modality": "TEXT", + "tokenCount": 54, + }, + { + "modality": "IMAGE", + "tokenCount": 39, + }, + ], + "candidatesTokenCount": 17, + "totalTokenCount": 110, + }, + } + mock_response.headers = {} + + from litellm.types.utils import ImageResponse + + model_response = ImageResponse() + result = self.config.transform_image_generation_response( + model="gemini-2.5-flash-image", + raw_response=mock_response, + model_response=model_response, + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert len(result.data) == 1 + assert result.data[0].b64_json == "base64_encoded_image_data" + assert result.data[0].url is None + assert result.usage.input_tokens == 93 + assert result.usage.input_tokens_details.text_tokens == 54 + assert result.usage.input_tokens_details.image_tokens == 39 + assert result.usage.output_tokens == 17 + assert result.usage.total_tokens == 110 + + def test_transform_image_generation_response_multiple_images(self): + """Test response transformation with multiple images""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.json.return_value = { + "candidates": [ + { + "content": { + "parts": [ + { + "inlineData": { + "mimeType": "image/png", + "data": "image1", + } + }, + { + "inlineData": { + "mimeType": "image/png", + "data": "image2", + } + }, + ] + } + } + ] + } + mock_response.headers = {} + + from litellm.types.utils import ImageResponse + + model_response = ImageResponse() + result = self.config.transform_image_generation_response( + model="gemini-2.5-flash-image", + raw_response=mock_response, + model_response=model_response, + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert len(result.data) == 2 + assert result.data[0].b64_json == "image1" + assert result.data[1].b64_json == "image2" + + def test_transform_image_generation_response_signature(self): + """Test response transformation includes thoughtSignature for Gemini 3 Pro""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.json.return_value = { + "candidates": [ + { + "content": { + "parts": [ + { + "inlineData": { + "mimeType": "image/png", + "data": "base64_encoded_image_data", + }, + "thoughtSignature": "test_signature_abc123", + } + ] + } + } + ] + } + mock_response.headers = {} + + from litellm.types.utils import ImageResponse + + model_response = ImageResponse() + result = self.config.transform_image_generation_response( + model="gemini-3-pro-image-preview", + raw_response=mock_response, + model_response=model_response, + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert len(result.data) == 1 + assert result.data[0].b64_json == "base64_encoded_image_data" + assert result.data[0].provider_specific_fields["thought_signature"] == "test_signature_abc123" + + def test_transform_image_generation_response_tracks_web_search_requests(self): + """Grounding queries are carried onto usage so search spend can be billed""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.json.return_value = { + "candidates": [ + { + "content": { + "parts": [ + { + "inlineData": { + "mimeType": "image/png", + "data": "base64_encoded_image_data", + } + } + ] + }, + "groundingMetadata": {"webSearchQueries": ["eiffel tower", "paris skyline"]}, + } + ], + "usageMetadata": { + "promptTokenCount": 93, + "candidatesTokenCount": 17, + "totalTokenCount": 110, + }, + } + mock_response.headers = {} + + from litellm.types.utils import ImageResponse + + result = self.config.transform_image_generation_response( + model="gemini-2.5-flash-image", + raw_response=mock_response, + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert result.usage.web_search_requests == 2 + + +class TestVertexAIImagenImageGenerationConfig: + def setup_method(self): + """Set up test fixtures""" + self.config = VertexAIImagenImageGenerationConfig() + + def test_get_supported_openai_params(self): + """Test get_supported_openai_params returns correct params""" + supported = self.config.get_supported_openai_params("imagegeneration@006") + assert "n" in supported + assert "size" in supported + + def test_map_openai_params_n(self): + """Test mapping n parameter to sampleCount""" + non_default_params = {"n": 3} + optional_params = {} + result = self.config.map_openai_params(non_default_params, optional_params, "imagegeneration@006", False) + assert result.get("sampleCount") == 3 + + def test_map_openai_params_size(self): + """Test mapping size parameter to aspectRatio""" + non_default_params = {"size": "1024x1024"} + optional_params = {} + result = self.config.map_openai_params(non_default_params, optional_params, "imagegeneration@006", False) + assert result.get("aspectRatio") == "1:1" + + def test_map_size_to_aspect_ratio(self): + """Test size to aspect ratio mapping""" + assert self.config._map_size_to_aspect_ratio("1024x1024") == "1:1" + assert self.config._map_size_to_aspect_ratio("1792x1024") == "16:9" + assert self.config._map_size_to_aspect_ratio("unknown") == "1:1" # default + + def test_transform_image_generation_request_basic(self): + """Test basic request transformation""" + request = self.config.transform_image_generation_request( + model="imagegeneration@006", + prompt="A cat", + optional_params={}, + litellm_params={}, + headers={}, + ) + assert "instances" in request + assert "parameters" in request + assert request["instances"][0]["prompt"] == "A cat" + assert request["parameters"]["sampleCount"] == 1 + + def test_transform_image_generation_request_with_params(self): + """Test request transformation with parameters""" + request = self.config.transform_image_generation_request( + model="imagegeneration@006", + prompt="A cat", + optional_params={"sampleCount": 2, "aspectRatio": "16:9"}, + litellm_params={}, + headers={}, + ) + assert request["parameters"]["sampleCount"] == 2 + assert request["parameters"]["aspectRatio"] == "16:9" + + def test_transform_image_generation_request_labels_from_metadata(self): + """Billing labels from litellm_params.metadata.requester_metadata on predict body.""" + request = self.config.transform_image_generation_request( + model="imagegeneration@006", + prompt="A cat", + optional_params={}, + litellm_params={"metadata": {"requester_metadata": {"team": "platform", "env": "prod"}}}, + headers={}, + ) + assert request["labels"] == {"team": "platform", "env": "prod"} + assert "labels" not in request["parameters"] + + def test_transform_image_generation_response(self): + """Test response transformation""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.json.return_value = {"predictions": [{"bytesBase64Encoded": "base64_encoded_image_data"}]} + mock_response.headers = {} + + from litellm.types.utils import ImageResponse + + model_response = ImageResponse() + result = self.config.transform_image_generation_response( + model="imagegeneration@006", + raw_response=mock_response, + model_response=model_response, + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert len(result.data) == 1 + assert result.data[0].b64_json == "base64_encoded_image_data" + assert result.data[0].url is None + + def test_transform_image_generation_response_multiple_images(self): + """Test response transformation with multiple images""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.json.return_value = { + "predictions": [ + {"bytesBase64Encoded": "image1"}, + {"bytesBase64Encoded": "image2"}, + ] + } + mock_response.headers = {} + + from litellm.types.utils import ImageResponse + + model_response = ImageResponse() + result = self.config.transform_image_generation_response( + model="imagegeneration@006", + raw_response=mock_response, + model_response=model_response, + logging_obj=MagicMock(), + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert len(result.data) == 2 + assert result.data[0].b64_json == "image1" + assert result.data[1].b64_json == "image2" + + +class TestGetVertexAIImageGenerationConfig: + """Test the router function that selects the correct config""" + + def test_get_gemini_model_config(self): + """Test that Gemini models return Gemini config""" + config = get_vertex_ai_image_generation_config("gemini-2.5-flash-image") + assert isinstance(config, VertexAIGeminiImageGenerationConfig) + + config = get_vertex_ai_image_generation_config("gemini-3-pro-image-preview") + assert isinstance(config, VertexAIGeminiImageGenerationConfig) + + config = get_vertex_ai_image_generation_config("vertex_ai/gemini-2.5-flash-image") + assert isinstance(config, VertexAIGeminiImageGenerationConfig) + + def test_get_imagen_model_config(self): + """Test that Imagen models return Imagen config""" + config = get_vertex_ai_image_generation_config("imagegeneration@006") + assert isinstance(config, VertexAIImagenImageGenerationConfig) + + config = get_vertex_ai_image_generation_config("imagen-4.0-generate-001") + assert isinstance(config, VertexAIImagenImageGenerationConfig) + + config = get_vertex_ai_image_generation_config("vertex_ai/imagegeneration@006") + assert isinstance(config, VertexAIImagenImageGenerationConfig) + + def test_get_non_gemini_model_config(self): + """Test that non-Gemini models default to Imagen config""" + config = get_vertex_ai_image_generation_config("some-other-model") + assert isinstance(config, VertexAIImagenImageGenerationConfig) + + +class TestVertexAIImageGenerationIntegration: + """Integration tests for Vertex AI image generation""" + + + def test_gemini_get_complete_url(self): + """Test Gemini config URL generation""" + config = VertexAIGeminiImageGenerationConfig() + url = config.get_complete_url( + api_base=None, + api_key=None, + model="gemini-2.5-flash-image", + optional_params={}, + litellm_params={ + "vertex_project": "test-project", + "vertex_location": "us-central1", + }, + ) + assert "test-project" in url + assert "us-central1" in url + assert "gemini-2.5-flash-image" in url + assert "generateContent" in url + + def test_imagen_get_complete_url(self): + """Test Imagen config URL generation""" + config = VertexAIImagenImageGenerationConfig() + url = config.get_complete_url( + api_base=None, + api_key=None, + model="imagegeneration@006", + optional_params={}, + litellm_params={ + "vertex_project": "test-project", + "vertex_location": "us-central1", + }, + ) + assert "test-project" in url + assert "us-central1" in url + assert "imagegeneration@006" in url + assert "predict" in url diff --git a/tests/test_litellm/rerank_api/__init__.py b/tests/unit/llms/vertex_ai/rerank/__init__.py similarity index 100% rename from tests/test_litellm/rerank_api/__init__.py rename to tests/unit/llms/vertex_ai/rerank/__init__.py diff --git a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py b/tests/unit/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py rename to tests/unit/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py diff --git a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py b/tests/unit/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py rename to tests/unit/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py diff --git a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_userlabels_e2e.py b/tests/unit/llms/vertex_ai/rerank/test_vertex_ai_rerank_userlabels_e2e.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_userlabels_e2e.py rename to tests/unit/llms/vertex_ai/rerank/test_vertex_ai_rerank_userlabels_e2e.py diff --git a/tests/test_litellm/llms/vertex_ai/test_bge_embedding.py b/tests/unit/llms/vertex_ai/test_bge_embedding.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_bge_embedding.py rename to tests/unit/llms/vertex_ai/test_bge_embedding.py diff --git a/tests/test_litellm/llms/vertex_ai/test_bge_response_transformation.py b/tests/unit/llms/vertex_ai/test_bge_response_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_bge_response_transformation.py rename to tests/unit/llms/vertex_ai/test_bge_response_transformation.py diff --git a/tests/test_litellm/llms/vertex_ai/test_gemini_batch_embeddings.py b/tests/unit/llms/vertex_ai/test_gemini_batch_embeddings.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_gemini_batch_embeddings.py rename to tests/unit/llms/vertex_ai/test_gemini_batch_embeddings.py diff --git a/tests/test_litellm/llms/vertex_ai/test_gemini_empty_properties.py b/tests/unit/llms/vertex_ai/test_gemini_empty_properties.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_gemini_empty_properties.py rename to tests/unit/llms/vertex_ai/test_gemini_empty_properties.py diff --git a/tests/test_litellm/llms/vertex_ai/test_gemini_header_forwarding.py b/tests/unit/llms/vertex_ai/test_gemini_header_forwarding.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_gemini_header_forwarding.py rename to tests/unit/llms/vertex_ai/test_gemini_header_forwarding.py diff --git a/tests/test_litellm/llms/vertex_ai/test_http_status_201.py b/tests/unit/llms/vertex_ai/test_http_status_201.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_http_status_201.py rename to tests/unit/llms/vertex_ai/test_http_status_201.py diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex.py b/tests/unit/llms/vertex_ai/test_vertex.py similarity index 97% rename from tests/test_litellm/llms/vertex_ai/test_vertex.py rename to tests/unit/llms/vertex_ai/test_vertex.py index e3007bac7f3..ab8bf123ab2 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex.py +++ b/tests/unit/llms/vertex_ai/test_vertex.py @@ -1193,7 +1193,6 @@ def test_logprobs(): def test_process_gemini_media(): """Test the _process_gemini_media function for different image sources""" - from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_media from litellm.types.llms.vertex_ai import FileDataType # Test GCS URI @@ -1271,7 +1270,6 @@ def test_process_gemini_media(): assert base64_result["inline_data"]["data"] == "/9j/4AAQSkZJRg..." - def test_get_image_mime_type_from_url(): """Test the _get_image_mime_type_from_url function for different image URLs""" from litellm.llms.vertex_ai.gemini.transformation import ( @@ -1372,46 +1370,6 @@ def encoded_images(): return [encode_image_to_base64(path) for path in image_paths] -@pytest.fixture -def mock_convert_url_to_base64(): - with patch( - "litellm.litellm_core_utils.prompt_templates.factory.convert_url_to_base64", - ) as mock: - # Setup the mock to return a valid image object - mock.return_value = "data:image/jpeg;base64,/9j/4AAQSkZJRg..." - yield mock - - -@pytest.fixture -def mock_blob(): - return Mock(spec=BlobType) - - -@pytest.mark.parametrize( - "http_url", - [ - "http://img1.etsystatic.com/260/0/7813604/il_fullxfull.4226713999_q86e.jpg", - "http://example.com/image.jpg", - "http://subdomain.domain.com/path/to/image.png", - ], -) -def test_process_gemini_media_http_url( - http_url: str, mock_convert_url_to_base64: Mock, mock_blob: Mock -) -> None: - """ - Test that _process_gemini_media correctly handles HTTP URLs. - - Args: - http_url: Test HTTP URL - mock_convert_to_anthropic: Mocked convert_to_anthropic_image_obj function - mock_blob: Mocked BlobType instance - - Vertex AI supports image urls. Ensure no network requests are made. - """ - expected_image_data = "data:image/jpeg;base64,/9j/4AAQSkZJRg..." - mock_convert_url_to_base64.return_value = expected_image_data - # Act - result = _process_gemini_media(http_url) # assert result["file_data"]["file_uri"] == http_url diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_batch_transformation.py b/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_vertex_ai_batch_transformation.py rename to tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/unit/llms/vertex_ai/test_vertex_ai_common_utils.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py rename to tests/unit/llms/vertex_ai/test_vertex_ai_common_utils.py diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_psc_endpoint_support.py b/tests/unit/llms/vertex_ai/test_vertex_ai_psc_endpoint_support.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_vertex_ai_psc_endpoint_support.py rename to tests/unit/llms/vertex_ai/test_vertex_ai_psc_endpoint_support.py diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py b/tests/unit/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py rename to tests/unit/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py b/tests/unit/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py rename to tests/unit/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_global_url_support.py b/tests/unit/llms/vertex_ai/test_vertex_global_url_support.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_vertex_global_url_support.py rename to tests/unit/llms/vertex_ai/test_vertex_global_url_support.py diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_image_generation.py b/tests/unit/llms/vertex_ai/test_vertex_image_generation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_vertex_image_generation.py rename to tests/unit/llms/vertex_ai/test_vertex_image_generation.py diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py b/tests/unit/llms/vertex_ai/test_vertex_llm_base.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py rename to tests/unit/llms/vertex_ai/test_vertex_llm_base.py diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py b/tests/unit/llms/vertex_ai/test_vertex_model_garden_openapi.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py rename to tests/unit/llms/vertex_ai/test_vertex_model_garden_openapi.py diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_passthrough_logging_handler.py b/tests/unit/llms/vertex_ai/test_vertex_passthrough_logging_handler.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/test_vertex_passthrough_logging_handler.py rename to tests/unit/llms/vertex_ai/test_vertex_passthrough_logging_handler.py diff --git a/tests/test_litellm/types/__init__.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/__init__.py similarity index 100% rename from tests/test_litellm/types/__init__.py rename to tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/__init__.py diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_anthropic_image_url_handling.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_anthropic_image_url_handling.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_anthropic_image_url_handling.py rename to tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_anthropic_image_url_handling.py diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py rename to tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py rename to tests/unit/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py diff --git a/tests/test_litellm/types/proxy/__init__.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/gemma/__init__.py similarity index 100% rename from tests/test_litellm/types/proxy/__init__.py rename to tests/unit/llms/vertex_ai/vertex_ai_partner_models/gemma/__init__.py diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py rename to tests/unit/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py diff --git a/tests/test_litellm/types/proxy/policy_engine/__init__.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/gpt_oss/__init__.py similarity index 100% rename from tests/test_litellm/types/proxy/policy_engine/__init__.py rename to tests/unit/llms/vertex_ai/vertex_ai_partner_models/gpt_oss/__init__.py diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gpt_oss/test_vertex_ai_gpt_oss_transformation.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/gpt_oss/test_vertex_ai_gpt_oss_transformation.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gpt_oss/test_vertex_ai_gpt_oss_transformation.py rename to tests/unit/llms/vertex_ai/vertex_ai_partner_models/gpt_oss/test_vertex_ai_gpt_oss_transformation.py diff --git a/tests/test_litellm/vector_stores/__init__.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/qwen/__init__.py similarity index 100% rename from tests/test_litellm/vector_stores/__init__.py rename to tests/unit/llms/vertex_ai/vertex_ai_partner_models/qwen/__init__.py diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/qwen/test_vertex_ai_qwen_global_endpoint.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/qwen/test_vertex_ai_qwen_global_endpoint.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/qwen/test_vertex_ai_qwen_global_endpoint.py rename to tests/unit/llms/vertex_ai/vertex_ai_partner_models/qwen/test_vertex_ai_qwen_global_endpoint.py diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py similarity index 100% rename from tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py rename to tests/unit/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py diff --git a/tests/test_litellm/llms/volcengine/embedding/__init__.py b/tests/unit/llms/volcengine/embedding/__init__.py similarity index 100% rename from tests/test_litellm/llms/volcengine/embedding/__init__.py rename to tests/unit/llms/volcengine/embedding/__init__.py diff --git a/tests/test_litellm/llms/volcengine/test_volcengine.py b/tests/unit/llms/volcengine/test_volcengine.py similarity index 100% rename from tests/test_litellm/llms/volcengine/test_volcengine.py rename to tests/unit/llms/volcengine/test_volcengine.py diff --git a/tests/test_litellm/videos/__init__.py b/tests/unit/llms/wandb/__init__.py similarity index 100% rename from tests/test_litellm/videos/__init__.py rename to tests/unit/llms/wandb/__init__.py diff --git a/tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py b/tests/unit/llms/wandb/test_wandb_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/wandb/test_wandb_chat_transformation.py rename to tests/unit/llms/wandb/test_wandb_chat_transformation.py diff --git a/tests/test_litellm/llms/xai/test_xai_audio_transcription_transformation.py b/tests/unit/llms/xai/test_xai_audio_transcription_transformation.py similarity index 100% rename from tests/test_litellm/llms/xai/test_xai_audio_transcription_transformation.py rename to tests/unit/llms/xai/test_xai_audio_transcription_transformation.py diff --git a/tests/test_litellm/llms/xai/test_xai_chat_transformation.py b/tests/unit/llms/xai/test_xai_chat_transformation.py similarity index 100% rename from tests/test_litellm/llms/xai/test_xai_chat_transformation.py rename to tests/unit/llms/xai/test_xai_chat_transformation.py diff --git a/tests/test_litellm/llms/xai/test_xai_cost_calculator.py b/tests/unit/llms/xai/test_xai_cost_calculator.py similarity index 100% rename from tests/test_litellm/llms/xai/test_xai_cost_calculator.py rename to tests/unit/llms/xai/test_xai_cost_calculator.py diff --git a/tests/test_litellm/llms/xai/test_xai_key_fallback.py b/tests/unit/llms/xai/test_xai_key_fallback.py similarity index 100% rename from tests/test_litellm/llms/xai/test_xai_key_fallback.py rename to tests/unit/llms/xai/test_xai_key_fallback.py diff --git a/tests/test_litellm/llms/xai/test_xai_model_registry.py b/tests/unit/llms/xai/test_xai_model_registry.py similarity index 100% rename from tests/test_litellm/llms/xai/test_xai_model_registry.py rename to tests/unit/llms/xai/test_xai_model_registry.py diff --git a/tests/test_litellm/llms/xai/test_xai_oauth.py b/tests/unit/llms/xai/test_xai_oauth.py similarity index 100% rename from tests/test_litellm/llms/xai/test_xai_oauth.py rename to tests/unit/llms/xai/test_xai_oauth.py diff --git a/tests/unit/messages/test_dispatch.py b/tests/unit/messages/test_dispatch.py index 88ef849f0e2..3d5059b200f 100644 --- a/tests/unit/messages/test_dispatch.py +++ b/tests/unit/messages/test_dispatch.py @@ -22,6 +22,8 @@ from litellm.rust_bridge.messages.entrypoints import ( NativeMessages, ) from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse +from pydantic import TypeAdapter +from litellm.messages import dispatch MESSAGES: Final = [{"role": "user", "content": "hi"}] PYTHON_RULES: Final[Rules] = () @@ -284,3 +286,137 @@ async def test_anthropic_acreate_routes_through_dispatch(monkeypatch: pytest.Mon NATIVE_AMESSAGES.reset() assert result is expected assert [request.model for request in captured] == ["claude-sonnet-4-5"] + + +@pytest.mark.asyncio +async def test_public_anthropic_messages_keeps_the_python_result() -> None: + response: Final = await litellm.anthropic_messages( + model="anthropic/claude-sonnet-4-5", messages=MESSAGES, max_tokens=10, mock_response="ok" + ) + + assert isinstance(response, dict) + content: Final = TypeAdapter(list[dict[str, object]]).validate_python(response.get("content", [])) + assert content[0]["text"] == "ok" + + +def test_sync_messages_request_projects_public_arguments() -> None: + rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),) + expected: Final = AnthropicMessagesResponse(model="claude-test") + + def native( + request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] + ) -> AnthropicMessagesResponse: + assert request.model == "claude-test" + assert request.messages == MESSAGES + assert request.max_tokens == 10 + assert request.custom_llm_provider == "anthropic" + return expected + + binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None) + binding.override(native) + response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision + (), + { + "model": "claude-test", + "messages": MESSAGES, + "max_tokens": 10, + "custom_llm_provider": "anthropic", + }, + python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"), + binding=binding, + native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + rules=rules, + ) + + assert response is expected + + +def test_messages_binding_error_delegates_unchanged_to_python() -> None: + rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),) + expected: Final = AnthropicMessagesResponse(model="claude-test") + + def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse: + return expected + + def native( + request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] + ) -> AnthropicMessagesResponse: + pytest.fail("a call without max_tokens cannot project a request and must stay on Python") + + binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None) + binding.override(native) + response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision + (), + {"model": "claude-test", "messages": MESSAGES, "custom_llm_provider": "anthropic"}, + python=python, + binding=binding, + native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + rules=rules, + ) + + assert response is expected + + +@pytest.mark.asyncio +async def test_async_messages_falls_back_after_native_declines() -> None: + from litellm.rust_bridge.bindings import native_exception_types + + native_types: Final = native_exception_types() + if native_types is None: + pytest.skip("native bridge is unavailable") + declined, _ = native_types + expected: Final = AnthropicMessagesResponse(model="claude-test") + rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_OPT_OUT),) + + async def native( + request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] + ) -> AnthropicMessagesResponse: + raise declined("unsupported") + + async def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse: + return expected + + binding: Final[NativeBinding[NativeAmessages]] = NativeBinding("amessages", validate=lambda _: None) + binding.override(native) + response: Final = await dispatch._ADISPATCH.arun( # pyright: ignore[reportPrivateUsage] # test an explicit route decision + (), + {"model": "claude-test", "messages": MESSAGES, "max_tokens": 10}, + python=python, + binding=binding, + native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + rules=rules, + ) + + assert response is expected + + +def test_internal_is_async_marker_bypasses_native() -> None: + rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),) + expected: Final = AnthropicMessagesResponse(model="claude-test") + + def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse: + return expected + + def native( + request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] + ) -> AnthropicMessagesResponse: + pytest.fail("anthropic_messages' inner handler call must stay on Python") + + binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None) + binding.override(native) + response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision + (), + { + "model": "claude-test", + "messages": MESSAGES, + "max_tokens": 10, + "custom_llm_provider": "anthropic", + "is_async": True, + }, + python=python, + binding=binding, + native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + rules=rules, + ) + + assert response is expected diff --git a/tests/test_litellm/rag/test_main.py b/tests/unit/rag/test_main.py similarity index 100% rename from tests/test_litellm/rag/test_main.py rename to tests/unit/rag/test_main.py diff --git a/tests/unit/rerank_api/__init__.py b/tests/unit/rerank_api/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/rerank_api/test_main.py b/tests/unit/rerank_api/test_main.py similarity index 100% rename from tests/test_litellm/rerank_api/test_main.py rename to tests/unit/rerank_api/test_main.py diff --git a/tests/test_litellm/test_a2a_registry_lookup.py b/tests/unit/test_a2a_registry_lookup.py similarity index 100% rename from tests/test_litellm/test_a2a_registry_lookup.py rename to tests/unit/test_a2a_registry_lookup.py diff --git a/tests/test_litellm/test_acompletion_session_reuse_e2e.py b/tests/unit/test_acompletion_session_reuse_e2e.py similarity index 100% rename from tests/test_litellm/test_acompletion_session_reuse_e2e.py rename to tests/unit/test_acompletion_session_reuse_e2e.py diff --git a/tests/test_litellm/test_add_deployment_no_master_key.py b/tests/unit/test_add_deployment_no_master_key.py similarity index 100% rename from tests/test_litellm/test_add_deployment_no_master_key.py rename to tests/unit/test_add_deployment_no_master_key.py diff --git a/tests/test_litellm/test_aembedding_session_reuse_e2e.py b/tests/unit/test_aembedding_session_reuse_e2e.py similarity index 100% rename from tests/test_litellm/test_aembedding_session_reuse_e2e.py rename to tests/unit/test_aembedding_session_reuse_e2e.py diff --git a/tests/test_litellm/test_anthropic_beta_headers_filtering.py b/tests/unit/test_anthropic_beta_headers_filtering.py similarity index 100% rename from tests/test_litellm/test_anthropic_beta_headers_filtering.py rename to tests/unit/test_anthropic_beta_headers_filtering.py diff --git a/tests/test_litellm/test_anthropic_skills_transformation.py b/tests/unit/test_anthropic_skills_transformation.py similarity index 100% rename from tests/test_litellm/test_anthropic_skills_transformation.py rename to tests/unit/test_anthropic_skills_transformation.py diff --git a/tests/test_litellm/test_assert_ci_coverage.py b/tests/unit/test_assert_ci_coverage.py similarity index 100% rename from tests/test_litellm/test_assert_ci_coverage.py rename to tests/unit/test_assert_ci_coverage.py diff --git a/tests/test_litellm/test_assert_workflow_dir_hygiene.py b/tests/unit/test_assert_workflow_dir_hygiene.py similarity index 100% rename from tests/test_litellm/test_assert_workflow_dir_hygiene.py rename to tests/unit/test_assert_workflow_dir_hygiene.py diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/unit/test_audio_transcription_rust_bridge.py similarity index 100% rename from tests/test_litellm/test_audio_transcription_rust_bridge.py rename to tests/unit/test_audio_transcription_rust_bridge.py diff --git a/tests/test_litellm/test_auto_update_price_and_context_window_file.py b/tests/unit/test_auto_update_price_and_context_window_file.py similarity index 100% rename from tests/test_litellm/test_auto_update_price_and_context_window_file.py rename to tests/unit/test_auto_update_price_and_context_window_file.py diff --git a/tests/test_litellm/test_azure_ad_token_credential_resolution.py b/tests/unit/test_azure_ad_token_credential_resolution.py similarity index 100% rename from tests/test_litellm/test_azure_ad_token_credential_resolution.py rename to tests/unit/test_azure_ad_token_credential_resolution.py diff --git a/tests/test_litellm/test_azure_ai_grok_4_3_model_metadata.py b/tests/unit/test_azure_ai_grok_4_3_model_metadata.py similarity index 100% rename from tests/test_litellm/test_azure_ai_grok_4_3_model_metadata.py rename to tests/unit/test_azure_ai_grok_4_3_model_metadata.py diff --git a/tests/test_litellm/test_azure_ai_grok_4_6_model_metadata.py b/tests/unit/test_azure_ai_grok_4_6_model_metadata.py similarity index 100% rename from tests/test_litellm/test_azure_ai_grok_4_6_model_metadata.py rename to tests/unit/test_azure_ai_grok_4_6_model_metadata.py diff --git a/tests/test_litellm/test_baseten_glm_5_3_model_metadata.py b/tests/unit/test_baseten_glm_5_3_model_metadata.py similarity index 100% rename from tests/test_litellm/test_baseten_glm_5_3_model_metadata.py rename to tests/unit/test_baseten_glm_5_3_model_metadata.py diff --git a/tests/test_litellm/test_batch_completion_models_all_responses.py b/tests/unit/test_batch_completion_models_all_responses.py similarity index 100% rename from tests/test_litellm/test_batch_completion_models_all_responses.py rename to tests/unit/test_batch_completion_models_all_responses.py diff --git a/tests/test_litellm/test_bedrock_marengo_embed_3_model_metadata.py b/tests/unit/test_bedrock_marengo_embed_3_model_metadata.py similarity index 100% rename from tests/test_litellm/test_bedrock_marengo_embed_3_model_metadata.py rename to tests/unit/test_bedrock_marengo_embed_3_model_metadata.py diff --git a/tests/test_litellm/test_budget_ratchet_check.py b/tests/unit/test_budget_ratchet_check.py similarity index 100% rename from tests/test_litellm/test_budget_ratchet_check.py rename to tests/unit/test_budget_ratchet_check.py diff --git a/tests/test_litellm/test_chat_ui_responses_session.py b/tests/unit/test_chat_ui_responses_session.py similarity index 100% rename from tests/test_litellm/test_chat_ui_responses_session.py rename to tests/unit/test_chat_ui_responses_session.py diff --git a/tests/test_litellm/test_check_licenses.py b/tests/unit/test_check_licenses.py similarity index 100% rename from tests/test_litellm/test_check_licenses.py rename to tests/unit/test_check_licenses.py diff --git a/tests/test_litellm/test_check_mcp_operation_boundary.py b/tests/unit/test_check_mcp_operation_boundary.py similarity index 100% rename from tests/test_litellm/test_check_mcp_operation_boundary.py rename to tests/unit/test_check_mcp_operation_boundary.py diff --git a/tests/test_litellm/test_check_migrations_no_data_rewrites.py b/tests/unit/test_check_migrations_no_data_rewrites.py similarity index 100% rename from tests/test_litellm/test_check_migrations_no_data_rewrites.py rename to tests/unit/test_check_migrations_no_data_rewrites.py diff --git a/tests/test_litellm/test_check_py310_typing_imports.py b/tests/unit/test_check_py310_typing_imports.py similarity index 100% rename from tests/test_litellm/test_check_py310_typing_imports.py rename to tests/unit/test_check_py310_typing_imports.py diff --git a/tests/test_litellm/test_check_test_quality.py b/tests/unit/test_check_test_quality.py similarity index 100% rename from tests/test_litellm/test_check_test_quality.py rename to tests/unit/test_check_test_quality.py diff --git a/tests/test_litellm/test_check_type_discipline.py b/tests/unit/test_check_type_discipline.py similarity index 100% rename from tests/test_litellm/test_check_type_discipline.py rename to tests/unit/test_check_type_discipline.py diff --git a/tests/test_litellm/test_circleci_path_filter.py b/tests/unit/test_circleci_path_filter.py similarity index 100% rename from tests/test_litellm/test_circleci_path_filter.py rename to tests/unit/test_circleci_path_filter.py diff --git a/tests/test_litellm/test_circleci_rust_toolchain.py b/tests/unit/test_circleci_rust_toolchain.py similarity index 100% rename from tests/test_litellm/test_circleci_rust_toolchain.py rename to tests/unit/test_circleci_rust_toolchain.py diff --git a/tests/test_litellm/test_claude_fable_5_config.py b/tests/unit/test_claude_fable_5_config.py similarity index 100% rename from tests/test_litellm/test_claude_fable_5_config.py rename to tests/unit/test_claude_fable_5_config.py diff --git a/tests/test_litellm/test_claude_opus_4_6_config.py b/tests/unit/test_claude_opus_4_6_config.py similarity index 100% rename from tests/test_litellm/test_claude_opus_4_6_config.py rename to tests/unit/test_claude_opus_4_6_config.py diff --git a/tests/test_litellm/test_claude_opus_4_8_config.py b/tests/unit/test_claude_opus_4_8_config.py similarity index 100% rename from tests/test_litellm/test_claude_opus_4_8_config.py rename to tests/unit/test_claude_opus_4_8_config.py diff --git a/tests/test_litellm/test_claude_opus_5_config.py b/tests/unit/test_claude_opus_5_config.py similarity index 100% rename from tests/test_litellm/test_claude_opus_5_config.py rename to tests/unit/test_claude_opus_5_config.py diff --git a/tests/test_litellm/test_claude_sonnet_5_config.py b/tests/unit/test_claude_sonnet_5_config.py similarity index 100% rename from tests/test_litellm/test_claude_sonnet_5_config.py rename to tests/unit/test_claude_sonnet_5_config.py diff --git a/tests/test_litellm/test_cloudflare_workers_ai_model_metadata.py b/tests/unit/test_cloudflare_workers_ai_model_metadata.py similarity index 100% rename from tests/test_litellm/test_cloudflare_workers_ai_model_metadata.py rename to tests/unit/test_cloudflare_workers_ai_model_metadata.py diff --git a/tests/test_litellm/test_completion_timeout_resolution.py b/tests/unit/test_completion_timeout_resolution.py similarity index 100% rename from tests/test_litellm/test_completion_timeout_resolution.py rename to tests/unit/test_completion_timeout_resolution.py diff --git a/tests/test_litellm/test_component_entrypoint.py b/tests/unit/test_component_entrypoint.py similarity index 100% rename from tests/test_litellm/test_component_entrypoint.py rename to tests/unit/test_component_entrypoint.py diff --git a/tests/unit/test_compression.py b/tests/unit/test_compression.py new file mode 100644 index 00000000000..be718f03963 --- /dev/null +++ b/tests/unit/test_compression.py @@ -0,0 +1,649 @@ +""" +Unit tests for litellm.compress(). +""" + +import importlib + +import pytest + +import litellm +from litellm.compression.scoring.bm25 import bm25_score_messages +from litellm.compression.scoring.embedding_scorer import embedding_score_messages +from litellm.compression.content_detection import detect_content_type +from litellm.compression.message_stubbing import extract_key, stub_message +from litellm.compression.retrieval_tool import build_retrieval_tool +from litellm.types.utils import CallTypes + +CALL_TYPE = CallTypes.completion +ANTHROPIC_CALL_TYPE = CallTypes.anthropic_messages + + +# --------------------------------------------------------------------------- +# BM25 scorer +# --------------------------------------------------------------------------- + + +def test_bm25_relevance_ranking(): + query = "Fix the authentication bug in the login handler" + messages = [ + { + "role": "user", + "content": "def login_handler(): authentication check bug fix", + }, + {"role": "user", "content": "def render_template(name): css styling layout"}, + {"role": "user", "content": "def verify(): authentication token bug handler"}, + ] + scores = bm25_score_messages(query, messages) + # Messages sharing query terms should score higher than unrelated ones + assert scores[0] > scores[1] + assert scores[2] > scores[1] + + +def test_bm25_empty_query(): + scores = bm25_score_messages("", [{"role": "user", "content": "hello"}]) + assert scores == [0.0] + + +def test_bm25_empty_messages(): + scores = bm25_score_messages("query", []) + assert scores == [] + + +def test_bm25_empty_content(): + scores = bm25_score_messages("query", [{"role": "user", "content": ""}]) + assert scores == [0.0] + + +# --------------------------------------------------------------------------- +# Content detection +# --------------------------------------------------------------------------- + + +def test_detect_code(): + code = """ +import os +from pathlib import Path + +def main(): + class Foo: + pass + return Foo() +""" + assert detect_content_type(code) == "code" + + +def test_detect_json(): + assert detect_content_type('{"key": "value", "num": 42}') == "json" + assert detect_content_type("[1, 2, 3]") == "json" + + +def test_detect_text(): + assert detect_content_type("This is a plain text paragraph about dogs.") == "text" + + +def test_detect_empty(): + assert detect_content_type("") == "text" + + +# --------------------------------------------------------------------------- +# Message stubbing +# --------------------------------------------------------------------------- + + +def test_extract_key_with_filename(): + msg = {"role": "user", "content": "# auth.py\ndef authenticate():\n pass"} + used: set = set() + key = extract_key(msg, fallback_index=0, used_keys=used) + assert key == "auth.py" + + +def test_extract_key_fallback(): + msg = {"role": "user", "content": "Some random content without a filename"} + used: set = set() + key = extract_key(msg, fallback_index=5, used_keys=used) + assert key == "message_5" + + +def test_extract_key_duplicates(): + used: set = set() + msg = {"role": "user", "content": "# auth.py\ncode here"} + k1 = extract_key(msg, fallback_index=0, used_keys=used) + k2 = extract_key(msg, fallback_index=1, used_keys=used) + assert k1 == "auth.py" + assert k2 == "auth.py_2" + + +def test_stub_message(): + msg = {"role": "user", "content": "line1\nline2\nline3"} + stubbed = stub_message(msg, "test_key") + assert stubbed["role"] == "user" + assert "test_key" in stubbed["content"] + assert "litellm_content_retrieve" in stubbed["content"] + assert "3 lines" in stubbed["content"] + + +# --------------------------------------------------------------------------- +# Retrieval tool +# --------------------------------------------------------------------------- + + +def test_retrieval_tool_schema(): + tool = build_retrieval_tool(["auth.py", "utils.py"]) + assert tool["type"] == "function" + assert tool["function"]["name"] == "litellm_content_retrieve" + assert "key" in tool["function"]["parameters"]["properties"] + assert tool["function"]["parameters"]["properties"]["key"]["enum"] == [ + "auth.py", + "utils.py", + ] + assert tool["function"]["parameters"]["required"] == ["key"] + + +def test_retrieval_tool_description_lists_keys(): + tool = build_retrieval_tool(["foo.py", "bar.js"]) + desc = tool["function"]["description"] + assert "foo.py" in desc + assert "bar.js" in desc + + +# --------------------------------------------------------------------------- +# compress() — end-to-end +# --------------------------------------------------------------------------- + + +def test_compress_below_trigger_passthrough(): + messages = [{"role": "user", "content": "hello"}] + result = litellm.compress(messages, model="gpt-4o", call_type=CALL_TYPE) + assert result["messages"] == messages + assert result["cache"] == {} + assert result["tools"] == [] + assert result["compression_ratio"] == 0.0 + assert result["original_tokens"] == result["compressed_tokens"] + + +def test_compress_above_trigger(): + big_messages = [ + {"role": "system", "content": "You are a coding assistant."}, + { + "role": "user", + "content": "# auth.py\n" + "def authenticate():\n pass\n" * 2000, + }, + { + "role": "user", + "content": "# utils.py\n" + "def helper():\n pass\n" * 2000, + }, + { + "role": "user", + "content": "# readme.md\n" + "This is documentation. " * 2000, + }, + {"role": "user", "content": "Fix the bug in auth.py"}, + ] + + result = litellm.compress( + big_messages, + model="gpt-4o", + call_type=CALL_TYPE, + compression_trigger=1000, + compression_target=500, + ) + + assert result["compressed_tokens"] < result["original_tokens"] + assert result["compression_ratio"] > 0 + assert len(result["cache"]) > 0 + assert len(result["tools"]) == 1 + assert result["tools"][0]["function"]["name"] == "litellm_content_retrieve" + + +def test_compress_anthropic_list_content_is_boundary_stable(): + messages = [ + {"role": "system", "content": [{"type": "text", "text": "System prompt"}]}, + { + "role": "user", + "content": [ + {"type": "text", "text": "# a.py\n" + "alpha " * 2000}, + { + "type": "image_url", + "image_url": {"url": "https://example.com/a.png"}, + }, + ], + }, + { + "role": "user", + "content": [ + {"type": "text", "text": "# b.py\n" + "beta " * 2000}, + { + "type": "image_url", + "image_url": {"url": "https://example.com/b.png"}, + }, + ], + }, + { + "role": "user", + "content": [{"type": "text", "text": "Fix alpha bug in a.py"}], + }, + ] + + result = litellm.compress( + messages=messages, + model="claude-sonnet-4-20250514", + call_type=ANTHROPIC_CALL_TYPE, + compression_trigger=1000, + compression_target=500, + ) + + assert result["compressed_tokens"] < result["original_tokens"] + assert len(result["messages"]) == len(messages) + assert [m["role"] for m in result["messages"]] == [m["role"] for m in messages] + assert len(result["cache"]) > 0 + assert len(result["tools"]) == 1 + assert result["tools"][0]["type"] == "custom" + assert result["tools"][0]["name"] == "litellm_content_retrieve" + assert "input_schema" in result["tools"][0] + + +def test_compress_preserves_system_message(): + messages = [ + {"role": "system", "content": "System prompt. " * 500}, + {"role": "user", "content": "Large file content. " * 5000}, + {"role": "user", "content": "Fix the bug"}, + ] + result = litellm.compress( + messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000 + ) + assert result["messages"][0]["role"] == "system" + assert "System prompt" in result["messages"][0]["content"] + + +def test_compress_preserves_last_user_message(): + messages = [ + {"role": "user", "content": "Big context " * 5000}, + {"role": "user", "content": "Fix the bug in auth.py"}, + ] + result = litellm.compress( + messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000 + ) + last_user = [m for m in result["messages"] if m["role"] == "user"][-1] + assert "Fix the bug in auth.py" in last_user["content"] + + +def test_compress_preserves_last_assistant_message(): + messages = [ + {"role": "user", "content": "Big context " * 5000}, + {"role": "assistant", "content": "I'll help with that. " * 2000}, + {"role": "user", "content": "Now fix the bug"}, + ] + result = litellm.compress( + messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000 + ) + assistant_msgs = [m for m in result["messages"] if m["role"] == "assistant"] + assert len(assistant_msgs) >= 1 + # The last assistant message should be preserved (not stubbed) + last_assistant = assistant_msgs[-1] + assert "I'll help with that" in last_assistant["content"] + + +def test_cache_keys_match_stubs(): + messages = [ + {"role": "user", "content": "# auth.py\n" + "code " * 5000}, + {"role": "user", "content": "Fix it"}, + ] + result = litellm.compress( + messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000 + ) + if result["tools"]: + tool_desc = result["tools"][0]["function"]["description"] + for key in result["cache"]: + assert key in tool_desc + + +def test_compress_default_target(): + """compression_target defaults to compression_trigger // 2.""" + messages = [ + {"role": "user", "content": "content " * 5000}, + {"role": "user", "content": "query"}, + ] + result = litellm.compress( + messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=2000 + ) + # Should have compressed — target = 1000 + assert result["compressed_tokens"] <= result["original_tokens"] + + +def test_compress_nested_tool_result_extracts_text_only(): + messages = [ + {"role": "system", "content": [{"type": "text", "text": "System rules"}]}, + { + "role": "user", + "content": [ + {"type": "text", "text": "prefix"}, + { + "type": "tool_result", + "tool_use_id": "toolu_1", + "content": [ + {"type": "text", "text": "nested text fragment"}, + { + "type": "image_url", + "image_url": { + "url": "https://example.com/secret-tool.png", + }, + }, + ], + }, + { + "type": "image_url", + "image_url": {"url": "https://example.com/top.png"}, + }, + {"type": "text", "text": " " + ("irrelevant " * 3000)}, + ], + }, + { + "role": "user", + "content": [{"type": "text", "text": "final query that must remain"}], + }, + ] + + result = litellm.compress( + messages=messages, + model="claude-sonnet-4-20250514", + call_type=ANTHROPIC_CALL_TYPE, + compression_trigger=500, + compression_target=100, + ) + + cached_text = " ".join(result["cache"].values()) + assert "nested text fragment" in cached_text + assert "https://example.com/secret-tool.png" not in cached_text + assert "https://example.com/top.png" not in cached_text + + +def test_compress_default_call_type_is_completion(): + result = litellm.compress( + messages=[ + {"role": "user", "content": "Large context " * 4000}, + {"role": "user", "content": "query"}, + ], + model="gpt-4o", + compression_trigger=1000, + compression_target=500, + ) + + assert result["compressed_tokens"] <= result["original_tokens"] + assert isinstance(result["tools"], list) + + +def test_compress_forwards_embedding_model_params(monkeypatch): + captured = {} + + def fake_embedding_score_messages( + query, messages, model, cache=None, embedding_model_params=None + ): + captured["query"] = query + captured["model"] = model + captured["embedding_model_params"] = embedding_model_params + return [0.0] * len(messages) + + monkeypatch.setattr( + "litellm.compression.scoring.embedding_scorer.embedding_score_messages", + fake_embedding_score_messages, + ) + + result = litellm.compress( + messages=[ + {"role": "user", "content": "Authentication code " * 2000}, + {"role": "user", "content": "Fix auth"}, + ], + model="gpt-4o", + call_type=CALL_TYPE, + compression_trigger=1000, + embedding_model="text-embedding-3-small", + embedding_model_params={"api_base": "https://example-embeddings.test"}, + ) + + assert result["compressed_tokens"] <= result["original_tokens"] + assert captured["model"] == "text-embedding-3-small" + assert captured["embedding_model_params"] == { + "api_base": "https://example-embeddings.test" + } + + +def test_embedding_scorer_forwards_embedding_model_params(monkeypatch): + captured = {} + + class _MockResponse: + data = [ + {"embedding": [1.0, 0.0]}, + {"embedding": [1.0, 0.0]}, + {"embedding": [0.0, 1.0]}, + ] + + def fake_embedding(**kwargs): + captured.update(kwargs) + return _MockResponse() + + monkeypatch.setattr(litellm, "embedding", fake_embedding) + + scores = embedding_score_messages( + query="auth", + messages=[ + {"role": "user", "content": "auth code"}, + {"role": "user", "content": "cooking recipe"}, + ], + model="text-embedding-3-small", + embedding_model_params={"api_base": "https://example-embeddings.test"}, + ) + + assert len(scores) == 2 + assert captured["model"] == "text-embedding-3-small" + assert captured["api_base"] == "https://example-embeddings.test" + + +# --------------------------------------------------------------------------- +# Embedding scorer — integration test (skipped without API key) +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "final_user_message, expected_content", + [ + ("How to cook?", "Unrelated cooking recipes "), + ("Fix auth", "Authentication code "), + ], +) +def test_simple_compression(final_user_message, expected_content): + messages = [ + {"role": "user", "content": "Authentication code " * 2000}, + {"role": "user", "content": "Unrelated cooking recipes " * 2000}, + {"role": "user", "content": final_user_message}, + ] + result = litellm.compress( + messages, model="gpt-4o", call_type=CALL_TYPE, compression_trigger=1000 + ) + if expected_content == "Unrelated cooking recipes ": + assert "Unrelated cooking recipes " in result["messages"][1]["content"] + assert "Authentication code " not in result["messages"][0]["content"] + elif expected_content == "Authentication code ": + assert "Authentication code " in result["messages"][0]["content"] + assert "Unrelated cooking recipes " not in result["messages"][1]["content"] + else: + raise ValueError(f"Unexpected expected_content: {expected_content}") + + +def test_compress_anthropic_drops_irrelevant_tool_exchange_span(monkeypatch): + compress_module = importlib.import_module("litellm.compression.compress") + + def fake_bm25_score_messages(query, messages): + assert "final query" in query + assert len(messages) == 5 + # Prefer idx=0 and de-prioritize the tool exchange span (idx=1,2) + return [0.95, 0.01, 0.02, 0.8, 1.0] + + def fake_token_counter(model, messages=None, text=None): + if messages is not None: + return 1000 + if text is None: + return 0 + if "final query" in text: + return 50 + if "assistant_tail" in text: + return 20 + if "other_blob" in text: + return 220 + if "tool_payload_relevant" in text: + return 200 + if text == "": + return 1 + return 10 + + monkeypatch.setattr( + compress_module, "bm25_score_messages", fake_bm25_score_messages + ) + monkeypatch.setattr(compress_module, "token_counter", fake_token_counter) + + messages = [ + {"role": "user", "content": "other_blob " * 300}, + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_drop", + "name": "litellm_content_retrieve", + "input": {"key": "message_1"}, + } + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_drop", + "content": [{"type": "text", "text": "tool_payload_relevant"}], + } + ], + }, + {"role": "assistant", "content": "assistant_tail"}, + {"role": "user", "content": "final query"}, + ] + + result = litellm.compress( + messages=messages, + model="claude-sonnet-4-20250514", + call_type=ANTHROPIC_CALL_TYPE, + compression_trigger=100, + compression_target=280, + ) + + # idx=1,2 should be dropped atomically (no orphan tool blocks left behind) + assert len(result["messages"]) == 3 + assert result["messages"][0]["role"] == "user" + assert "other_blob" in result["messages"][0]["content"] + assert result["messages"][1]["content"] == "assistant_tail" + assert result["messages"][2]["content"] == "final query" + assert result["cache"] == {} + + +def test_compress_anthropic_keeps_relevant_tool_exchange_span(monkeypatch): + compress_module = importlib.import_module("litellm.compression.compress") + + def fake_bm25_score_messages(query, messages): + assert "final query" in query + assert len(messages) == 5 + # Prefer the tool exchange span over idx=0 + return [0.05, 0.01, 0.92, 0.8, 1.0] + + def fake_token_counter(model, messages=None, text=None): + if messages is not None: + return 1000 + if text is None: + return 0 + if "final query" in text: + return 50 + if "assistant_tail" in text: + return 20 + if "other_blob" in text: + return 220 + if "tool_payload_relevant" in text: + return 200 + if text == "": + return 1 + return 10 + + monkeypatch.setattr( + compress_module, "bm25_score_messages", fake_bm25_score_messages + ) + monkeypatch.setattr(compress_module, "token_counter", fake_token_counter) + + messages = [ + {"role": "user", "content": "other_blob " * 300}, + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_keep", + "name": "litellm_content_retrieve", + "input": {"key": "message_1"}, + } + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_keep", + "content": [{"type": "text", "text": "tool_payload_relevant"}], + } + ], + }, + {"role": "assistant", "content": "assistant_tail"}, + {"role": "user", "content": "final query"}, + ] + + result = litellm.compress( + messages=messages, + model="claude-sonnet-4-20250514", + call_type=ANTHROPIC_CALL_TYPE, + compression_trigger=100, + compression_target=280, + ) + + assert len(result["messages"]) == 5 + assert result["messages"][1]["role"] == "assistant" + assert result["messages"][2]["role"] == "user" + # idx=0 should be compressed instead + assert "litellm_content_retrieve" in result["messages"][0]["content"] + assert len(result["cache"]) == 1 + + +def test_compress_anthropic_malformed_tool_sequence_passes_through(): + messages = [ + {"role": "user", "content": "other_blob " * 300}, + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_broken", + "name": "litellm_content_retrieve", + "input": {"key": "message_1"}, + } + ], + }, + {"role": "user", "content": [{"type": "text", "text": "missing tool_result"}]}, + {"role": "user", "content": "final query"}, + ] + + result = litellm.compress( + messages=messages, + model="claude-sonnet-4-20250514", + call_type=ANTHROPIC_CALL_TYPE, + compression_trigger=100, + compression_target=280, + ) + + assert result["messages"] == messages + assert result["cache"] == {} + assert result["tools"] == [] + assert result["compression_skipped_reason"] == "invalid_anthropic_tool_sequence" diff --git a/tests/test_litellm/test_conftest_isolation.py b/tests/unit/test_conftest_isolation.py similarity index 100% rename from tests/test_litellm/test_conftest_isolation.py rename to tests/unit/test_conftest_isolation.py diff --git a/tests/test_litellm/test_constants.py b/tests/unit/test_constants.py similarity index 100% rename from tests/test_litellm/test_constants.py rename to tests/unit/test_constants.py diff --git a/tests/test_litellm/test_container_router.py b/tests/unit/test_container_router.py similarity index 100% rename from tests/test_litellm/test_container_router.py rename to tests/unit/test_container_router.py diff --git a/tests/test_litellm/test_cost_calculation_log_level.py b/tests/unit/test_cost_calculation_log_level.py similarity index 100% rename from tests/test_litellm/test_cost_calculation_log_level.py rename to tests/unit/test_cost_calculation_log_level.py diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/unit/test_cost_calculator.py similarity index 100% rename from tests/test_litellm/test_cost_calculator.py rename to tests/unit/test_cost_calculator.py diff --git a/tests/test_litellm/test_cost_map_guard.py b/tests/unit/test_cost_map_guard.py similarity index 100% rename from tests/test_litellm/test_cost_map_guard.py rename to tests/unit/test_cost_map_guard.py diff --git a/tests/test_litellm/test_count_tokens_public_api.py b/tests/unit/test_count_tokens_public_api.py similarity index 100% rename from tests/test_litellm/test_count_tokens_public_api.py rename to tests/unit/test_count_tokens_public_api.py diff --git a/tests/test_litellm/test_dashscope_image_generation.py b/tests/unit/test_dashscope_image_generation.py similarity index 99% rename from tests/test_litellm/test_dashscope_image_generation.py rename to tests/unit/test_dashscope_image_generation.py index 1dd0b322623..6f91fe9a0e0 100644 --- a/tests/test_litellm/test_dashscope_image_generation.py +++ b/tests/unit/test_dashscope_image_generation.py @@ -2,7 +2,7 @@ Unit tests for DashScope image generation support (qwen-image-2.0, qwen-image-2.0-pro, qwen-image-3.0, qwen-image-3.0-pro). -Run in docker: pytest tests/test_litellm/test_dashscope_image_generation.py -v +Run in docker: pytest tests/unit/test_dashscope_image_generation.py -v """ from unittest.mock import MagicMock, patch diff --git a/tests/test_litellm/test_daybreak_model_metadata.py b/tests/unit/test_daybreak_model_metadata.py similarity index 100% rename from tests/test_litellm/test_daybreak_model_metadata.py rename to tests/unit/test_daybreak_model_metadata.py diff --git a/tests/test_litellm/test_deepseek_model_metadata.py b/tests/unit/test_deepseek_model_metadata.py similarity index 100% rename from tests/test_litellm/test_deepseek_model_metadata.py rename to tests/unit/test_deepseek_model_metadata.py diff --git a/tests/test_litellm/test_default_branch.py b/tests/unit/test_default_branch.py similarity index 100% rename from tests/test_litellm/test_default_branch.py rename to tests/unit/test_default_branch.py diff --git a/tests/test_litellm/test_detect_changes.py b/tests/unit/test_detect_changes.py similarity index 100% rename from tests/test_litellm/test_detect_changes.py rename to tests/unit/test_detect_changes.py diff --git a/tests/test_litellm/test_dockerfile_apk_repository.py b/tests/unit/test_dockerfile_apk_repository.py similarity index 100% rename from tests/test_litellm/test_dockerfile_apk_repository.py rename to tests/unit/test_dockerfile_apk_repository.py diff --git a/tests/test_litellm/test_dockerfile_bedrock_realtime_extra.py b/tests/unit/test_dockerfile_bedrock_realtime_extra.py similarity index 100% rename from tests/test_litellm/test_dockerfile_bedrock_realtime_extra.py rename to tests/unit/test_dockerfile_bedrock_realtime_extra.py diff --git a/tests/test_litellm/test_dockerfile_non_root.py b/tests/unit/test_dockerfile_non_root.py similarity index 100% rename from tests/test_litellm/test_dockerfile_non_root.py rename to tests/unit/test_dockerfile_non_root.py diff --git a/tests/test_litellm/test_drop_params_env_var.py b/tests/unit/test_drop_params_env_var.py similarity index 100% rename from tests/test_litellm/test_drop_params_env_var.py rename to tests/unit/test_drop_params_env_var.py diff --git a/tests/test_litellm/test_e2e_egress_sentinel.py b/tests/unit/test_e2e_egress_sentinel.py similarity index 100% rename from tests/test_litellm/test_e2e_egress_sentinel.py rename to tests/unit/test_e2e_egress_sentinel.py diff --git a/tests/test_litellm/test_eager_tiktoken_load.py b/tests/unit/test_eager_tiktoken_load.py similarity index 100% rename from tests/test_litellm/test_eager_tiktoken_load.py rename to tests/unit/test_eager_tiktoken_load.py diff --git a/tests/test_litellm/test_env_key_doc_gate.py b/tests/unit/test_env_key_doc_gate.py similarity index 100% rename from tests/test_litellm/test_env_key_doc_gate.py rename to tests/unit/test_env_key_doc_gate.py diff --git a/tests/test_litellm/test_exception_exports.py b/tests/unit/test_exception_exports.py similarity index 100% rename from tests/test_litellm/test_exception_exports.py rename to tests/unit/test_exception_exports.py diff --git a/tests/test_litellm/test_exception_header_preservation.py b/tests/unit/test_exception_header_preservation.py similarity index 100% rename from tests/test_litellm/test_exception_header_preservation.py rename to tests/unit/test_exception_header_preservation.py diff --git a/tests/test_litellm/test_exception_mapping_request_attribute.py b/tests/unit/test_exception_mapping_request_attribute.py similarity index 100% rename from tests/test_litellm/test_exception_mapping_request_attribute.py rename to tests/unit/test_exception_mapping_request_attribute.py diff --git a/tests/test_litellm/test_filter_out_litellm_params.py b/tests/unit/test_filter_out_litellm_params.py similarity index 100% rename from tests/test_litellm/test_filter_out_litellm_params.py rename to tests/unit/test_filter_out_litellm_params.py diff --git a/tests/test_litellm/test_fireworks_serverless_model_costs.py b/tests/unit/test_fireworks_serverless_model_costs.py similarity index 100% rename from tests/test_litellm/test_fireworks_serverless_model_costs.py rename to tests/unit/test_fireworks_serverless_model_costs.py diff --git a/tests/test_litellm/test_gate_slot_lock.py b/tests/unit/test_gate_slot_lock.py similarity index 100% rename from tests/test_litellm/test_gate_slot_lock.py rename to tests/unit/test_gate_slot_lock.py diff --git a/tests/test_litellm/test_gemini_3_1_flash_lite_image_pricing.py b/tests/unit/test_gemini_3_1_flash_lite_image_pricing.py similarity index 100% rename from tests/test_litellm/test_gemini_3_1_flash_lite_image_pricing.py rename to tests/unit/test_gemini_3_1_flash_lite_image_pricing.py diff --git a/tests/test_litellm/test_gemini_tts_native_audio_pricing.py b/tests/unit/test_gemini_tts_native_audio_pricing.py similarity index 100% rename from tests/test_litellm/test_gemini_tts_native_audio_pricing.py rename to tests/unit/test_gemini_tts_native_audio_pricing.py diff --git a/tests/test_litellm/test_get_blog_posts.py b/tests/unit/test_get_blog_posts.py similarity index 100% rename from tests/test_litellm/test_get_blog_posts.py rename to tests/unit/test_get_blog_posts.py diff --git a/tests/test_litellm/test_git_hooks.py b/tests/unit/test_git_hooks.py similarity index 100% rename from tests/test_litellm/test_git_hooks.py rename to tests/unit/test_git_hooks.py diff --git a/tests/test_litellm/test_gpt_5_4_model_metadata.py b/tests/unit/test_gpt_5_4_model_metadata.py similarity index 100% rename from tests/test_litellm/test_gpt_5_4_model_metadata.py rename to tests/unit/test_gpt_5_4_model_metadata.py diff --git a/tests/test_litellm/test_gpt_5_5_model_metadata.py b/tests/unit/test_gpt_5_5_model_metadata.py similarity index 100% rename from tests/test_litellm/test_gpt_5_5_model_metadata.py rename to tests/unit/test_gpt_5_5_model_metadata.py diff --git a/tests/test_litellm/test_gpt_image_cost_calculator.py b/tests/unit/test_gpt_image_cost_calculator.py similarity index 100% rename from tests/test_litellm/test_gpt_image_cost_calculator.py rename to tests/unit/test_gpt_image_cost_calculator.py diff --git a/tests/test_litellm/test_gpt_realtime_mode.py b/tests/unit/test_gpt_realtime_mode.py similarity index 100% rename from tests/test_litellm/test_gpt_realtime_mode.py rename to tests/unit/test_gpt_realtime_mode.py diff --git a/tests/test_litellm/test_groq_streaming_encoding.py b/tests/unit/test_groq_streaming_encoding.py similarity index 100% rename from tests/test_litellm/test_groq_streaming_encoding.py rename to tests/unit/test_groq_streaming_encoding.py diff --git a/tests/test_litellm/test_guardrail_exception_status_codes.py b/tests/unit/test_guardrail_exception_status_codes.py similarity index 100% rename from tests/test_litellm/test_guardrail_exception_status_codes.py rename to tests/unit/test_guardrail_exception_status_codes.py diff --git a/tests/test_litellm/test_lazy_imports.py b/tests/unit/test_lazy_imports.py similarity index 100% rename from tests/test_litellm/test_lazy_imports.py rename to tests/unit/test_lazy_imports.py diff --git a/tests/test_litellm/test_lint_workflow_diff_gates.py b/tests/unit/test_lint_workflow_diff_gates.py similarity index 100% rename from tests/test_litellm/test_lint_workflow_diff_gates.py rename to tests/unit/test_lint_workflow_diff_gates.py diff --git a/tests/test_litellm/test_litellm_params_reserved_keys.py b/tests/unit/test_litellm_params_reserved_keys.py similarity index 100% rename from tests/test_litellm/test_litellm_params_reserved_keys.py rename to tests/unit/test_litellm_params_reserved_keys.py diff --git a/tests/test_litellm/test_logging.py b/tests/unit/test_logging.py similarity index 100% rename from tests/test_litellm/test_logging.py rename to tests/unit/test_logging.py diff --git a/tests/test_litellm/test_lowest_latency_zero_tokens.py b/tests/unit/test_lowest_latency_zero_tokens.py similarity index 100% rename from tests/test_litellm/test_lowest_latency_zero_tokens.py rename to tests/unit/test_lowest_latency_zero_tokens.py diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py new file mode 100644 index 00000000000..effc038f85b --- /dev/null +++ b/tests/unit/test_main.py @@ -0,0 +1,4124 @@ +import asyncio +import base64 +from datetime import datetime +import contextlib +import copy +import json +import logging +import os +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Final + +import httpx +import pytest +import respx + + +import urllib.parse +from importlib import import_module +from unittest.mock import MagicMock, patch + +import litellm +from litellm import main as litellm_main +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging +from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices, Usage + + +@pytest.fixture(autouse=True) +def clear_client_cache(): + """ + Clear the HTTP client cache before each test to ensure mocks are used. + This prevents cached real clients from being reused across tests. + """ + cache = getattr(litellm, "in_memory_llm_clients_cache", None) + if cache is not None: + cache.flush_cache() + yield + if cache is not None: + cache.flush_cache() + + +@pytest.fixture(autouse=True) +def add_api_keys_to_env(monkeypatch): + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-api03-1234567890") + monkeypatch.setenv("OPENAI_API_KEY", "sk-openai-api03-1234567890") + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "my-fake-aws-access-key-id") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "my-fake-aws-secret-access-key") + monkeypatch.setenv("AWS_REGION", "us-east-1") + # Keep these transformation tests on the simple access-key path. A leaked + # session token or role/web-identity env var pushes Bedrock auth down a + # different branch and fails before the mocked HTTP client is exercised. + monkeypatch.delenv("AWS_SESSION_TOKEN", raising=False) + monkeypatch.delenv("AWS_ROLE_ARN", raising=False) + monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False) + + +@pytest.fixture +def openai_api_response(): + mock_response_data = { + "id": "chatcmpl-B0W3vmiM78Xkgx7kI7dr7PC949DMS", + "choices": [ + { + "finish_reason": "stop", + "index": 0, + "logprobs": None, + "message": { + "content": "", + "refusal": None, + "role": "assistant", + "audio": None, + "function_call": None, + "tool_calls": None, + }, + } + ], + "created": 1739462947, + "model": "gpt-4o-mini-2024-07-18", + "object": "chat.completion", + "service_tier": "default", + "system_fingerprint": "fp_bd83329f63", + "usage": { + "completion_tokens": 1, + "prompt_tokens": 121, + "total_tokens": 122, + "completion_tokens_details": { + "accepted_prediction_tokens": 0, + "audio_tokens": 0, + "reasoning_tokens": 0, + "rejected_prediction_tokens": 0, + }, + "prompt_tokens_details": {"audio_tokens": 0, "cached_tokens": 0}, + }, + } + + return mock_response_data + + +def test_completion_missing_role(openai_api_response): + from openai import OpenAI + + from litellm.types.utils import ModelResponse + + client = OpenAI(api_key="test_api_key") + + mock_raw_response = MagicMock() + mock_raw_response.headers = { + "x-request-id": "123", + "openai-organization": "org-123", + "x-ratelimit-limit-requests": "100", + "x-ratelimit-remaining-requests": "99", + } + mock_raw_response.parse.return_value = ModelResponse(**openai_api_response) + + print(f"openai_api_response: {openai_api_response}") + + with patch.object( + client.chat.completions.with_raw_response, "create", MagicMock(return_value=mock_raw_response) + ) as mock_create: + litellm.completion( + model="gpt-4o-mini", + messages=[ + {"role": "user", "content": "Hey"}, + { + "content": "", + "tool_calls": [ + { + "id": "call_m0vFJjQmTH1McvaHBPR2YFwY", + "function": { + "arguments": '{"input": "dksjsdkjdhskdjshdskhjkhlk"}', + "name": "tool_name", + }, + "type": "function", + "index": 0, + }, + { + "id": "call_Vw6RaqV2n5aaANXEdp5pYxo2", + "function": { + "arguments": '{"input": "jkljlkjlkjlkjlk"}', + "name": "tool_name", + }, + "type": "function", + "index": 1, + }, + { + "id": "call_hBIKwldUEGlNh6NlSXil62K4", + "function": { + "arguments": '{"input": "jkjlkjlkjlkj;lj"}', + "name": "tool_name", + }, + "type": "function", + "index": 2, + }, + ], + }, + ], + client=client, + ) + + mock_create.assert_called_once() + + +@pytest.mark.parametrize("model", ["gpt-4o-mini"]) +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_url_with_format_param_openai(model, sync_mode): + from openai import AsyncOpenAI, OpenAI + + from litellm import acompletion, completion + + if sync_mode: + client = OpenAI() + else: + client = AsyncOpenAI() + + args = { + "model": model, + "messages": [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": { + "url": "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png", + "format": "image/png", + }, + }, + {"type": "text", "text": "Describe this image"}, + ], + } + ], + } + with patch.object( + client.chat.completions.with_raw_response, "create" + ) as mock_client: + try: + if sync_mode: + response = completion(**args, client=client) + else: + response = await acompletion(**args, client=client) + print(response) + except Exception as e: + print(e) + + mock_client.assert_called() + + print(mock_client.call_args.kwargs) + + json_str = json.dumps(mock_client.call_args.kwargs) + + assert "format" not in json_str + + +def test_bedrock_latency_optimized_inference(): + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + client = HTTPHandler() + with patch.object(client, "post") as mock_post: + try: + response = litellm.completion( + model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + messages=[{"role": "user", "content": "Hello, how are you?"}], + performanceConfig={"latency": "optimized"}, + client=client, + ) + except Exception as e: + print(e) + + mock_post.assert_called_once() + json_data = json.loads(mock_post.call_args.kwargs["data"]) + assert json_data["performanceConfig"]["latency"] == "optimized" + + +@pytest.mark.parametrize( + ("custom_llm_provider", "model", "expected"), + [ + ("anthropic", "claude-sonnet-5", True), + ("bedrock", "us.anthropic.claude-sonnet-5-20260501-v1:0", True), + ("bedrock", "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123", True), + ("bedrock", "us.amazon.nova-2-lite-v1:0", False), + ("vertex_ai", "claude-sonnet-5", True), + ("vertex_ai", "gemini-3.8-flash", False), + ("azure_ai", "claude-sonnet-4-6", True), + ("azure_ai", "gpt-5.6", False), + ("openai", "gpt-5.6", False), + ("gemini", "gemini-3.8-flash", False), + ], +) +def test_is_claude_tool_target(custom_llm_provider: str, model: str, expected: bool): + assert litellm_main._is_claude_tool_target(custom_llm_provider=custom_llm_provider, model=model) is expected + + +@pytest.mark.parametrize("key", ["input_examples", "eager_input_streaming"]) +def test_drop_anthropic_only_tool_keys_strips_tool_and_function_levels(key: str): + tools = [ + {"type": "function", "name": "example_tool", key: True, "function": {"name": "example_tool", key: True}}, + "opaque_tool", + ] + + cleaned = litellm_main._drop_anthropic_only_tool_keys(tools=tools) + + assert cleaned == [ + {"type": "function", "name": "example_tool", "function": {"name": "example_tool"}}, + "opaque_tool", + ] + assert tools[0][key] is True + assert tools[0]["function"][key] is True + + +def test_completion_strips_eager_input_streaming_before_openai(respx_mock: respx.MockRouter, openai_api_response): + api_base: Final = "http://localhost:12346/v1" + mock_route: Final = respx_mock.post(url__regex=rf"{api_base}/chat/completions.*").mock( + return_value=httpx.Response(status_code=200, json=openai_api_response) + ) + + litellm.completion( + model="openai/gpt-5.6", + messages=[{"role": "user", "content": "Write the file"}], + tools=[ + { + "type": "function", + "function": {"name": "write_file", "parameters": {"type": "object", "properties": {}}}, + "eager_input_streaming": True, + } + ], + api_base=api_base, + api_key="fake_openai_api_key", + ) + + assert mock_route.called + sent_tool: Final = json.loads(respx_mock.calls[0].request.content)["tools"][0] + assert "eager_input_streaming" not in sent_tool + assert sent_tool["function"]["name"] == "write_file" + + +def test_custom_provider_with_extra_headers(): + + with patch.object( + litellm.llms.custom_httpx.http_handler.HTTPHandler, "post" + ) as mock_post: + response = litellm.completion( + model="custom/custom", + messages=[{"role": "user", "content": "Hello, how are you?"}], + headers={"X-Custom-Header": "custom-value"}, + api_base="https://example.com/api/v1", + ) + + mock_post.assert_called_once() + assert mock_post.call_args[1]["headers"]["X-Custom-Header"] == "custom-value" + + +def test_custom_provider_with_extra_body(): + + with patch.object( + litellm.llms.custom_httpx.http_handler.HTTPHandler, "post" + ) as mock_post: + response = litellm.completion( + model="custom/custom", + messages=[{"role": "user", "content": "Hello, how are you?"}], + extra_body={ + "X-Custom-BodyValue": "custom-value", + "X-Custom-BodyValue2": "custom-value2", + }, + api_base="https://example.com/api/v1", + ) + mock_post.assert_called_once() + + assert mock_post.call_args[1]["json"]["X-Custom-BodyValue"] == "custom-value" + assert mock_post.call_args[1]["json"] == { + "model": "custom", + "params": { + "prompt": ["Hello, how are you?"], + "max_tokens": None, + "temperature": None, + "top_p": None, + "top_k": None, + }, + "X-Custom-BodyValue": "custom-value", + "X-Custom-BodyValue2": "custom-value2", + } + + # test that extra_body is not passed if not provided + with patch.object( + litellm.llms.custom_httpx.http_handler.HTTPHandler, "post" + ) as mock_post: + response = litellm.completion( + model="custom/custom", + messages=[{"role": "user", "content": "Hello, how are you?"}], + api_base="https://example.com/api/v1", + ) + mock_post.assert_called_once() + assert mock_post.call_args[1]["json"] == { + "model": "custom", + "params": { + "prompt": ["Hello, how are you?"], + "max_tokens": None, + "temperature": None, + "top_p": None, + "top_k": None, + }, + } + + +@pytest.fixture(autouse=True) +def set_openrouter_api_key(): + original_api_key = os.environ.get("OPENROUTER_API_KEY") + os.environ["OPENROUTER_API_KEY"] = "fake-key-for-testing" + yield + if original_api_key is not None: + os.environ["OPENROUTER_API_KEY"] = original_api_key + else: + del os.environ["OPENROUTER_API_KEY"] + + +@pytest.mark.asyncio +async def test_extra_body_with_fallback( + respx_mock: respx.MockRouter, set_openrouter_api_key, monkeypatch +): + """ + test regression for https://github.com/BerriAI/litellm/issues/8425. + + This was perhaps a wider issue with the acompletion function not passing kwargs such as extra_body correctly when fallbacks are specified. + """ + + # Save original state to restore after test + original_disable_aiohttp = litellm.disable_aiohttp_transport + + try: + # since this uses respx, we need to set use_aiohttp_transport to False + # Set both the global variable and environment variable to ensure it takes effect + litellm.disable_aiohttp_transport = True + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + # Flush cache to ensure no stale aiohttp clients are used + litellm.in_memory_llm_clients_cache.flush_cache() + + # Set up test parameters + model = "openrouter/deepseek/deepseek-chat" + messages = [{"role": "user", "content": "Hello, world!"}] + extra_body = { + "provider": { + "order": ["DeepSeek"], + "allow_fallbacks": False, + "require_parameters": True, + } + } + fallbacks = [{"model": "openrouter/google/gemini-flash-1.5-8b"}] + + # Set up mock to respond to any POST request to the OpenRouter endpoint + # This ensures it works for both primary and fallback models + mock_route = respx_mock.post("https://openrouter.ai/api/v1/chat/completions") + mock_route.return_value = httpx.Response( + 200, + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": model, + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "Hello from mocked response!", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 9, + "completion_tokens": 12, + "total_tokens": 21, + }, + }, + ) + + response = await litellm.acompletion( + model=model, + messages=messages, + extra_body=extra_body, + fallbacks=fallbacks, + api_key="fake-openrouter-api-key", + ) + + # Verify the response + assert response is not None + assert ( + len(respx_mock.calls) > 0 + ), "Mock was not called - check if aiohttp transport is properly disabled" + + # Get the request from the mock + request: httpx.Request = respx_mock.calls[0].request + request_body = request.read() + request_body = json.loads(request_body) + + # Verify basic parameters + assert request_body["model"] == "deepseek/deepseek-chat" + assert request_body["messages"] == messages + + # Verify the extra_body parameters remain under the provider key + assert request_body["provider"]["order"] == ["DeepSeek"] + assert request_body["provider"]["allow_fallbacks"] is False + assert request_body["provider"]["require_parameters"] is True + finally: + # Restore original state to prevent test pollution + litellm.disable_aiohttp_transport = original_disable_aiohttp + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.parametrize("env_base", ["OPENAI_BASE_URL", "OPENAI_API_BASE"]) +@pytest.mark.asyncio +@pytest.mark.flaky(retries=3, delay=1) +async def test_openai_env_base( + respx_mock: respx.MockRouter, env_base, openai_api_response, monkeypatch +): + "This tests OpenAI env variables are honored, including legacy OPENAI_API_BASE" + # Ensure aiohttp transport is disabled to use httpx which respx can mock + litellm.disable_aiohttp_transport = True + + expected_base_url = "http://localhost:12345/v1" + + # Assign the environment variable based on env_base, and use a fake API key. + monkeypatch.setenv(env_base, expected_base_url) + monkeypatch.setenv("OPENAI_API_KEY", "fake_openai_api_key") + + model = "gpt-4o" + messages = [{"role": "user", "content": "Hello, how are you?"}] + + # Configure respx mock to intercept the request + mock_route = respx_mock.post( + url__regex=r"http://localhost:12345/v1/chat/completions.*" + ).mock( + return_value=httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": model, + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "Hello from mocked response!", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 9, + "completion_tokens": 12, + "total_tokens": 21, + }, + }, + ) + ) + + try: + response = await litellm.acompletion(model=model, messages=messages) + + # verify we had a response + assert response.choices[0].message.content == "Hello from mocked response!" + + # Verify the mock was called + assert ( + mock_route.called + ), "Mock route was not called - request may have bypassed respx" + finally: + # Clean up to avoid affecting other tests + litellm.disable_aiohttp_transport = False + + +def build_database_url(username, password, host, dbname): + username_enc = urllib.parse.quote_plus(username) + password_enc = urllib.parse.quote_plus(password) + dbname_enc = urllib.parse.quote_plus(dbname) + return f"postgresql://{username_enc}:{password_enc}@{host}/{dbname_enc}" + + +def test_build_database_url(): + url = build_database_url("user@name", "p@ss:word", "localhost", "db/name") + assert url == "postgresql://user%40name:p%40ss%3Aword@localhost/db%2Fname" + + +def test_bedrock_llama(): + litellm._turn_on_debug() + from litellm.types.utils import CallTypes + from litellm.utils import return_raw_request + + model = "bedrock/invoke/us.meta.llama4-scout-17b-instruct-v1:0" + + request = return_raw_request( + endpoint=CallTypes.completion, + kwargs={ + "model": model, + "messages": [ + {"role": "user", "content": "hi"}, + ], + }, + ) + print(request) + + assert ( + request["raw_request_body"]["prompt"] + == "<|begin_of_text|><|start_header_id|>user<|end_header_id|>\n\nhi<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n" + ) + + +def _mocked_openai_chat_response(model: str) -> httpx.Response: + return httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": model, + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "Hello from mocked response!", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 9, + "completion_tokens": 12, + "total_tokens": 21, + }, + }, + ) + + +def test_return_raw_request_does_not_call_provider(respx_mock: respx.MockRouter): + """Regression for #33952: return_raw_request must transform without contacting the provider. + + Previously return_raw_request invoked the real endpoint with a fake key and relied on the + provider rejecting it, which sent an unintended inference request and (in the async proxy + route) blocked the event loop on provider I/O. + """ + from litellm.types.utils import CallTypes + from litellm.utils import return_raw_request + + model = "gpt-4o" + route = respx_mock.post("https://api.openai.com/v1/chat/completions").mock( + return_value=_mocked_openai_chat_response(model) + ) + + request = return_raw_request( + endpoint=CallTypes.completion, + kwargs={ + "model": model, + "messages": [{"role": "user", "content": "hi"}], + }, + ) + + assert route.call_count == 0 + assert request.get("error") is None + assert request["raw_request_body"]["model"] == model + assert request["raw_request_body"]["messages"] == [ + {"role": "user", "content": "hi"} + ] + + +def test_completion_forwards_verbosity_in_raw_request(respx_mock: respx.MockRouter): + """Regression test: completion() must forward the verbosity param to the provider request body.""" + from litellm.types.utils import CallTypes + from litellm.utils import return_raw_request + + model = "gpt-5.2" + messages = [{"role": "user", "content": "hi"}] + respx_mock.post("https://api.openai.com/v1/chat/completions").mock( + return_value=_mocked_openai_chat_response(model) + ) + + request = return_raw_request( + endpoint=CallTypes.completion, + kwargs={ + "model": model, + "messages": messages, + "verbosity": "high", + }, + ) + + assert request["raw_request_body"]["verbosity"] == "high" + assert request["raw_request_body"]["model"] == model + assert request["raw_request_body"]["messages"] == messages + + +@pytest.mark.asyncio +async def test_acompletion_forwards_verbosity_to_provider_request( + respx_mock: respx.MockRouter, monkeypatch +): + """Regression test: acompletion() must forward the verbosity param to the provider request body.""" + original_disable_aiohttp = litellm.disable_aiohttp_transport + try: + litellm.disable_aiohttp_transport = True + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + litellm.in_memory_llm_clients_cache.flush_cache() + + model = "gpt-5.2" + messages = [{"role": "user", "content": "hi"}] + mock_route = respx_mock.post("https://api.openai.com/v1/chat/completions").mock( + return_value=_mocked_openai_chat_response(model) + ) + + response = await litellm.acompletion( + model=model, + messages=messages, + verbosity="low", + api_key="fake-openai-api-key", + ) + + assert response.choices[0].message.content == "Hello from mocked response!" + assert mock_route.called + request_body = json.loads(respx_mock.calls[0].request.read()) + assert request_body["verbosity"] == "low" + assert request_body["model"] == model + assert request_body["messages"] == messages + finally: + litellm.disable_aiohttp_transport = original_disable_aiohttp + litellm.in_memory_llm_clients_cache.flush_cache() + + +def test_responses_api_bridge_check_strips_responses_prefix(): + """Test that responses_api_bridge_check strips 'responses/' prefix and sets mode.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 4096} + + model_info, model = responses_api_bridge_check( + model="responses/gpt-4-responses", + custom_llm_provider="openai", + ) + + assert model == "gpt-4-responses" + assert model_info["mode"] == "responses" + + +def test_responses_api_bridge_check_gpt_5_4_pro(): + """Test that gpt-5.4-pro routes through responses API bridge, not chat completions. + + Regression test for https://github.com/BerriAI/litellm/issues/23014 + gpt-5.4-pro is a responses-only model and must not be sent to /v1/chat/completions. + """ + from litellm.main import responses_api_bridge_check + + for model_name in ["gpt-5.4-pro", "gpt-5.4-pro-2026-03-05"]: + model_info, model = responses_api_bridge_check( + model=model_name, + custom_llm_provider="openai", + ) + assert ( + model_info.get("mode") == "responses" + ), f"{model_name} should have mode='responses', got '{model_info.get('mode')}'" + + +def test_responses_api_bridge_check_gpt_5_4_tools_plus_reasoning_routes_to_responses(): + """gpt-5.4 with both tools and reasoning_effort should route to Responses API.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.4", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort="xhigh", + ) + + assert model == "gpt-5.4" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_gpt_6_astra_tools_with_default_reasoning_routes_to_responses(): + from litellm.main import responses_api_bridge_check + + model_info, model = responses_api_bridge_check( + model="gpt-6-astra", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + ) + + assert model == "gpt-6-astra" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_gpt_5_5_tools_plus_reasoning_routes_to_responses(): + """gpt-5.5+ with both tools and reasoning_effort should route to Responses API.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.5-pro", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort="xhigh", + ) + + assert model == "gpt-5.5-pro" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_azure_gpt_5_4_tools_plus_reasoning_routes_to_responses(): + """Azure gpt-5.4 with both tools and reasoning_effort should route to Responses API.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.4", + custom_llm_provider="azure", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort="high", + ) + + assert model == "gpt-5.4" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_azure_gpt_5_4_tools_with_default_reasoning_routes_to_responses(): + """ + Azure gpt-5.4 with tools and UNSET reasoning_effort must bridge: OpenAI enables + reasoning by default for gpt-5.4+, and Chat Completions rejects function tools + whenever reasoning is on. + """ + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.4", + custom_llm_provider="azure", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + ) + + assert model == "gpt-5.4" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_gpt_5_4_tools_with_default_reasoning_routes_to_responses(): + """ + gpt-5.4 with tools and UNSET reasoning_effort must bridge: OpenAI enables reasoning + by default for gpt-5.4+, and Chat Completions rejects function tools whenever + reasoning is on ("use /v1/responses or set reasoning_effort to 'none'"). + """ + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.4", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + ) + + assert model == "gpt-5.4" + assert model_info.get("mode") == "responses" + + +@pytest.mark.parametrize( + "model_name, expected_mode", + [ + pytest.param("gpt-5.6-sol", "responses", id="above-boundary-bridges"), + pytest.param("gpt-5.1", None, id="below-boundary-stays-chat"), + ], +) +def test_responses_api_bridge_check_gpt_5_6_tools_with_default_reasoning_routes_to_responses( + monkeypatch, model_name, expected_mode +): + """ + gpt-5.6 must bridge on function tools alone. The bridge used to require an explicit + reasoning_effort, so a gpt-5.6 call carrying tools and no effort was rejected with + "Function tools with reasoning_effort are not supported for gpt-5.6-sol in + /v1/chat/completions". + + Paired with a model below the gpt-5.4 boundary, which must still stay on chat. The + gate parses the version and drops any suffix, so the family members bridge + identically and only the boundary distinguishes behaviour. + """ + import litellm + from litellm.main import responses_api_bridge_check + + monkeypatch.delenv("OPENAI_BASE_URL", raising=False) + monkeypatch.delenv("OPENAI_API_BASE", raising=False) + monkeypatch.setattr(litellm, "api_base", None) + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model=model_name, + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + ) + + assert model == model_name + assert model_info.get("mode") == expected_mode + + +def test_responses_api_bridge_check_gpt_5_4_tools_with_reasoning_none_stays_chat(): + """ + Explicit reasoning_effort "none" is OpenAI's documented escape hatch that keeps + function tools servable on Chat Completions; the bridge must not fire. + """ + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.4", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort="none", + ) + + assert model == "gpt-5.4" + assert model_info.get("mode") != "responses" + + +def test_responses_api_bridge_check_reasoning_none_with_summary_still_routes_to_responses(): + """A reasoning summary is Responses-only regardless of effort value.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.4", + custom_llm_provider="openai", + reasoning_effort="none", + reasoning_summary="detailed", + ) + + assert model == "gpt-5.4" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_gpt_5_4_custom_tools_only_stays_chat(): + """ + Chat Completions serves custom (grammar) tools natively with reasoning on; only + FUNCTION tools trigger the OpenAI rejection. Custom-only requests must stay on chat + so responses keep the native custom tool_call shape instead of the bridge's + function-shaped mapping. + """ + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "custom", "custom": {"name": "ApplyPatch", "description": "V4A patch"}}], + reasoning_effort=None, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") != "responses" + + +def test_responses_api_bridge_check_gpt_5_4_mixed_function_and_custom_tools_routes_to_responses(): + """One function tool in the mix is enough to make chat unservable with reasoning on.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[ + {"type": "custom", "custom": {"name": "ApplyPatch"}}, + {"type": "function", "function": {"name": "shell"}}, + ], + reasoning_effort=None, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_gpt_5_4_flat_function_tool_routes_to_responses(): + """Responses-style flat function tool defs still count as function tools.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "name": "shell", "parameters": {"type": "object"}}], + reasoning_effort=None, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") == "responses" + + +@pytest.mark.parametrize( + "custom_llm_provider, model_name, api_base", + [ + pytest.param("openai", "gpt-5.6", None, id="openai"), + pytest.param("azure_ai", "gpt-6-astra", "https://myproject.services.ai.azure.com", id="azure-ai-foundry"), + ], +) +def test_responses_api_bridge_check_function_tool_without_body_stays_chat( + monkeypatch, custom_llm_provider, model_name, api_base +): + import litellm + from litellm.main import responses_api_bridge_check + + monkeypatch.delenv("OPENAI_BASE_URL", raising=False) + monkeypatch.delenv("OPENAI_API_BASE", raising=False) + monkeypatch.setattr(litellm, "api_base", None) + + model_info, model = responses_api_bridge_check( + model=model_name, + custom_llm_provider=custom_llm_provider, + tools=[{"type": "function"}], + reasoning_effort=None, + api_base=api_base, + ) + + assert model == model_name + assert model_info.get("mode") != "responses" + + +def test_responses_api_bridge_check_dict_effort_none_stays_chat(): + """The escape hatch must honor litellm's dict form: {"effort": "none"} means reasoning off.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort={"effort": "none"}, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") != "responses" + + +def test_responses_api_bridge_check_dict_effort_active_routes_to_responses(): + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort={"effort": "low"}, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_dict_effort_none_with_summary_routes_to_responses(): + """A summary inside the dict form is Responses-only even when effort is none.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort={"effort": "none", "summary": "concise"}, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") == "responses" + + +@pytest.mark.parametrize("blank_api_base", [None, "", " ", "\t"]) +def test_responses_api_bridge_check_blank_api_base_is_default_openai(blank_api_base): + """ + A blank api_base (None, empty, or whitespace) resolves to the default OpenAI + endpoint downstream, which enforces the reasoning+tools constraint, so gpt-5.4+ + function-tool requests with unset reasoning_effort must still auto-bridge. + """ + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + api_base=blank_api_base, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_custom_api_base_with_unset_effort_stays_chat(): + """ + Chat-only OpenAI-compatible backends registered under the openai provider with a + custom api_base and gpt-5.4+ model names serve tools-without-reasoning fine and + have no /responses route; the unset-effort arm must not reroute them. + """ + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + api_base="http://vllm.internal:8000/v1", + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") != "responses" + + +def test_responses_api_bridge_check_custom_api_base_via_global_with_unset_effort_stays_chat(monkeypatch): + """ + A custom base set through the litellm.api_base global (not the call arg) is resolved the + same way the chat handler resolves it, so the unset-effort arm must not reroute a chat-only + backend to a /responses route it lacks. Regression guard: the gate previously inspected only + the call-level api_base and bridged these requests. + """ + import litellm + from litellm.main import responses_api_bridge_check + + monkeypatch.setattr(litellm, "api_base", "http://vllm.internal:8000/v1") + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + api_base=None, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") != "responses" + + +@pytest.mark.parametrize("env_var", ["OPENAI_BASE_URL", "OPENAI_API_BASE"]) +def test_responses_api_bridge_check_custom_api_base_via_env_with_unset_effort_stays_chat(monkeypatch, env_var): + """ + A custom base set via OPENAI_BASE_URL/OPENAI_API_BASE env is resolved identically to the chat + handler, so the unset-effort arm leaves the request on chat instead of bridging it. + """ + import litellm + from litellm.main import responses_api_bridge_check + + monkeypatch.setattr(litellm, "api_base", None) + monkeypatch.delenv("OPENAI_BASE_URL", raising=False) + monkeypatch.delenv("OPENAI_API_BASE", raising=False) + monkeypatch.setenv(env_var, "http://vllm.internal:8000/v1") + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + api_base=None, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") != "responses" + + +@pytest.mark.parametrize( + "api_base", + [ + "https://southcentralus.privatelink.api.openai.com/v1", + "https://privatelink.corp.api.openai.com/v1", + "https://api.openai.com:443/v1", + "https://api.openai.com/v1/", + "HTTPS://API.OPENAI.COM/v1", + ], +) +def test_responses_api_bridge_check_openai_backed_custom_api_base_with_unset_effort_routes_to_responses(api_base): + """ + A custom api_base whose host is api.openai.com or a subdomain of it (a PrivateLink hostname, a + port-qualified or trailing-slash default) still reaches the real OpenAI backend, which rejects + function tools with reasoning on Chat Completions, so the unset-effort arm must bridge exactly as + it does for the literal default URL. Regression guard for GH #39353. + """ + from litellm.main import responses_api_bridge_check + + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + api_base=api_base, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") == "responses" + + +@pytest.mark.parametrize( + "api_base", + [ + "https://api.openai.com.evil.example/v1", + "https://notapi.openai.com/v1", + "https://gateway.example/v1?upstream=api.openai.com", + "https://openai.internal.example/api.openai.com/v1", + ], +) +def test_responses_api_bridge_check_lookalike_custom_api_base_with_unset_effort_stays_chat(api_base): + """Only the host decides: api.openai.com appearing elsewhere in the URL is still a foreign backend.""" + from litellm.main import responses_api_bridge_check + + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + api_base=api_base, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") != "responses" + + +def test_responses_api_bridge_check_privatelink_api_base_via_env_with_unset_effort_routes_to_responses(monkeypatch): + """A PrivateLink base set through OPENAI_BASE_URL resolves the way the chat handler's does and still bridges.""" + import litellm + from litellm.main import responses_api_bridge_check + + monkeypatch.setattr(litellm, "api_base", None) + monkeypatch.delenv("OPENAI_API_BASE", raising=False) + monkeypatch.setenv("OPENAI_BASE_URL", "https://southcentralus.privatelink.api.openai.com/v1") + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + api_base=None, + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_custom_api_base_with_explicit_effort_still_routes(): + """Explicit reasoning_effort keeps its pre-existing bridging behavior on any api_base.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.6", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort="high", + api_base="http://vllm.internal:8000/v1", + ) + + assert model == "gpt-5.6" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_azure_with_api_base_and_unset_effort_routes(): + """Azure OpenAI always sets api_base and does enforce the constraint; keep bridging.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.4", + custom_llm_provider="azure", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + api_base="https://myresource.openai.azure.com", + ) + + assert model == "gpt-5.4" + assert model_info.get("mode") == "responses" + + +_FOUNDRY_API_BASE: Final = "https://myproject.services.ai.azure.com" +_FOUNDRY_FUNCTION_TOOL: Final = ({"type": "function", "function": {"name": "get_weather"}},) + + +@pytest.mark.parametrize( + "model_name, api_base, reasoning_effort", + [ + pytest.param("gpt-6-astra", _FOUNDRY_API_BASE, None, id="gpt-6-unset-effort"), + pytest.param("gpt-6-astra", _FOUNDRY_API_BASE, "low", id="gpt-6-explicit-effort"), + pytest.param("gpt-6-astra", "https://myresource.openai.azure.com", None, id="gpt-6-azure-openai-host"), + pytest.param("gpt-5.6-sol", _FOUNDRY_API_BASE, "low", id="gpt-5.6-explicit-effort"), + pytest.param("gpt-5.6-sol", _FOUNDRY_API_BASE, {"effort": "high"}, id="gpt-5.6-explicit-effort-dict"), + ], +) +def test_responses_api_bridge_check_azure_ai_foundry_rejected_tools_route_to_responses( + model_name, api_base, reasoning_effort +): + from litellm.main import responses_api_bridge_check + + model_info, model = responses_api_bridge_check( + model=model_name, + custom_llm_provider="azure_ai", + tools=_FOUNDRY_FUNCTION_TOOL, + reasoning_effort=reasoning_effort, + api_base=api_base, + ) + + assert model == model_name + assert model_info.get("mode") == "responses" + + +@pytest.mark.parametrize( + "model_name, api_base, reasoning_effort", + [ + pytest.param("gpt-6-astra", _FOUNDRY_API_BASE, "none", id="explicit-none-stays-chat"), + pytest.param("gpt-5.6-sol", _FOUNDRY_API_BASE, None, id="gpt-5.6-unset-effort-stays-chat"), + pytest.param("gpt-5.6-sol", _FOUNDRY_API_BASE, "none", id="gpt-5.6-explicit-none-stays-chat"), + pytest.param("gpt-5.5", _FOUNDRY_API_BASE, "high", id="gpt-5.5-explicit-effort-stays-chat"), + pytest.param("gpt-5.4-mini", _FOUNDRY_API_BASE, None, id="gpt-5.4-mini-unset-effort-stays-chat"), + pytest.param("gpt-5.4-mini", _FOUNDRY_API_BASE, "low", id="gpt-5.4-mini-explicit-effort-stays-chat"), + pytest.param("gpt-6-astra", "https://myproject.models.ai.azure.com", None, id="serverless-host-stays-chat"), + pytest.param("Mistral-large-2411", _FOUNDRY_API_BASE, None, id="non-gpt-5-model-stays-chat"), + pytest.param("claude-opus-4-1", _FOUNDRY_API_BASE, None, id="claude-on-foundry-stays-chat"), + ], +) +def test_responses_api_bridge_check_azure_ai_without_foundry_responses_route_stays_chat( + model_name, api_base, reasoning_effort +): + from litellm.main import responses_api_bridge_check + + model_info, model = responses_api_bridge_check( + model=model_name, + custom_llm_provider="azure_ai", + tools=_FOUNDRY_FUNCTION_TOOL, + reasoning_effort=reasoning_effort, + api_base=api_base, + ) + + assert model == model_name + assert model_info.get("mode") != "responses" + + +def test_responses_api_bridge_check_older_gpt_5_tools_without_reasoning_stays_chat(): + """Pre-5.4 GPT-5 names keep the old boundary: tools alone never bridge.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.1", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + ) + + assert model == "gpt-5.1" + assert model_info.get("mode") != "responses" + + +def test_responses_api_bridge_check_gpt_5_4_reasoning_summary_without_tools_routes_to_responses(): + """gpt-5.4+ with reasoning_effort + reasoningSummary but no tools should bridge (AI SDK).""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.4", + custom_llm_provider="openai", + tools=None, + reasoning_effort="medium", + reasoning_summary="auto", + ) + + assert model == "gpt-5.4" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_gpt_5_reasoning_summary_routes_to_responses(): + """Bare ``gpt-5`` with reasoning_effort + reasoningSummary should bridge (not 5.4+).""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5", + custom_llm_provider="openai", + tools=None, + reasoning_effort="medium", + reasoning_summary="auto", + ) + + assert model == "gpt-5" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_gpt_5_tools_without_summary_stays_chat(): + """gpt-5 with tools + reasoning_effort but no summary should stay on chat.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5", + custom_llm_provider="openai", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort="medium", + reasoning_summary=None, + ) + + assert model == "gpt-5" + assert model_info.get("mode") != "responses" + + +@patch("litellm.completion_extras.responses_api_bridge.completion") +def test_gpt_5_4_responses_bridge_preserves_reasoning_summary_dict( + mock_responses_completion, +): + """When routed to Responses, preserve reasoning_effort summary dict.""" + mock_responses_completion.return_value = MagicMock() + + import litellm + + litellm.completion( + model="gpt-5.4", + messages=[{"role": "user", "content": "What is the capital of France?"}], + tools=[ + { + "type": "function", + "function": { + "name": "get_capital", + "description": "Get the capital of a country", + "parameters": { + "type": "object", + "properties": {"country": {"type": "string"}}, + }, + }, + } + ], + reasoning_effort={"effort": "xhigh", "summary": "detailed"}, + api_key="fake-key", + ) + + assert mock_responses_completion.called is True + optional_params = mock_responses_completion.call_args.kwargs["optional_params"] + assert optional_params["reasoning_effort"] == { + "effort": "xhigh", + "summary": "detailed", + } + + +@pytest.mark.parametrize("reasoning_effort", ["high", {"effort": "high"}]) +def test_responses_bridge_preserves_reasoning_effort_with_drop_params( + reasoning_effort, + restore_model_registry, + respx_mock: respx.MockRouter, + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + response_body: Final = { + "id": "resp_test", + "object": "response", + "created_at": 1734366691, + "status": "completed", + "model": "test-responses-bridge", + "output": [ + { + "type": "message", + "id": "msg_1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "Done.", "annotations": []}], + } + ], + "parallel_tool_calls": True, + "usage": { + "input_tokens": 1, + "output_tokens": 1, + "total_tokens": 2, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + "error": None, + "incomplete_details": None, + "instructions": None, + "metadata": None, + "temperature": None, + "tool_choice": "auto", + "tools": [], + "top_p": None, + "max_output_tokens": None, + "previous_response_id": None, + "reasoning": None, + "truncation": None, + "user": None, + } + response_route: Final = respx_mock.post("https://api.perplexity.ai/v1/responses").respond(json=response_body) + model: Final = "perplexity/test-responses-bridge" + litellm.register_model( + { + model: { + "litellm_provider": "perplexity", + "mode": "responses", + "supports_reasoning": False, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + } + }, + persist_across_reloads=False, + ) + + litellm.completion( + model=model, + messages=[{"role": "user", "content": "hello"}], + reasoning_effort=reasoning_effort, + drop_params=True, + api_key="fake-key", + api_base="https://api.perplexity.ai", + ) + + request_body: Final = json.loads(response_route.calls[0].request.content) + assert request_body["reasoning"] == {"effort": "high"} + + +_FOUNDRY_RESPONSES_FUNCTION_CALL_BODY: Final = { + "id": "resp_foundry", + "object": "response", + "created_at": 1789852145, + "status": "completed", + "model": "gpt-6-astra", + "output": [ + { + "id": "fc_1", + "type": "function_call", + "status": "completed", + "arguments": '{"city":"Paris"}', + "call_id": "call_1", + "name": "get_weather", + } + ], + "parallel_tool_calls": True, + "usage": { + "input_tokens": 53, + "output_tokens": 18, + "total_tokens": 71, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + "error": None, + "incomplete_details": None, + "instructions": None, + "metadata": {}, + "temperature": 1.0, + "tool_choice": "auto", + "tools": [], + "top_p": 1.0, + "max_output_tokens": 200, + "previous_response_id": None, + "reasoning": {"effort": "medium", "summary": None}, + "truncation": "disabled", + "user": None, +} + + +def test_completion_bridges_azure_ai_foundry_gpt_5_4_plus_function_tools_to_responses( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + responses_route: Final = respx_mock.post(f"{_FOUNDRY_API_BASE}/openai/v1/responses").respond( + json=_FOUNDRY_RESPONSES_FUNCTION_CALL_BODY + ) + + response: Final = litellm.completion( + model="azure_ai/gpt-6-astra", + messages=[{"role": "user", "content": "What is the weather in Paris? Use the tool."}], + tools=[ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather for a city", + "parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}, + }, + } + ], + max_tokens=200, + api_base=_FOUNDRY_API_BASE, + api_key="fake-foundry-key", + ) + + assert [str(call.request.url) for call in respx_mock.calls] == [f"{_FOUNDRY_API_BASE}/openai/v1/responses"] + request: Final = responses_route.calls[0].request + request_body: Final = json.loads(request.content) + assert request_body["tools"][0]["type"] == "function" + assert request_body["tools"][0]["name"] == "get_weather" + assert request.headers["api-key"] == "fake-foundry-key" + assert response.choices[0].finish_reason == "tool_calls" + assert response.choices[0].message.tool_calls[0].function.name == "get_weather" + + +@pytest.mark.parametrize( + "model, model_info, expected_model_param, expected_base_model_param", + [ + ("gemini/gemini-3.1-pro", None, "gemini-3.1-pro", None), + ( + "gemini/gemini-3.1-pro", + {"base_model": "gemini-3.1-pro-preview"}, + "gemini-3.1-pro", + "gemini-3.1-pro-preview", + ), + ], +) +def test_completion_optional_params_base_model( + model: str, + model_info: dict | None, + expected_model_param: str, + expected_base_model_param: str | None, +): + """``model_info.base_model`` must reach ``get_optional_params`` as ``base_model`` + (an additive capability hint), without overwriting ``model`` with the label. + + Regression for #29618: overwriting ``model`` with a friendly ``base_model`` + label made Bedrock drop ``tools``/``tool_choice`` under ``drop_params``.""" + with patch("litellm.main.get_optional_params") as mock_get_optional_params: + mock_get_optional_params.return_value = MagicMock() + + import litellm + + kwargs = { + "model": model, + "messages": [{"role": "user", "content": "What is the capital of France?"}], + "api_key": "fake-key", + "mock_response": "Hey, how's it going?", + } + if model_info is not None: + kwargs["model_info"] = model_info + + litellm.completion(**kwargs) + + assert mock_get_optional_params.called is True + call_kwargs = mock_get_optional_params.call_args.kwargs + assert call_kwargs["model"] == expected_model_param + assert call_kwargs["base_model"] == expected_base_model_param + + +@patch("litellm.completion_extras.responses_api_bridge.completion") +def test_gpt_5_4_responses_bridge_merges_reasoning_summary_kwarg_without_tools( + mock_responses_completion, +): + """reasoningSummary without tools should route and merge into reasoning_effort dict.""" + mock_responses_completion.return_value = MagicMock() + + import litellm + + litellm.completion( + model="gpt-5.4", + messages=[{"role": "user", "content": "ok"}], + reasoning_effort="medium", + reasoningSummary="auto", + api_key="fake-key", + ) + + assert mock_responses_completion.called is True + optional_params = mock_responses_completion.call_args.kwargs["optional_params"] + assert optional_params["reasoning_effort"] == { + "effort": "medium", + "summary": "auto", + } + assert "reasoningSummary" not in optional_params + assert "reasoning_summary" not in optional_params + + +@patch("litellm.completion_extras.responses_api_bridge.completion") +def test_responses_bridge_preserves_reasoning_summary_without_effort( + mock_responses_completion, +): + """Reasoning summary should survive responses routing even without effort.""" + mock_responses_completion.return_value = MagicMock() + + import litellm + + with patch.object(litellm, "route_all_chat_openai_to_responses", True): + litellm.completion( + model="gpt-4o", + messages=[{"role": "user", "content": "ok"}], + reasoningSummary="auto", + api_key="fake-key", + ) + + assert mock_responses_completion.called is True + optional_params = mock_responses_completion.call_args.kwargs["optional_params"] + assert optional_params["reasoning_effort"] == {"summary": "auto"} + assert "reasoningSummary" not in optional_params + assert "reasoning_summary" not in optional_params + + +@patch("litellm.completion_extras.responses_api_bridge.completion") +def test_gpt_5_responses_bridge_tools_and_reasoning_summary( + mock_responses_completion, +): + """Bare gpt-5 with tools + reasoningSummary should bridge (OpenCode-style).""" + mock_responses_completion.return_value = MagicMock() + + import litellm + + litellm.completion( + model="gpt-5", + messages=[{"role": "user", "content": "ok"}], + tools=[ + { + "type": "function", + "function": { + "name": "apply_patch", + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + tool_choice="auto", + reasoning_effort="medium", + reasoningSummary="auto", + stream=True, + api_key="fake-key", + ) + + assert mock_responses_completion.called is True + optional_params = mock_responses_completion.call_args.kwargs["optional_params"] + assert optional_params.get("reasoning_effort") == { + "effort": "medium", + "summary": "auto", + } + + +def test_responses_api_bridge_check_handles_exception(): + """Test that responses_api_bridge_check handles exceptions and still processes responses/ models.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.side_effect = Exception("Model not found") + + model_info, model = responses_api_bridge_check( + model="responses/custom-model", custom_llm_provider="custom" + ) + + assert model == "custom-model" + assert model_info["mode"] == "responses" + + +def test_responses_api_bridge_check_global_flag_routes_openai(): + """When route_all_chat_openai_to_responses is True, any OpenAI model routes to responses.""" + from litellm.main import responses_api_bridge_check + + with patch.object(litellm, "route_all_chat_openai_to_responses", True): + model_info, model = responses_api_bridge_check( + model="gpt-4o", + custom_llm_provider="openai", + ) + + assert model == "gpt-4o" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_global_flag_does_not_affect_azure(): + """route_all_chat_openai_to_responses should not affect Azure models.""" + from litellm.main import responses_api_bridge_check + + with patch.object(litellm, "route_all_chat_openai_to_responses", True): + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 4096} + model_info, model = responses_api_bridge_check( + model="gpt-4o", + custom_llm_provider="azure", + ) + + assert model_info.get("mode") != "responses" + + +def test_responses_api_bridge_check_global_flag_default_false(): + """By default, route_all_chat_openai_to_responses is False and doesn't affect routing.""" + from litellm.main import responses_api_bridge_check + + with patch.object(litellm, "route_all_chat_openai_to_responses", False): + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 4096} + model_info, model = responses_api_bridge_check( + model="gpt-4o", + custom_llm_provider="openai", + ) + + assert model_info.get("mode") != "responses" + + +@pytest.mark.asyncio +async def test_async_mock_delay(): + """Use asyncio await for mock delay on acompletion""" + import time + + from litellm import acompletion + + start_time = time.time() + result = await acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "Hey, how's it going?"}], + mock_delay=0.01, + mock_response="Hello world", + ) + end_time = time.time() + delay = end_time - start_time + assert delay >= 0.01 + + +def test_stream_chunk_builder_keeps_tool_calls_carried_only_by_a_later_choice_of_a_multi_choice_chunk(): + from litellm import stream_chunk_builder + from litellm.types.utils import ( + ChatCompletionDeltaToolCall, + Delta, + Function, + ModelResponseStream, + StreamingChoices, + ) + + def chunk(choices: list[StreamingChoices]) -> ModelResponseStream: + return ModelResponseStream( + id="chatcmpl-multi-choice", + created=1751934860, + model="gpt-4.1-mini", + object="chat.completion.chunk", + choices=choices, + ) + + chunks = [ + chunk( + [ + StreamingChoices(index=0, delta=Delta(role="assistant", content="hello")), + StreamingChoices( + index=1, + delta=Delta( + role="assistant", + tool_calls=[ + ChatCompletionDeltaToolCall( + id="call_1", + index=0, + type="function", + function=Function(name="lookup_fruit", arguments='{"fruit":'), + ) + ], + ), + ), + ] + ), + chunk( + [ + StreamingChoices(index=0, delta=Delta(content=" world"), finish_reason="stop"), + StreamingChoices( + index=1, + delta=Delta( + tool_calls=[ChatCompletionDeltaToolCall(index=0, function=Function(arguments='"kiwi"}'))] + ), + finish_reason="tool_calls", + ), + ] + ), + ] + + response = stream_chunk_builder(chunks=chunks) + + tool_calls = response.choices[0].message.tool_calls + assert tool_calls is not None + assert [(call.id, call.function.name, call.function.arguments) for call in tool_calls] == [ + ("call_1", "lookup_fruit", '{"fruit":"kiwi"}') + ] + + +def test_stream_chunk_builder_thinking_blocks(): + from litellm import stream_chunk_builder + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + chunks = [ + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + reasoning_content="I need to summar", + thinking_blocks=[ + { + "type": "thinking", + "thinking": "I need to summar", + "signature": None, + } + ], + provider_specific_fields={ + "thinking_blocks": [ + { + "type": "thinking", + "thinking": "I need to summar", + "signature": None, + } + ] + }, + content="", + role="assistant", + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + reasoning_content="ize the previous agent's thinking process into a", + thinking_blocks=[ + { + "type": "thinking", + "thinking": "ize the previous agent's thinking process into a", + "signature": None, + } + ], + provider_specific_fields={ + "thinking_blocks": [ + { + "type": "thinking", + "thinking": "ize the previous agent's thinking process into a", + "signature": None, + } + ] + }, + content="", + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + reasoning_content=" short description. Based on the input data provide", + thinking_blocks=[ + { + "type": "thinking", + "thinking": " short description. Based on the input data provide", + "signature": None, + } + ], + provider_specific_fields={ + "thinking_blocks": [ + { + "type": "thinking", + "thinking": " short description. Based on the input data provide", + "signature": None, + } + ] + }, + content="", + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + reasoning_content="d, it seems the agent was planning to refine their search", + thinking_blocks=[ + { + "type": "thinking", + "thinking": "d, it seems the agent was planning to refine their search", + "signature": None, + } + ], + provider_specific_fields={ + "thinking_blocks": [ + { + "type": "thinking", + "thinking": "d, it seems the agent was planning to refine their search", + "signature": None, + } + ] + }, + content="", + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + reasoning_content=" to focus more on technical aspects of home automation and home", + thinking_blocks=[ + { + "type": "thinking", + "thinking": " to focus more on technical aspects of home automation and home", + "signature": None, + } + ], + provider_specific_fields={ + "thinking_blocks": [ + { + "type": "thinking", + "thinking": " to focus more on technical aspects of home automation and home", + "signature": None, + } + ] + }, + content="", + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + reasoning_content=" energy system management.\n\nI'll create a brief", + thinking_blocks=[ + { + "type": "thinking", + "thinking": " energy system management.\n\nI'll create a brief", + "signature": None, + } + ], + provider_specific_fields={ + "thinking_blocks": [ + { + "type": "thinking", + "thinking": " energy system management.\n\nI'll create a brief", + "signature": None, + } + ] + }, + content="", + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + reasoning_content=" summary of what the agent was doing.", + thinking_blocks=[ + { + "type": "thinking", + "thinking": " summary of what the agent was doing.", + "signature": None, + } + ], + provider_specific_fields={ + "thinking_blocks": [ + { + "type": "thinking", + "thinking": " summary of what the agent was doing.", + "signature": None, + } + ] + }, + content="", + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + reasoning_content="", + thinking_blocks=[ + { + "type": "thinking", + "thinking": "", + "signature": "ErUBCkYIBRgCIkAKBSMkB2+MBF643wiWxlERsGXVdlhbPx9lnTIbygzjFIeZ5uhTV+HNWDon9vQV4hmXvAKwQfwS8vkNFB366l05Egzt2U18IpRrZRyQn1UaDDdYvKHYP8Ps1IbWjSIw8eSYOU9gtqNcwR6D0wY7iOPx2GliDEatLI5rSs96CByoTIoADL2M5bX8KP0jEpbHKh0ccYryigdH/3J8EiFt/BmGUceVASP5l9r22dFWiBgC", + } + ], + provider_specific_fields={ + "thinking_blocks": [ + { + "type": "thinking", + "thinking": "", + "signature": "ErUBCkYIBRgCIkAKBSMkB2+MBF643wiWxlERsGXVdlhbPx9lnTIbygzjFIeZ5uhTV+HNWDon9vQV4hmXvAKwQfwS8vkNFB366l05Egzt2U18IpRrZRyQn1UaDDdYvKHYP8Ps1IbWjSIw8eSYOU9gtqNcwR6D0wY7iOPx2GliDEatLI5rSs96CByoTIoADL2M5bX8KP0jEpbHKh0ccYryigdH/3J8EiFt/BmGUceVASP5l9r22dFWiBgC", + } + ] + }, + content="", + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=1, + delta=Delta( + provider_specific_fields=None, + content='{"a', + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=1, + delta=Delta( + provider_specific_fields=None, + content='gent_doing"', + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=1, + delta=Delta( + provider_specific_fields=None, + content=': "Re', + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=1, + delta=Delta( + provider_specific_fields=None, + content="searching", + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=1, + delta=Delta( + provider_specific_fields=None, + content=" technic", + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=1, + delta=Delta( + provider_specific_fields=None, + content="al aspect", + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=1, + delta=Delta( + provider_specific_fields=None, + content="s of home au", + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason=None, + index=1, + delta=Delta( + provider_specific_fields=None, + content='tomation"}', + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + citations=None, + ), + ModelResponseStream( + id="chatcmpl-e8febeb7-cf7d-4947-9417-59ae5e6989f9", + created=1751934860, + model="claude-3-7-sonnet-latest", + object="chat.completion.chunk", + system_fingerprint=None, + choices=[ + StreamingChoices( + finish_reason="tool_calls", + index=0, + delta=Delta( + provider_specific_fields=None, + content=None, + role=None, + function_call=None, + tool_calls=None, + audio=None, + ), + logprobs=None, + ) + ], + provider_specific_fields=None, + ), + ] + + response = stream_chunk_builder(chunks=chunks) + print(response) + + assert response is not None + assert response.choices[0].message.content is not None + assert response.choices[0].message.thinking_blocks is not None + + +from litellm.llms.openai.openai import OpenAIChatCompletion + + +def throw_retryable_error(*_, **__): + raise RuntimeError("BOOM") + + +@pytest.mark.asyncio +async def test_retrying() -> None: + litellm.num_retries = 10 + with ( + patch.object( + OpenAIChatCompletion, + "make_openai_chat_completion_request", + side_effect=throw_retryable_error, + ) as mock_request, + pytest.raises(litellm.InternalServerError, match="LiteLLM Retried: 10 times"), + ): + await litellm.acompletion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Hello"}], + ) + + +def test_anthropic_disable_url_suffix_env_var(): + """Test that LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX prevents /v1/messages suffix.""" + import os + from unittest.mock import MagicMock, patch + + from litellm import completion + + # Test with environment variable disabled (default behavior) + with patch.dict(os.environ, {"ANTHROPIC_API_BASE": "https://api.example.com"}): + actual_api_base = None + + with patch("litellm.main.anthropic_chat_completions") as mock_anthropic: + + def capture_completion(**kwargs): + nonlocal actual_api_base + actual_api_base = kwargs.get("api_base") + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + return mock_response + + mock_anthropic.completion = capture_completion + + # This should append /v1/messages + completion( + model="anthropic/claude-3-sonnet", + messages=[{"role": "user", "content": "test"}], + api_key="test-key", + ) + + # Verify the api_base has /v1/messages appended + assert actual_api_base.endswith("/v1/messages") + assert actual_api_base == "https://api.example.com/v1/messages" + + # Test with environment variable enabled + with patch.dict( + os.environ, + { + "ANTHROPIC_API_BASE": "https://api.example.com/custom/path", + "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX": "true", + }, + ): + actual_api_base = None + + with patch("litellm.main.anthropic_chat_completions") as mock_anthropic: + + def capture_completion(**kwargs): + nonlocal actual_api_base + actual_api_base = kwargs.get("api_base") + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + return mock_response + + mock_anthropic.completion = capture_completion + + # This should NOT append /v1/messages + completion( + model="anthropic/claude-3-sonnet", + messages=[{"role": "user", "content": "test"}], + api_key="test-key", + ) + + # Verify the api_base does not have /v1/messages appended + assert actual_api_base == "https://api.example.com/custom/path" + assert not actual_api_base.endswith("/v1/messages") + + +def test_anthropic_text_disable_url_suffix_env_var(): + """Test that LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX prevents /v1/complete suffix for anthropic_text.""" + import os + from unittest.mock import MagicMock, patch + + from litellm import completion + + # Test with environment variable disabled (default behavior) + with patch.dict(os.environ, {"ANTHROPIC_API_BASE": "https://api.example.com"}): + actual_api_base = None + + with patch("litellm.main.base_llm_http_handler") as mock_handler: + + def capture_completion(**kwargs): + nonlocal actual_api_base + actual_api_base = kwargs.get("api_base") + return MagicMock() + + mock_handler.completion = capture_completion + + # This should append /v1/complete + completion( + model="anthropic_text/claude-instant-1", + messages=[{"role": "user", "content": "test"}], + api_key="test-key", + ) + + # Verify the api_base has /v1/complete appended + assert actual_api_base.endswith("/v1/complete") + assert actual_api_base == "https://api.example.com/v1/complete" + + # Test with environment variable enabled + with patch.dict( + os.environ, + { + "ANTHROPIC_API_BASE": "https://api.example.com/custom/complete", + "LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX": "true", + }, + ): + actual_api_base = None + + with patch("litellm.main.base_llm_http_handler") as mock_handler: + + def capture_completion(**kwargs): + nonlocal actual_api_base + actual_api_base = kwargs.get("api_base") + return MagicMock() + + mock_handler.completion = capture_completion + + # This should NOT append /v1/complete + completion( + model="anthropic_text/claude-instant-1", + messages=[{"role": "user", "content": "test"}], + api_key="test-key", + ) + + # Verify the api_base does not have /v1/complete appended + assert actual_api_base == "https://api.example.com/custom/complete" + assert not actual_api_base.endswith("/v1/complete") + + +def test_image_edit_merges_headers_and_extra_headers(): + from litellm.images.main import base_llm_http_handler + + combined_headers = { + "x-test-header-one": "value-1", + "x-test-header-two": "value-2", + } + + mock_image_edit_config = MagicMock() + mock_image_edit_config.get_supported_openai_params.return_value = set() + mock_image_edit_config.map_openai_params.side_effect = lambda **kwargs: dict( + kwargs["image_edit_optional_params"] + ) + + with ( + patch( + "litellm.images.main.ProviderConfigManager.get_provider_image_edit_config", + return_value=mock_image_edit_config, + ) as mock_config, + patch.object( + base_llm_http_handler, + "image_edit_handler", + return_value="ok", + ) as mock_handler, + ): + response = litellm.image_edit( + image=MagicMock(name="image"), + prompt="test", + model="azure/gpt-image-1", + headers={"x-test-header-one": "value-1"}, + extra_headers={ + "x-test-header-two": "value-2", + }, + ) + + assert response == "ok" + mock_config.assert_called_once() + + handler_kwargs = mock_handler.call_args.kwargs + assert handler_kwargs["extra_headers"] == combined_headers + assert "extra_headers" not in handler_kwargs["image_edit_optional_request_params"] + + +@pytest.mark.parametrize("metadata_key", ("metadata", "litellm_metadata")) +@pytest.mark.parametrize("input_tokens", (51234, 0)) +def test_mock_completion_usage_reports_admission_input_tokens(metadata_key: str, input_tokens: int): + response = litellm.completion( + model="anthropic/claude-sonnet-5", + messages=[{"role": "user", "content": "hello"}], + mock_response="ok", + api_key="mock", + **{metadata_key: {"user_api_key_budget_reservation": {"reserved_cost": 1.0, "input_tokens": input_tokens}}}, + ) + + assert response.usage.prompt_tokens == input_tokens + assert response.usage.total_tokens == input_tokens + response.usage.completion_tokens + + +def test_mock_completion_usage_falls_back_to_default_without_admission_count(): + response = litellm.completion( + model="anthropic/claude-sonnet-5", + messages=[{"role": "user", "content": "hello"}], + mock_response="ok", + api_key="mock", + metadata={"user_api_key_budget_reservation": {"reserved_cost": 1.0}}, + ) + + assert response.usage.prompt_tokens == litellm_main.DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT + + +_AZURE_AI_CUSTOM_PRICED_DEPLOYMENT: Final = { + "model_name": "azure-ai-custom-priced", + "litellm_params": { + "model": "azure_ai/gpt-5.6", + "api_key": "mock", + "api_base": "https://example.services.ai.azure.com", + "mock_response": "ok", + "input_cost_per_token": 3e-6, + "output_cost_per_token": 7e-6, + "cache_read_input_token_cost": 1e-7, + "cache_creation_input_token_cost": 5e-7, + }, + "model_info": {"id": "azure-ai-custom-priced-deployment-id"}, +} + + +def _expected_custom_price(response: litellm.ModelResponse) -> float: + params: Final = _AZURE_AI_CUSTOM_PRICED_DEPLOYMENT["litellm_params"] + return ( + response.usage.prompt_tokens * params["input_cost_per_token"] + + response.usage.completion_tokens * params["output_cost_per_token"] + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("use_async", (False, True)) +async def test_mock_completion_prices_azure_ai_router_deployment_with_custom_pricing(use_async: bool): + router: Final = litellm.Router(model_list=[_AZURE_AI_CUSTOM_PRICED_DEPLOYMENT]) + messages: Final = [{"role": "user", "content": "hello"}] + + response: Final = ( + await router.acompletion(model="azure-ai-custom-priced", messages=messages) + if use_async + else router.completion(model="azure-ai-custom-priced", messages=messages) + ) + + assert response._hidden_params["response_cost"] == pytest.approx(_expected_custom_price(response)) + assert response._hidden_params["custom_llm_provider"] == "azure_ai" + + +@pytest.mark.parametrize( + ("model", "expected_provider"), + (("anthropic/claude-sonnet-5", "anthropic"), ("no-such-provider-model", None)), +) +def test_mock_completion_infers_provider_when_called_directly_without_one(model: str, expected_provider: str | None): + response: Final = litellm.mock_completion( + model=model, + messages=[{"role": "user", "content": "hello"}], + mock_response="ok", + ) + + assert response.choices[0].message.content == "ok" + assert response._hidden_params.get("custom_llm_provider") == expected_provider + + +_ADMISSION_INPUT_TOKENS: Final = 51234 + + +def _admission_metadata(input_tokens: int) -> dict[str, object]: # mutable-ok: logging writes into metadata + return {"user_api_key_budget_reservation": {"reserved_cost": 1.0, "input_tokens": input_tokens}} + + +_ADMISSION_METADATA: Final = _admission_metadata(_ADMISSION_INPUT_TOKENS) +_MOCK_STREAM_MESSAGES: Final = [{"role": "user", "content": "hello " * 200}] +_STREAM_CHUNK_BUILDER_TOKEN_COUNTER: Final = "litellm.litellm_core_utils.streaming_chunk_builder_utils.token_counter" + + +def _prompt_token_counter_calls(token_counter: MagicMock) -> list[object]: + return [call for call in token_counter.call_args_list if call.kwargs.get("messages") is not None] + + +def _client_usage_chunks(chunks: list[ModelResponseStream]) -> list[Usage]: + return [chunk.usage for chunk in chunks if getattr(chunk, "usage", None) is not None] + + +@pytest.mark.parametrize("n", (None, 2)) +def test_mock_completion_stream_usage_reports_admission_input_tokens_without_tokenizer_fallback(n: int | None): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks: Final = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + n=n, + stream_options={"include_usage": True}, + metadata=_ADMISSION_METADATA, + ) + ) + + usage_chunks: Final = _client_usage_chunks(chunks) + assert len(usage_chunks) == 1 + assert usage_chunks[0].prompt_tokens == _ADMISSION_INPUT_TOKENS + assert usage_chunks[0].completion_tokens == litellm_main.DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT + assert usage_chunks[0].total_tokens == _ADMISSION_INPUT_TOKENS + usage_chunks[0].completion_tokens + assert _prompt_token_counter_calls(token_counter) == [] + assert all(chunk.choices for chunk in chunks[:-1]) + assert {chunk.id for chunk in chunks} == {chunks[0].id} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("n", (None, 2)) +async def test_mock_acompletion_stream_usage_reports_admission_input_tokens_without_tokenizer_fallback( + n: int | None, +): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + response: Final = await litellm.acompletion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + n=n, + stream_options={"include_usage": True}, + litellm_metadata=_ADMISSION_METADATA, + ) + chunks: Final = [chunk async for chunk in response] + + usage_chunks: Final = _client_usage_chunks(chunks) + assert len(usage_chunks) == 1 + assert usage_chunks[0].prompt_tokens == _ADMISSION_INPUT_TOKENS + assert usage_chunks[0].total_tokens == _ADMISSION_INPUT_TOKENS + usage_chunks[0].completion_tokens + assert _prompt_token_counter_calls(token_counter) == [] + assert all(chunk.choices for chunk in chunks[:-1]) + assert {chunk.id for chunk in chunks} == {chunks[0].id} + + +def test_mock_completion_stream_without_include_usage_hides_usage_chunk_but_logs_admission_count(): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks: Final = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + metadata=_ADMISSION_METADATA, + ) + ) + + assert _client_usage_chunks(chunks) == [] + assert all(len(chunk.choices) == 1 for chunk in chunks) + assert chunks[-1]._hidden_params["usage"].prompt_tokens == _ADMISSION_INPUT_TOKENS + assert _prompt_token_counter_calls(token_counter) == [] + + +def test_mock_completion_stream_with_empty_stream_options_completes_and_logs_admission_count(): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks: Final = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={}, + metadata=_ADMISSION_METADATA, + ) + ) + + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" + assert _client_usage_chunks(chunks) == [] + assert _prompt_token_counter_calls(token_counter) == [] + + +@pytest.mark.asyncio +async def test_mock_acompletion_stream_with_empty_stream_options_completes_and_logs_admission_count(): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + response: Final = await litellm.acompletion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={}, + litellm_metadata=_ADMISSION_METADATA, + ) + chunks: Final = [chunk async for chunk in response] + + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" + assert _client_usage_chunks(chunks) == [] + assert _prompt_token_counter_calls(token_counter) == [] + + +def test_mock_completion_stream_without_admission_count_falls_back_to_tokenizer(): + expected_prompt_tokens: Final = litellm.token_counter(model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES) + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks: Final = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + metadata={"user_api_key_budget_reservation": {"reserved_cost": 1.0}}, + ) + ) + + usage_chunks: Final = _client_usage_chunks(chunks) + assert len(usage_chunks) == 1 + assert usage_chunks[0].prompt_tokens == expected_prompt_tokens + assert usage_chunks[0].total_tokens == expected_prompt_tokens + usage_chunks[0].completion_tokens + assert len(_prompt_token_counter_calls(token_counter)) >= 1 + + +@pytest.mark.asyncio +async def test_mock_acompletion_stream_without_admission_count_falls_back_to_tokenizer(): + expected_prompt_tokens: Final = litellm.token_counter(model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES) + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + response: Final = await litellm.acompletion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + ) + chunks: Final = [chunk async for chunk in response] + + usage_chunks: Final = _client_usage_chunks(chunks) + assert len(usage_chunks) == 1 + assert usage_chunks[0].prompt_tokens == expected_prompt_tokens + assert len(_prompt_token_counter_calls(token_counter)) >= 1 + + +def _usage_triple(usage: Usage) -> tuple[int, int, int]: + return (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) + + +@pytest.mark.parametrize("input_tokens", (_ADMISSION_INPUT_TOKENS, 0)) +def test_mock_completion_stream_and_non_stream_report_the_same_admission_usage(input_tokens: int): + metadata: Final = _admission_metadata(input_tokens) + non_stream: Final = litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + metadata=metadata, + ) + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks: Final = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + metadata=metadata, + ) + ) + + assert _usage_triple(non_stream.usage) == _usage_triple(_client_usage_chunks(chunks)[0]) + assert non_stream.usage.prompt_tokens == input_tokens + assert _prompt_token_counter_calls(token_counter) == [] + + +@pytest.mark.asyncio +async def test_mock_acompletion_stream_reports_zero_admission_input_tokens_without_tokenizer_fallback(): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + response: Final = await litellm.acompletion( + model="openai/gpt-5.4-mini", + messages=[{"role": "user", "content": ""}], + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + litellm_metadata=_admission_metadata(0), + ) + chunks: Final = [chunk async for chunk in response] + + usage_chunks: Final = _client_usage_chunks(chunks) + assert len(usage_chunks) == 1 + assert _usage_triple(usage_chunks[0]) == (0, usage_chunks[0].completion_tokens, usage_chunks[0].completion_tokens) + assert _prompt_token_counter_calls(token_counter) == [] + + +def test_mock_text_completion_stream_and_non_stream_report_the_same_zero_admission_usage(): + metadata: Final = _admission_metadata(0) + non_stream: Final = litellm.text_completion( + model="openai/gpt-5.4-mini", prompt="", mock_response="ok", api_key="mock", metadata=metadata + ) + chunks: Final = list( + litellm.text_completion( + model="openai/gpt-5.4-mini", + prompt="", + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + metadata=metadata, + ) + ) + + stream_usages: Final = tuple(chunk.usage for chunk in chunks if getattr(chunk, "usage", None) is not None) + assert len(stream_usages) == 1 + assert _usage_triple(non_stream.usage) == _usage_triple(stream_usages[0]) + assert non_stream.usage.prompt_tokens == 0 + + +def test_mock_completion_stream_with_model_response(): + """Test that mock_completion correctly handles stream=True with a ModelResponse as mock_response.""" + from litellm import completion + from litellm.types.utils import Choices, Message, ModelResponse, Usage + + # Create a ModelResponse object + mock_model_response = ModelResponse( + id="chatcmpl-test-123", + created=1234567890, + model="gpt-4o-mini", + object="chat.completion", + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message( + content="This is a test response", + role="assistant", + ), + ) + ], + usage=Usage( + prompt_tokens=10, + completion_tokens=20, + total_tokens=30, + ), + ) + + # Call completion with stream=True and mock_response as ModelResponse + response = completion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Hello"}], + stream=True, + mock_response=mock_model_response, + ) + + # Verify that the response is a stream + assert response is not None + + # Collect all chunks from the stream + chunks = [] + for chunk in response: + chunks.append(chunk) + print(f"Chunk: {chunk}") + + # Verify we got chunks + assert len(chunks) > 0 + + # Verify the content is streamed correctly + accumulated_content = "" + for chunk in chunks: + if ( + hasattr(chunk.choices[0].delta, "content") + and chunk.choices[0].delta.content + ): + accumulated_content += chunk.choices[0].delta.content + + assert "This is a test response" in accumulated_content or len(chunks) > 0 + + +@pytest.mark.asyncio +async def test_async_mock_completion_stream_with_model_response(): + """Test that async mock_completion correctly handles stream=True with a ModelResponse as mock_response.""" + from litellm import acompletion + from litellm.types.utils import Choices, Message, ModelResponse, Usage + + # Create a ModelResponse object + mock_model_response = ModelResponse( + id="chatcmpl-test-456", + created=1234567890, + model="gpt-4o-mini", + object="chat.completion", + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message( + content="This is an async test response", + role="assistant", + ), + ) + ], + usage=Usage( + prompt_tokens=15, + completion_tokens=25, + total_tokens=40, + ), + ) + + # Call acompletion with stream=True and mock_response as ModelResponse + response = await acompletion( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "Hello async"}], + stream=True, + mock_response=mock_model_response, + ) + + # Verify that the response is a stream + assert response is not None + + # Collect all chunks from the stream + chunks = [] + async for chunk in response: + chunks.append(chunk) + print(f"Async Chunk: {chunk}") + + # Verify we got chunks + assert len(chunks) > 0 + + # Verify the content is streamed correctly + accumulated_content = "" + for chunk in chunks: + if ( + hasattr(chunk.choices[0].delta, "content") + and chunk.choices[0].delta.content + ): + accumulated_content += chunk.choices[0].delta.content + + assert "This is an async test response" in accumulated_content or len(chunks) > 0 + + +class TestCallTypesOCR: + """Test that OCR call types are properly defined in CallTypes enum. + + Fixes https://github.com/BerriAI/litellm/issues/17381 + """ + + def test_ocr_call_type_exists(self): + """Test that CallTypes.ocr exists and has correct value.""" + from litellm.types.utils import CallTypes + + assert hasattr(CallTypes, "ocr") + assert CallTypes.ocr.value == "ocr" + + def test_aocr_call_type_exists(self): + """Test that CallTypes.aocr exists and has correct value.""" + from litellm.types.utils import CallTypes + + assert hasattr(CallTypes, "aocr") + assert CallTypes.aocr.value == "aocr" + + def test_ocr_call_type_from_string(self): + """Test that CallTypes can be constructed from 'ocr' string.""" + from litellm.types.utils import CallTypes + + call_type = CallTypes("ocr") + assert call_type == CallTypes.ocr + + def test_aocr_call_type_from_string(self): + """Test that CallTypes can be constructed from 'aocr' string. + + This is the actual use case that was failing - the OCR endpoint + uses route_type='aocr' and guardrails try to instantiate + CallTypes('aocr'). + """ + from litellm.types.utils import CallTypes + + call_type = CallTypes("aocr") + assert call_type == CallTypes.aocr + + +def test_stream_chunk_builder_text_completion_combines_text_and_usage(): + from litellm.main import stream_chunk_builder_text_completion + from litellm.types.utils import TextCompletionResponse + + chunks = [ + TextCompletionResponse( + id="cmpl-1", + object="text_completion", + created=1, + model="gpt-3.5-turbo-instruct", + choices=[{"text": "Hello", "index": 0, "logprobs": None, "finish_reason": None}], + ), + TextCompletionResponse( + id="cmpl-1", + object="text_completion", + created=1, + model="gpt-3.5-turbo-instruct", + choices=[{"text": " world", "index": 0, "logprobs": None, "finish_reason": "stop"}], + ), + ] + + response = stream_chunk_builder_text_completion( + chunks=chunks, messages=[{"role": "user", "content": "say hello"}] + ) + + assert response.choices[0].text == "Hello world" + assert response.choices[0].finish_reason == "stop" + assert response.usage.prompt_tokens > 0 + assert response.usage.completion_tokens > 0 + assert response.usage.total_tokens == response.usage.prompt_tokens + response.usage.completion_tokens + + +def test_completion_forwards_store_and_prompt_cache_key_to_openai(): + """ + Regression test for https://github.com/BerriAI/litellm/issues/33184 + + store and prompt_cache_key are documented OpenAI chat completion params that + were accepted as supported but silently dropped before the provider request + was built, because they were not named parameters of completion() and + get_optional_params() the way safety_identifier is. + """ + from openai import OpenAI + + client = OpenAI(api_key="fake-api-key") + + with patch.object(client.chat.completions.with_raw_response, "create") as mock_client: + try: + litellm.completion( + model="openai/gpt-4o", + messages=[{"role": "user", "content": "Hello"}], + store=False, + prompt_cache_key="test-cache-key", + client=client, + ) + except Exception as e: + print(e) + + mock_client.assert_called_once() + request_body = mock_client.call_args.kwargs + assert request_body["store"] is False + assert request_body["prompt_cache_key"] == "test-cache-key" + + +@pytest.mark.asyncio +async def test_acompletion_forwards_store_and_prompt_cache_key_to_openai(): + """ + Async variant of the store/prompt_cache_key forwarding regression test for + https://github.com/BerriAI/litellm/issues/33184 + """ + from openai import AsyncOpenAI + + client = AsyncOpenAI(api_key="fake-api-key") + + with patch.object(client.chat.completions.with_raw_response, "create") as mock_client: + try: + await litellm.acompletion( + model="openai/gpt-4o", + messages=[{"role": "user", "content": "Hello"}], + store=False, + prompt_cache_key="test-cache-key", + client=client, + ) + except Exception as e: + print(e) + + mock_client.assert_called_once() + request_body = mock_client.call_args.kwargs + assert request_body["store"] is False + assert request_body["prompt_cache_key"] == "test-cache-key" + + +def test_completion_omits_store_and_prompt_cache_key_when_not_passed(): + """ + When store and prompt_cache_key are not passed, they must not appear in the + outbound request body (guards against always forwarding None defaults). + """ + from openai import OpenAI + + client = OpenAI(api_key="fake-api-key") + + with patch.object(client.chat.completions.with_raw_response, "create") as mock_client: + try: + litellm.completion( + model="openai/gpt-4o", + messages=[{"role": "user", "content": "Hello"}], + client=client, + ) + except Exception as e: + print(e) + + mock_client.assert_called_once() + request_body = mock_client.call_args.kwargs + assert "store" not in request_body + assert "prompt_cache_key" not in request_body + + +def test_completion_forwards_store_and_prompt_cache_key_to_mcp_gateway(): + """ + Regression test for the MCP gateway early-return in completion(): store and + prompt_cache_key are named params, so they no longer travel via **kwargs and + must be forwarded explicitly like safety_identifier and service_tier. + """ + with patch.object( + import_module("litellm.responses.mcp.chat_completions_handler"), "acompletion_with_mcp" + ) as mock_mcp: + result = litellm.completion( + model="openai/gpt-4o", + messages=[{"role": "user", "content": "Hello"}], + tools=[{"type": "mcp", "server_url": "litellm_proxy"}], + store=False, + prompt_cache_key="test-cache-key", + ) + + result.close() + mock_mcp.assert_called_once() + call_kwargs = mock_mcp.call_args.kwargs + assert call_kwargs["store"] is False + assert call_kwargs["prompt_cache_key"] == "test-cache-key" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "aws_credential_kwargs", + [ + { + "aws_session_name": "litellm-gcp", + "aws_role_name": "arn:aws:iam::123456789012:role/litellm-bedrock-role", + "aws_web_identity_token": "oidc/google/108963886734710037768", + }, + { + "aws_access_key_id": "AKIASTATICKEYFORTEST", + "aws_secret_access_key": "static-secret-key", + "aws_session_token": "static-session-token", + }, + ], + ids=["web_identity", "static_keys"], +) +async def test_acompletion_forwards_aws_credentials_through_responses_bridge( + respx_mock: respx.MockRouter, monkeypatch, aws_credential_kwargs: dict +): + from botocore.credentials import Credentials + + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + original_disable_aiohttp = litellm.disable_aiohttp_transport + try: + litellm.disable_aiohttp_transport = True + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + litellm.in_memory_llm_clients_cache.flush_cache() + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + + get_credentials_mock = MagicMock(return_value=Credentials("fake-key", "fake-secret")) + monkeypatch.setattr(BaseAWSLLM, "get_credentials", get_credentials_mock) + + respx_mock.post("https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses").respond( + json={ + "id": "resp_123", + "object": "response", + "created_at": 1760144904, + "status": "completed", + "model": "openai.gpt-5.4", + "output": [ + { + "type": "message", + "id": "msg_1", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "ok", "annotations": []}], + } + ], + } + ) + + response = await litellm.acompletion( + model="bedrock_mantle/openai.gpt-5.4", + messages=[{"role": "user", "content": "hi"}], + api_base="https://bedrock-mantle.us-east-2.api.aws/v1", + aws_region_name="us-east-2", + num_retries=0, + **aws_credential_kwargs, + ) + + assert response.choices[0].message.content == "ok" + credential_kwargs = get_credentials_mock.call_args.kwargs + assert credential_kwargs["aws_region_name"] == "us-east-2" + for key, value in aws_credential_kwargs.items(): + assert credential_kwargs[key] == value + authorization = respx_mock.calls.last.request.headers["Authorization"] + assert authorization.startswith("AWS4-HMAC-SHA256") + assert "fake-key" in authorization + finally: + litellm.disable_aiohttp_transport = original_disable_aiohttp + litellm.in_memory_llm_clients_cache.flush_cache() + + +_GEMINI_RESPONSE_BODY = { + "candidates": [{"content": {"parts": [{"text": "hello"}], "role": "model"}, "finishReason": "STOP"}], + "usageMetadata": {"promptTokenCount": 2, "candidatesTokenCount": 1, "totalTokenCount": 3}, +} + + +def _gemini_client_returning_a_reply(): + """An injected HTTP client whose post() answers like generativelanguage does.""" + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + client = HTTPHandler() + request = httpx.Request("POST", "https://generativelanguage.googleapis.com/") + post = MagicMock(return_value=httpx.Response(200, json=_GEMINI_RESPONSE_BODY, request=request)) + return client, post + + +@pytest.fixture +def restore_model_registry(): + """litellm.model_cost and the provider name sets are module-global. + + register_model merges into the existing entry in place, hence the deep copy. + """ + model_cost = copy.deepcopy(litellm.model_cost) + openai_models = set(litellm.open_ai_chat_completion_models) + yield + litellm.model_cost.clear() + litellm.model_cost.update(model_cost) + litellm.open_ai_chat_completion_models.clear() + litellm.open_ai_chat_completion_models.update(openai_models) + + +def test_openai_model_name_does_not_outrank_explicit_provider(): + """`gemini/gpt-4o` goes to Google, not to litellm's OpenAI handler. + + completion() checks `model in litellm.open_ai_chat_completion_models` ahead of + the gemini branch, so the call used to reach the OpenAI handler carrying + VertexGeminiConfig, whose transform_request raises NotImplementedError. + """ + assert "gpt-4o" in litellm.open_ai_chat_completion_models + client, post = _gemini_client_returning_a_reply() + + with patch.object(client, "post", new=post): + response = litellm.completion( + model="gemini/gpt-4o", + messages=[{"role": "user", "content": "hello"}], + api_key="test-api-key", + client=client, + ) + + assert "generativelanguage.googleapis.com" in post.call_args.kwargs["url"] + assert "models/gpt-4o" in post.call_args.kwargs["url"] + assert response.choices[0].message.content == "hello" + + +def test_mislabelled_pricing_entry_does_not_reroute_provider(restore_model_registry): + """register_model is the other way into the same failure. + + An entry claiming litellm_provider "openai" adds its name to + open_ai_chat_completion_models, so one mislabelled price reroutes every later + call to that model in the process. + """ + litellm.register_model( + { + "gemini-2.5-pro": { + "litellm_provider": "openai", + "mode": "chat", + "input_cost_per_token": 1e-06, + "output_cost_per_token": 4e-06, + } + } + ) + assert "gemini-2.5-pro" in litellm.open_ai_chat_completion_models + client, post = _gemini_client_returning_a_reply() + + with patch.object(client, "post", new=post): + response = litellm.completion( + model="gemini/gemini-2.5-pro", + messages=[{"role": "user", "content": "hello"}], + api_key="test-api-key", + client=client, + ) + + assert "generativelanguage.googleapis.com" in post.call_args.kwargs["url"] + assert response.choices[0].message.content == "hello" + + +def test_openai_model_without_a_provider_still_routes_to_openai(): + from openai import OpenAI + + client = OpenAI(api_key="fake-key") + raw_response = client.chat.completions.with_raw_response + with patch.object(raw_response, "create") as mock_create, contextlib.suppress(Exception): + litellm.completion( + model="gpt-4o", + messages=[{"role": "user", "content": "hello"}], + client=client, + ) + + mock_create.assert_called() + + +def _openai_chat_create_kwargs(client, **completion_kwargs): + with patch.object(client.chat.completions.with_raw_response, "create") as mock_client: + with contextlib.suppress(Exception): + litellm.completion( + messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], + cache_control_injection_points=[{"location": "message", "role": "system"}], + client=client, + **completion_kwargs, + ) + + mock_client.assert_called_once() + return mock_client.call_args.kwargs + + +@pytest.fixture +def _no_openai_api_base_override(monkeypatch): + monkeypatch.delenv("OPENAI_BASE_URL", raising=False) + monkeypatch.delenv("OPENAI_API_BASE", raising=False) + monkeypatch.setattr(litellm, "api_base", None) + + +@pytest.mark.usefixtures("_no_openai_api_base_override") +def test_completion_custom_api_base_sends_no_prompt_cache_breakpoint_for_gpt_5_6(): + from openai import OpenAI + + client = OpenAI(api_key="fake-api-key", base_url="http://127.0.0.1:9/v1") + request_body = _openai_chat_create_kwargs(client, model="gpt-5.6", api_base="http://127.0.0.1:9/v1") + + assert request_body["messages"][0] == {"role": "system", "content": "sys", "cache_control": {"type": "ephemeral"}} + assert "prompt_cache_breakpoint" not in json.dumps(request_body["messages"]) + assert "prompt_cache_options" not in json.dumps(request_body) + + +@pytest.mark.usefixtures("_no_openai_api_base_override") +def test_completion_custom_base_url_sends_no_prompt_cache_breakpoint_for_gpt_5_6(): + from openai import OpenAI + + client = OpenAI(api_key="fake-api-key", base_url="http://127.0.0.1:9/v1") + request_body = _openai_chat_create_kwargs(client, model="gpt-5.6", base_url="http://127.0.0.1:9/v1") + + assert request_body["messages"][0] == {"role": "system", "content": "sys", "cache_control": {"type": "ephemeral"}} + assert "prompt_cache_breakpoint" not in json.dumps(request_body["messages"]) + assert "prompt_cache_options" not in json.dumps(request_body) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("_no_openai_api_base_override") +async def test_acompletion_custom_base_url_sends_no_prompt_cache_breakpoint_for_gpt_5_6(): + from openai import AsyncOpenAI + + client = AsyncOpenAI(api_key="fake-api-key", base_url="http://127.0.0.1:9/v1") + with patch.object(client.chat.completions.with_raw_response, "create") as mock_create: + with contextlib.suppress(Exception): + await litellm.acompletion( + model="gpt-5.6", + messages=[{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}], + cache_control_injection_points=[{"location": "message", "role": "system"}], + client=client, + base_url="http://127.0.0.1:9/v1", + ) + + mock_create.assert_called_once() + request_body = mock_create.call_args.kwargs + + assert request_body["messages"][0] == {"role": "system", "content": "sys", "cache_control": {"type": "ephemeral"}} + assert "prompt_cache_breakpoint" not in json.dumps(request_body["messages"]) + assert "prompt_cache_options" not in json.dumps(request_body) + + +@pytest.mark.usefixtures("_no_openai_api_base_override") +def test_completion_default_api_base_sends_prompt_cache_breakpoint_for_gpt_5_6(): + from openai import OpenAI + + client = OpenAI(api_key="fake-api-key") + request_body = _openai_chat_create_kwargs(client, model="gpt-5.6") + + assert request_body["messages"][0]["content"] == [ + {"type": "text", "text": "sys", "prompt_cache_breakpoint": {"mode": "explicit"}} + ] + assert request_body["extra_body"]["prompt_cache_options"] == {"mode": "explicit"} + + +_SUBSCRIPTION_OAUTH_CREDENTIAL = "Bearer sk-ant-oat01-fake-subscription-token-for-testing-0123456789" + + +def _scoped_headers_for_oauth_request(): + from litellm.types.utils import ProviderSpecificHeader + + return [ + ProviderSpecificHeader( + custom_llm_provider="anthropic,bedrock,vertex_ai", + extra_headers={"anthropic-version": "2023-06-01"}, + ), + ProviderSpecificHeader( + custom_llm_provider="anthropic", + extra_headers={"authorization": _SUBSCRIPTION_OAUTH_CREDENTIAL}, + ), + ] + + +def _run_anthropic_hop_with_shared_headers(shared_headers): + litellm.completion( + model="anthropic/claude-3-5-sonnet-20240620", + messages=[{"role": "user", "content": "Say OK"}], + extra_headers=shared_headers, + provider_specific_header=_scoped_headers_for_oauth_request(), + api_key="sk-fake-anthropic-key", + mock_response="OK", + ) + + +def test_completion_does_not_mutate_caller_supplied_headers(): + shared_headers = {"x-tenant": "acme"} + + _run_anthropic_hop_with_shared_headers(shared_headers) + + assert shared_headers == {"x-tenant": "acme"} + + +def test_anthropic_oauth_credential_does_not_persist_into_next_provider_hop(): + shared_headers = {"x-tenant": "acme"} + + _run_anthropic_hop_with_shared_headers(shared_headers) + + leaked = [name for name, value in shared_headers.items() if value == _SUBSCRIPTION_OAUTH_CREDENTIAL] + assert leaked == [] + assert "anthropic-version" not in shared_headers + + +STREAM_COST_MODEL = "gpt-4o" +STREAMED_USAGE = {"prompt_tokens": 137, "completion_tokens": 42, "total_tokens": 179} + + +def _text_chunk(content, finish_reason=None, usage=None): + chunk = { + "id": "chatcmpl-stream-cost", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": STREAM_COST_MODEL, + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": content}, + "finish_reason": finish_reason, + } + ], + } + if usage is not None: + chunk["usage"] = usage + return chunk + + +def _priced_at(prompt_tokens, completion_tokens): + prices = litellm.model_cost[STREAM_COST_MODEL] + return ( + prompt_tokens * prices["input_cost_per_token"] + + completion_tokens * prices["output_cost_per_token"] + ) + + +@pytest.fixture +def local_cost_map(monkeypatch): + """The prices these tests assert are the checked-in ones. Setting the environment + variable alone does not reload the map, so pin the map itself. + + Prices are read through two separate lru_caches, so pinning ``model_cost`` is not + enough on its own: an entry warmed against the network-fetched map keeps its old + prices and billing reads those while the assertions read the pinned map. + ``_invalidate_model_cost_lowercase_map`` clears both caches, where + ``get_model_info.cache_clear`` reaches only one. Invalidate on the way in and out + so entries never leak across tests in either direction.""" + from litellm.utils import _invalidate_model_cost_lowercase_map + + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + _invalidate_model_cost_lowercase_map() + yield + _invalidate_model_cost_lowercase_map() + + +def test_a_streamed_response_bills_the_usage_the_provider_reported(local_cost_map): + rebuilt = litellm.stream_chunk_builder( + chunks=[ + _text_chunk("Hello"), + _text_chunk(" there"), + _text_chunk(None, finish_reason="stop", usage=STREAMED_USAGE), + ], + messages=[{"role": "user", "content": "hi"}], + ) + + assert rebuilt.choices[0].message.content == "Hello there" + assert rebuilt.usage.prompt_tokens == STREAMED_USAGE["prompt_tokens"] + assert rebuilt.usage.completion_tokens == STREAMED_USAGE["completion_tokens"] + + cost = litellm.completion_cost(completion_response=rebuilt, model=STREAM_COST_MODEL) + + assert cost == pytest.approx(_priced_at(137, 42)) + + +def test_streaming_and_not_streaming_bill_the_same_usage_the_same(local_cost_map): + rebuilt = litellm.stream_chunk_builder( + chunks=[ + _text_chunk("Hello"), + _text_chunk(" there"), + _text_chunk(None, finish_reason="stop", usage=STREAMED_USAGE), + ], + messages=[{"role": "user", "content": "hi"}], + ) + whole = litellm.ModelResponse( + id="chatcmpl-stream-cost", + model=STREAM_COST_MODEL, + object="chat.completion", + created=1700000000, + choices=[ + { + "index": 0, + "message": {"role": "assistant", "content": "Hello there"}, + "finish_reason": "stop", + } + ], + usage=STREAMED_USAGE, + ) + + assert litellm.completion_cost( + completion_response=rebuilt, model=STREAM_COST_MODEL + ) == pytest.approx(litellm.completion_cost(completion_response=whole, model=STREAM_COST_MODEL)) + + +def test_a_stream_that_reported_no_usage_is_still_billed(local_cost_map): + rebuilt = litellm.stream_chunk_builder( + chunks=[ + _text_chunk("Hello"), + _text_chunk(" there"), + _text_chunk(None, finish_reason="stop"), + ], + messages=[{"role": "user", "content": "hi"}], + ) + + assert rebuilt.usage.prompt_tokens > 0 + assert rebuilt.usage.completion_tokens > 0 + + cost = litellm.completion_cost(completion_response=rebuilt, model=STREAM_COST_MODEL) + + assert cost > 0 + assert cost == pytest.approx( + _priced_at(rebuilt.usage.prompt_tokens, rebuilt.usage.completion_tokens) + ) + + +@pytest.mark.asyncio +async def test_acompletion_resolves_provider_from_api_base(): + response = await litellm.acompletion( + model="deepseek-chat", + api_base="https://api.deepseek.com/v1", + api_key="fake-key", + messages=[{"role": "user", "content": "hi"}], + mock_response="resolved", + ) + + assert response.choices[0].message.content == "resolved" + + +@dataclass(frozen=True, slots=True) +class _RecordedSpeechSuccess: + call_type: str | None + spend_metadata: Mapping[str, object] + response_cost: float | None + logged_response_cost: float | None + + +def _record_speech_success(payload: dict[str, object]) -> _RecordedSpeechSuccess: + call_type: Final = payload.get("call_type") + response_cost: Final = payload.get("response_cost") + logging_payload: Final = payload.get("standard_logging_object") + logged_cost: Final = logging_payload.get("response_cost") if isinstance(logging_payload, dict) else None + return _RecordedSpeechSuccess( + call_type=call_type if isinstance(call_type, str) else None, + spend_metadata=get_litellm_metadata_from_kwargs(payload), + response_cost=response_cost if isinstance(response_cost, float) else None, + logged_response_cost=logged_cost if isinstance(logged_cost, float) else None, + ) + + +class _SuccessEventRecorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.events: list[_RecordedSpeechSuccess] = [] # mutable-ok: test recorder of success-callback events + + async def async_log_success_event( + self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object + ) -> None: + self.events.append(_record_speech_success(kwargs)) + + +async def _wait_for_success_event(recorder: _SuccessEventRecorder, call_type: str) -> _RecordedSpeechSuccess: + for _ in range(100): + if (event := next((e for e in recorder.events if e.call_type == call_type), None)) is not None: + return event + await asyncio.sleep(0.05) + pytest.fail(f"no {call_type} success event; got {[e.call_type for e in recorder.events]}") + + +def _gemini_tts_generate_content_response() -> dict[str, object]: + return { + "candidates": [ + { + "content": { + "parts": [ + { + "inlineData": { + "mimeType": "audio/L16;codec=pcm;rate=24000", + "data": base64.b64encode(b"pcm-audio-bytes").decode(), + } + } + ], + "role": "model", + }, + "finishReason": "STOP", + "index": 0, + } + ], + "usageMetadata": { + "promptTokenCount": 5, + "candidatesTokenCount": 60, + "totalTokenCount": 65, + "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 5}], + "candidatesTokensDetails": [{"modality": "AUDIO", "tokenCount": 60}], + }, + "modelVersion": "gemini-2.5-flash-preview-tts", + } + + +@pytest.mark.asyncio +async def test_aspeech_gemini_bridge_keeps_proxy_metadata_for_spend_tracking( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.delenv("GEMINI_API_KEY", raising=False) + monkeypatch.delenv("GOOGLE_API_KEY", raising=False) + recorder: Final = _SuccessEventRecorder() + monkeypatch.setattr(litellm, "callbacks", [recorder]) + mock_route: Final = respx_mock.post( + url__regex=r"https://generativelanguage\.googleapis\.com/v1beta/models/gemini-2\.5-flash-preview-tts:generateContent.*" + ).mock(return_value=httpx.Response(200, json=_gemini_tts_generate_content_response())) + + await litellm.aspeech( + model="gemini/gemini-2.5-flash-preview-tts", + input="spend tracking check", + voice="Kore", + api_key="fake-gemini-key", + metadata={"user_api_key": "hashed-virtual-key", "user_api_key_user_id": "user-1"}, + ) + + assert mock_route.called + assert mock_route.calls.last.request.headers["x-goog-api-key"] == "fake-gemini-key" + speech_event: Final = await _wait_for_success_event(recorder, call_type="aspeech") + assert speech_event.spend_metadata["user_api_key"] == "hashed-virtual-key" + assert speech_event.spend_metadata["user_api_key_user_id"] == "user-1" + expected_prompt_cost, expected_completion_cost = litellm.cost_per_token( + model="gemini/gemini-2.5-flash-preview-tts", + usage_object=Usage(prompt_tokens=5, completion_tokens=60, total_tokens=65), + ) + expected_cost: Final = expected_prompt_cost + expected_completion_cost + assert expected_cost > 0 + assert speech_event.response_cost == pytest.approx(expected_cost) + assert speech_event.logged_response_cost == pytest.approx(expected_cost) + + +def _stream_builder_text_chunk(model: str, content: str, finish_reason: str | None = None) -> ModelResponseStream: + return ModelResponseStream( + id="chatcmpl-cost", + created=1724900000, + model=model, + object="chat.completion.chunk", + choices=[StreamingChoices(finish_reason=finish_reason, index=0, delta=Delta(content=content, role="assistant"))], + ) + + +def test_stream_chunk_builder_sets_hidden_response_cost_for_known_model(): + chunks: Final = [ + _stream_builder_text_chunk("gpt-4o", "Hello "), + _stream_builder_text_chunk("gpt-4o", "world.", finish_reason="stop"), + ] + + response: Final = litellm.stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "hi"}]) + + assert response is not None + prompt_cost, completion_cost = litellm.cost_per_token(model="gpt-4o", usage_object=response.usage) + expected_cost: Final = prompt_cost + completion_cost + assert expected_cost > 0 + assert response._hidden_params["response_cost"] == pytest.approx(expected_cost) + + +def test_stream_chunk_builder_unknown_model_leaves_response_cost_unset(): + chunks: Final = [ + _stream_builder_text_chunk("totally-unknown-model-xyz", "Hello "), + _stream_builder_text_chunk("totally-unknown-model-xyz", "world.", finish_reason="stop"), + ] + + response: Final = litellm.stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "hi"}]) + + assert response is not None + assert response._hidden_params.get("response_cost") is None + assert response.choices[0].message.content == "Hello world." + + +def test_stream_chunk_builder_prices_proxy_alias_via_model_map(): + chunks: Final = [ + _stream_builder_text_chunk("claude-opus-5", "Hello "), + _stream_builder_text_chunk("claude-opus-5", "world.", finish_reason="stop"), + ] + for chunk in chunks: + chunk._hidden_params = {"custom_llm_provider": "openai"} + + response: Final = litellm.stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "hi"}]) + + assert response is not None + assert response._hidden_params["custom_llm_provider"] == "openai" + prompt_cost, completion_cost = litellm.cost_per_token(model="claude-opus-5", usage_object=response.usage) + expected_cost: Final = prompt_cost + completion_cost + assert expected_cost > 0 + assert response._hidden_params["response_cost"] == pytest.approx(expected_cost) + + +def _stream_builder_logging_obj(model: str = "gpt-4o", custom_llm_provider: str = "openai") -> LiteLLMLogging: + logging_obj: Final = LiteLLMLogging( + model=model, + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="completion", + start_time=datetime.now(), + litellm_call_id="test-call-id", + function_id="test-function-id", + ) + logging_obj.update_environment_variables( + model=model, + user=None, + optional_params={}, + litellm_params={"custom_llm_provider": custom_llm_provider}, + custom_llm_provider=custom_llm_provider, + ) + return logging_obj + + +def test_stream_chunk_builder_stamps_streaming_usage_cost_by_default(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", False) + chunks: Final = [ + _stream_builder_text_chunk("gpt-4o", "Hello "), + _stream_builder_text_chunk("gpt-4o", "world.", finish_reason="stop"), + ] + + response: Final = litellm.stream_chunk_builder( + chunks=chunks, messages=[{"role": "user", "content": "hi"}], logging_obj=_stream_builder_logging_obj() + ) + + assert response is not None + usage_cost: Final = getattr(response.usage, "cost", None) + assert usage_cost is not None + assert usage_cost > 0 + assert response._hidden_params["response_cost"] == pytest.approx(usage_cost) + + +def test_stream_chunk_builder_skips_stamp_when_cost_is_unpriceable(): + import time as time_module + + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging + + logging_obj: Final = LiteLLMLogging( + model="us.anthropic.claude-opus-5", + messages=[{"role": "user", "content": "hi"}], + stream=True, + call_type="completion", + start_time=time_module.time(), + litellm_call_id="stream-builder-alias-unpriceable", + function_id="1", + ) + logging_obj.model_call_details["custom_llm_provider"] = "bedrock" + logging_obj.optional_params = {} + usage_chunk: Final = _stream_builder_text_chunk("bedrock-claude-opus-5", "") + usage_chunk.usage = Usage(prompt_tokens=40, completion_tokens=5, total_tokens=45) + chunks: Final = [ + _stream_builder_text_chunk("bedrock-claude-opus-5", "Hello ", finish_reason="stop"), + usage_chunk, + ] + + response: Final = litellm.stream_chunk_builder( + chunks=chunks, messages=[{"role": "user", "content": "hi"}], logging_obj=logging_obj + ) + + assert response is not None + assert getattr(response.usage, "cost", None) is None + assert response._hidden_params.get("response_cost") is None + + +def test_stream_chunk_builder_keeps_provider_reported_usage_cost(): + usage_chunk: Final = _stream_builder_text_chunk("gpt-4o", "") + usage_chunk.usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15, cost=0.5) + chunks: Final = [ + _stream_builder_text_chunk("gpt-4o", "Hello "), + _stream_builder_text_chunk("gpt-4o", "world.", finish_reason="stop"), + usage_chunk, + ] + + response: Final = litellm.stream_chunk_builder( + chunks=chunks, messages=[{"role": "user", "content": "hi"}], logging_obj=_stream_builder_logging_obj() + ) + + assert response is not None + assert getattr(response.usage, "cost", None) == pytest.approx(0.5) + assert response._hidden_params["response_cost"] == pytest.approx(0.5) + + +def test_stream_chunk_builder_prices_alias_from_openai_sdk_usage_chunk(): + from openai.types.completion_usage import CompletionUsage + + usage_chunk: Final = _stream_builder_text_chunk("mantle-claude", "") + usage_chunk.usage = CompletionUsage(prompt_tokens=20, completion_tokens=60, total_tokens=80, cost=0.000704) + assert type(usage_chunk.usage) is CompletionUsage + chunks: Final = [ + _stream_builder_text_chunk("mantle-claude", "Hello "), + _stream_builder_text_chunk("mantle-claude", "world.", finish_reason="stop"), + usage_chunk, + ] + + response: Final = litellm.stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "hi"}]) + + assert response is not None + assert response.usage.prompt_tokens == 20 + assert response.usage.completion_tokens == 60 + assert getattr(response.usage, "cost", None) == pytest.approx(0.000704) + assert response._hidden_params["response_cost"] == pytest.approx(0.000704) + + +def test_stream_chunk_builder_leaves_xai_reported_cost_to_the_calculator(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "cost_margin_config", {"xai": 0.5}) + usage_chunk: Final = _stream_builder_text_chunk("grok-4", "") + usage_chunk.usage = Usage(prompt_tokens=5, completion_tokens=2, total_tokens=7, cost=0.42) + chunks: Final = [ + _stream_builder_text_chunk("grok-4", "Hello "), + _stream_builder_text_chunk("grok-4", "world.", finish_reason="stop"), + usage_chunk, + ] + logging_obj: Final = _stream_builder_logging_obj(model="grok-4", custom_llm_provider="xai") + + response: Final = litellm.stream_chunk_builder( + chunks=chunks, messages=[{"role": "user", "content": "hi"}], logging_obj=logging_obj + ) + + assert response is not None + assert getattr(response.usage, "cost", None) == pytest.approx(0.42) + assert response._hidden_params.get("response_cost") is None + assert logging_obj._response_cost_calculator(result=response) == pytest.approx(0.63) + + +def test_speech_mistral_dispatches_and_decodes_audio(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("MISTRAL_API_KEY", "sk-mistral-test") + audio_bytes: Final = b"ID3-fake-mp3-bytes" + mock_route: Final = respx_mock.post("https://api.mistral.ai/v1/audio/speech").mock( + return_value=httpx.Response(200, json={"audio_data": base64.b64encode(audio_bytes).decode()}) + ) + + response: Final = litellm.speech( + model="mistral/voxtral-mini-tts-2603", + input="hello from litellm", + voice="en_paul_neutral", + response_format="wav", + speed=2, + instructions="sound cheerful", + ) + + assert mock_route.called + request_body: Final = json.loads(mock_route.calls.last.request.content) + assert request_body == { + "model": "voxtral-mini-tts-2603", + "input": "hello from litellm", + "voice_id": "en_paul_neutral", + "response_format": "wav", + } + assert mock_route.calls.last.request.headers["authorization"] == "Bearer sk-mistral-test" + assert response.content == audio_bytes + + +def test_speech_mistral_routes_to_configured_api_base(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("MISTRAL_API_KEY", "sk-mistral-test") + audio_bytes: Final = b"ID3-gateway-bytes" + gateway_route: Final = respx_mock.post("https://mistral.gateway.internal/v1/audio/speech").mock( + return_value=httpx.Response(200, json={"audio_data": base64.b64encode(audio_bytes).decode()}) + ) + + response: Final = litellm.speech( + model="mistral/voxtral-mini-tts-2603", + input="hello from litellm", + voice="en_paul_neutral", + api_base="https://mistral.gateway.internal", + ) + + assert gateway_route.called + assert response.content == audio_bytes + + +FOUNDRY_HOST: Final = "https://my-project.services.ai.azure.com" + + +def test_azure_ai_transcription_on_a_foundry_host_uses_the_azure_openai_deployment_route( + respx_mock: respx.MockRouter, +): + route: Final = respx_mock.post( + url__regex=r"https://my-project\.services\.ai\.azure\.com/openai/deployments/whisper-1/audio/transcriptions\?api-version=.+" + ).mock(return_value=httpx.Response(200, json={"text": "hello"})) + + response: Final = litellm.transcription( + model="azure_ai/whisper-1", + file=("tone.wav", b"RIFF\x00\x00\x00\x00WAVE", "audio/wav"), + api_base=FOUNDRY_HOST, + api_key="fake-key", + ) + + assert route.called + assert response.text == "hello" + + +def test_azure_ai_speech_on_a_foundry_host_uses_the_azure_openai_deployment_route( + respx_mock: respx.MockRouter, +): + route: Final = respx_mock.post( + url__regex=r"https://my-project\.services\.ai\.azure\.com/openai/deployments/tts-1/audio/speech\?api-version=.+" + ).mock(return_value=httpx.Response(200, content=b"mp3-bytes")) + + response: Final = litellm.speech( + model="azure_ai/tts-1", + input="hello", + voice="alloy", + api_base=FOUNDRY_HOST, + api_key="fake-key", + ) + + assert route.called + assert response.content == b"mp3-bytes" + + +FORWARDED_CLIENT_HEADERS: Final = {"x-forwarded-for": "10.0.0.1", "x-amzn-trace-id": "Root=1-lit7694"} + + +def _chat_completion_json() -> Mapping[str, object]: + return { + "id": "chatcmpl-lit7694", + "object": "chat.completion", + "created": 1, + "model": "gpt-5.4", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } + + +def _chat_completion_sse() -> bytes: + chunk: Final = { + "id": "chatcmpl-lit7694", + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-5.4", + "choices": [{"index": 0, "delta": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + } + return f"data: {json.dumps(chunk)}\n\ndata: [DONE]\n\n".encode() + + +@pytest.mark.parametrize("stream", [False, True]) +def test_bridged_responses_with_openai_http_handler_keeps_forwarded_headers_out_of_the_body( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, stream: bool +): + monkeypatch.setenv("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER", "true") + route: Final = respx_mock.post("https://api.openai.com/v1/chat/completions").mock( + return_value=httpx.Response(200, content=_chat_completion_sse(), headers={"content-type": "text/event-stream"}) + if stream + else httpx.Response(200, json=_chat_completion_json()) + ) + + response: Final = litellm.responses( + model="openai/gpt-5.4", + input="Reply with the single word ok", + stream=stream, + use_chat_completions_api=True, + headers=dict(FORWARDED_CLIENT_HEADERS), + api_key="sk-test", + ) + if stream: + list(response) + + assert route.called + request: Final = route.calls.last.request + body: Final = json.loads(request.content) + assert "extra_headers" not in body + assert body["model"] == "gpt-5.4" + assert {k: request.headers[k] for k in FORWARDED_CLIENT_HEADERS} == FORWARDED_CLIENT_HEADERS + + +@pytest.mark.parametrize("http2_on", [True, False]) +def test_aiohttp_openai_warns_only_when_http2_enabled( + monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, http2_on: bool +): + from litellm.main import base_llm_aiohttp_handler + + monkeypatch.setattr(litellm, "http2", http2_on) + monkeypatch.delenv("LITELLM_HTTP2", raising=False) + + handler_completion: Final = MagicMock(return_value=MagicMock()) + monkeypatch.setattr(base_llm_aiohttp_handler, "completion", handler_completion) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + litellm.completion( + model="aiohttp_openai/gpt-4o", + messages=[{"role": "user", "content": "hi"}], + api_key="sk-test", + ) + + assert handler_completion.called + warned: Final = "aiohttp_openai/ always uses aiohttp" in caplog.text + assert warned is http2_on + + +@pytest.mark.parametrize("tool_choice", [{"type": "bogus"}, {"name": "lookup_fruit"}, {"type": "file_search"}]) +def test_completion_rejects_untranslatable_tool_choice_with_a_400(tool_choice): + with pytest.raises(litellm.BadRequestError) as exc_info: + litellm.completion( + model="anthropic/claude-haiku-4-5", + messages=[{"role": "user", "content": "Which fruit is red?"}], + tools=[{"type": "function", "function": {"name": "lookup_fruit", "parameters": {"type": "object"}}}], + tool_choice=tool_choice, + api_key="sk-unused", + ) + assert exc_info.value.status_code == 400 + assert f"tool_choice={tool_choice}" in str(exc_info.value) diff --git a/tests/test_litellm/test_main_module_header.py b/tests/unit/test_main_module_header.py similarity index 100% rename from tests/test_litellm/test_main_module_header.py rename to tests/unit/test_main_module_header.py diff --git a/tests/test_litellm/test_mistral_medium_3_5_model_metadata.py b/tests/unit/test_mistral_medium_3_5_model_metadata.py similarity index 100% rename from tests/test_litellm/test_mistral_medium_3_5_model_metadata.py rename to tests/unit/test_mistral_medium_3_5_model_metadata.py diff --git a/tests/test_litellm/test_mistral_small_4_0_model_metadata.py b/tests/unit/test_mistral_small_4_0_model_metadata.py similarity index 100% rename from tests/test_litellm/test_mistral_small_4_0_model_metadata.py rename to tests/unit/test_mistral_small_4_0_model_metadata.py diff --git a/tests/test_litellm/test_mistral_zai_glm_5_2_model_metadata.py b/tests/unit/test_mistral_zai_glm_5_2_model_metadata.py similarity index 100% rename from tests/test_litellm/test_mistral_zai_glm_5_2_model_metadata.py rename to tests/unit/test_mistral_zai_glm_5_2_model_metadata.py diff --git a/tests/test_litellm/test_model_block_unblock.py b/tests/unit/test_model_block_unblock.py similarity index 100% rename from tests/test_litellm/test_model_block_unblock.py rename to tests/unit/test_model_block_unblock.py diff --git a/tests/test_litellm/test_model_cost_aliases.py b/tests/unit/test_model_cost_aliases.py similarity index 100% rename from tests/test_litellm/test_model_cost_aliases.py rename to tests/unit/test_model_cost_aliases.py diff --git a/tests/test_litellm/test_model_param_helper.py b/tests/unit/test_model_param_helper.py similarity index 100% rename from tests/test_litellm/test_model_param_helper.py rename to tests/unit/test_model_param_helper.py diff --git a/tests/test_litellm/test_model_prices_schema.py b/tests/unit/test_model_prices_schema.py similarity index 100% rename from tests/test_litellm/test_model_prices_schema.py rename to tests/unit/test_model_prices_schema.py diff --git a/tests/test_litellm/test_model_response_normalization.py b/tests/unit/test_model_response_normalization.py similarity index 100% rename from tests/test_litellm/test_model_response_normalization.py rename to tests/unit/test_model_response_normalization.py diff --git a/tests/test_litellm/test_muse_spark_1_1_model_metadata.py b/tests/unit/test_muse_spark_1_1_model_metadata.py similarity index 100% rename from tests/test_litellm/test_muse_spark_1_1_model_metadata.py rename to tests/unit/test_muse_spark_1_1_model_metadata.py diff --git a/tests/test_litellm/test_muse_spark_1_2_model_metadata.py b/tests/unit/test_muse_spark_1_2_model_metadata.py similarity index 100% rename from tests/test_litellm/test_muse_spark_1_2_model_metadata.py rename to tests/unit/test_muse_spark_1_2_model_metadata.py diff --git a/tests/test_litellm/test_muse_spark_1_3_model_metadata.py b/tests/unit/test_muse_spark_1_3_model_metadata.py similarity index 100% rename from tests/test_litellm/test_muse_spark_1_3_model_metadata.py rename to tests/unit/test_muse_spark_1_3_model_metadata.py diff --git a/tests/test_litellm/test_mutation_report.py b/tests/unit/test_mutation_report.py similarity index 100% rename from tests/test_litellm/test_mutation_report.py rename to tests/unit/test_mutation_report.py diff --git a/tests/test_litellm/test_nested_drop_params.py b/tests/unit/test_nested_drop_params.py similarity index 100% rename from tests/test_litellm/test_nested_drop_params.py rename to tests/unit/test_nested_drop_params.py diff --git a/tests/test_litellm/test_non_chat_routes_open_llm_spans.py b/tests/unit/test_non_chat_routes_open_llm_spans.py similarity index 100% rename from tests/test_litellm/test_non_chat_routes_open_llm_spans.py rename to tests/unit/test_non_chat_routes_open_llm_spans.py diff --git a/tests/test_litellm/test_openai_embedding_encoding_format_default.py b/tests/unit/test_openai_embedding_encoding_format_default.py similarity index 100% rename from tests/test_litellm/test_openai_embedding_encoding_format_default.py rename to tests/unit/test_openai_embedding_encoding_format_default.py diff --git a/tests/test_litellm/test_openai_service_tier_long_context_pricing.py b/tests/unit/test_openai_service_tier_long_context_pricing.py similarity index 100% rename from tests/test_litellm/test_openai_service_tier_long_context_pricing.py rename to tests/unit/test_openai_service_tier_long_context_pricing.py diff --git a/tests/test_litellm/test_pre_commit_lint.py b/tests/unit/test_pre_commit_lint.py similarity index 100% rename from tests/test_litellm/test_pre_commit_lint.py rename to tests/unit/test_pre_commit_lint.py diff --git a/tests/test_litellm/test_prisma_generate_if_needed.py b/tests/unit/test_prisma_generate_if_needed.py similarity index 100% rename from tests/test_litellm/test_prisma_generate_if_needed.py rename to tests/unit/test_prisma_generate_if_needed.py diff --git a/tests/test_litellm/test_process_helpers.py b/tests/unit/test_process_helpers.py similarity index 100% rename from tests/test_litellm/test_process_helpers.py rename to tests/unit/test_process_helpers.py diff --git a/tests/test_litellm/test_project_alias_tracking.py b/tests/unit/test_project_alias_tracking.py similarity index 100% rename from tests/test_litellm/test_project_alias_tracking.py rename to tests/unit/test_project_alias_tracking.py diff --git a/tests/test_litellm/test_project_tags_pydantic.py b/tests/unit/test_project_tags_pydantic.py similarity index 100% rename from tests/test_litellm/test_project_tags_pydantic.py rename to tests/unit/test_project_tags_pydantic.py diff --git a/tests/test_litellm/test_proxy_auth.py b/tests/unit/test_proxy_auth.py similarity index 100% rename from tests/test_litellm/test_proxy_auth.py rename to tests/unit/test_proxy_auth.py diff --git a/tests/test_litellm/test_rag_openai_ingestion.py b/tests/unit/test_rag_openai_ingestion.py similarity index 100% rename from tests/test_litellm/test_rag_openai_ingestion.py rename to tests/unit/test_rag_openai_ingestion.py diff --git a/tests/test_litellm/test_rate_limit_error_unification.py b/tests/unit/test_rate_limit_error_unification.py similarity index 100% rename from tests/test_litellm/test_rate_limit_error_unification.py rename to tests/unit/test_rate_limit_error_unification.py diff --git a/tests/test_litellm/test_read_rc_version.py b/tests/unit/test_read_rc_version.py similarity index 100% rename from tests/test_litellm/test_read_rc_version.py rename to tests/unit/test_read_rc_version.py diff --git a/tests/test_litellm/test_redact_string_in_error_paths.py b/tests/unit/test_redact_string_in_error_paths.py similarity index 100% rename from tests/test_litellm/test_redact_string_in_error_paths.py rename to tests/unit/test_redact_string_in_error_paths.py diff --git a/tests/test_litellm/test_redis.py b/tests/unit/test_redis.py similarity index 100% rename from tests/test_litellm/test_redis.py rename to tests/unit/test_redis.py diff --git a/tests/test_litellm/test_redis_credential_provider.py b/tests/unit/test_redis_credential_provider.py similarity index 100% rename from tests/test_litellm/test_redis_credential_provider.py rename to tests/unit/test_redis_credential_provider.py diff --git a/tests/test_litellm/test_register_model_custom_pricing.py b/tests/unit/test_register_model_custom_pricing.py similarity index 100% rename from tests/test_litellm/test_register_model_custom_pricing.py rename to tests/unit/test_register_model_custom_pricing.py diff --git a/tests/test_litellm/test_register_model_zero_cost_persistence.py b/tests/unit/test_register_model_zero_cost_persistence.py similarity index 100% rename from tests/test_litellm/test_register_model_zero_cost_persistence.py rename to tests/unit/test_register_model_zero_cost_persistence.py diff --git a/tests/test_litellm/test_replicate_model_key_format.py b/tests/unit/test_replicate_model_key_format.py similarity index 100% rename from tests/test_litellm/test_replicate_model_key_format.py rename to tests/unit/test_replicate_model_key_format.py diff --git a/tests/test_litellm/test_responses_api_bridge_non_stream.py b/tests/unit/test_responses_api_bridge_non_stream.py similarity index 100% rename from tests/test_litellm/test_responses_api_bridge_non_stream.py rename to tests/unit/test_responses_api_bridge_non_stream.py diff --git a/tests/test_litellm/test_responses_id_security.py b/tests/unit/test_responses_id_security.py similarity index 94% rename from tests/test_litellm/test_responses_id_security.py rename to tests/unit/test_responses_id_security.py index a6081670172..704a52fc202 100644 --- a/tests/test_litellm/test_responses_id_security.py +++ b/tests/unit/test_responses_id_security.py @@ -4,7 +4,7 @@ Tests for ResponsesIDSecurity hook. Tests the security hook that prevents user B from seeing response from user A. """ -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import MagicMock, patch import pytest from fastapi import HTTPException @@ -113,63 +113,6 @@ class TestDecryptResponseId: assert team_id is None -class TestEncryptResponseId: - """Test _encrypt_response_id function""" - - @pytest.mark.skip( - reason="Flaky on CI; disabling temporarily until responses_id_security is fixed" - ) - def test_encrypt_response_id_success( - self, responses_id_security, mock_user_api_key_dict - ): - """Test encrypting a response ID with user information""" - mock_response = ResponsesAPIResponse( - id="resp_123", created_at=1234567890, output=[], status="completed" - ) - - with patch( - "litellm.proxy.hooks.responses_id_security.encrypt_value_helper" - ) as mock_encrypt: - mock_encrypt.return_value = "encrypted_base64_value" - - with patch.object( - responses_id_security, "_get_signing_key", return_value="test-key" - ): - result = responses_id_security._encrypt_response_id( - mock_response, mock_user_api_key_dict - ) - - assert result.id == "resp_encrypted_base64_value" - assert result.id.startswith("resp_") - mock_encrypt.assert_called_once() - - @pytest.mark.skip( - reason="Flaky on CI; disabling temporarily until responses_id_security is fixed" - ) - def test_encrypt_response_id_maintains_prefix( - self, responses_id_security, mock_user_api_key_dict - ): - """Test that encrypted response ID maintains 'resp_' prefix""" - mock_response = ResponsesAPIResponse( - id="resp_456", created_at=1234567890, output=[], status="in_progress" - ) - - with patch( - "litellm.proxy.common_utils.encrypt_decrypt_utils._get_salt_key", - return_value="test-salt-key", - ): - with patch.object( - responses_id_security, "_get_signing_key", return_value="test-key" - ): - result = responses_id_security._encrypt_response_id( - mock_response, mock_user_api_key_dict - ) - - assert result.id.startswith("resp_") - # The encrypted ID should be different from the original - assert result.id != "resp_456" - - class TestCheckUserAccessToResponseId: """Test check_user_access_to_response_id function""" @@ -857,7 +800,6 @@ class TestAsyncPostCallSuccessHook: assert result == mock_response - _FABRICATED_PROVIDER_RESPONSE_ID = "resp_fabricatedprovideridaaaaaaaaaaaaaaaa" _FABRICATED_UNMANAGED_ID = "resp_fabricatedunmanagedidbbbbbbbbbbbbbbbb" _UNIT_TEST_SALT_KEY = "lit6837-unit-test-salt-key" diff --git a/tests/test_litellm/test_responses_streaming_container_ownership.py b/tests/unit/test_responses_streaming_container_ownership.py similarity index 100% rename from tests/test_litellm/test_responses_streaming_container_ownership.py rename to tests/unit/test_responses_streaming_container_ownership.py diff --git a/tests/test_litellm/test_retrieve_batch_bedrock_dispatch.py b/tests/unit/test_retrieve_batch_bedrock_dispatch.py similarity index 100% rename from tests/test_litellm/test_retrieve_batch_bedrock_dispatch.py rename to tests/unit/test_retrieve_batch_bedrock_dispatch.py diff --git a/tests/test_litellm/test_router.py b/tests/unit/test_router/test_router.py similarity index 100% rename from tests/test_litellm/test_router.py rename to tests/unit/test_router/test_router.py diff --git a/tests/test_litellm/test_router_block_helpers.py b/tests/unit/test_router_block_helpers.py similarity index 100% rename from tests/test_litellm/test_router_block_helpers.py rename to tests/unit/test_router_block_helpers.py diff --git a/tests/test_litellm/test_router_exception_redaction.py b/tests/unit/test_router_exception_redaction.py similarity index 100% rename from tests/test_litellm/test_router_exception_redaction.py rename to tests/unit/test_router_exception_redaction.py diff --git a/tests/test_litellm/test_router_google_genai.py b/tests/unit/test_router_google_genai.py similarity index 100% rename from tests/test_litellm/test_router_google_genai.py rename to tests/unit/test_router_google_genai.py diff --git a/tests/test_litellm/test_router_model_cost_isolation.py b/tests/unit/test_router_model_cost_isolation.py similarity index 100% rename from tests/test_litellm/test_router_model_cost_isolation.py rename to tests/unit/test_router_model_cost_isolation.py diff --git a/tests/test_litellm/test_router_order_fallback.py b/tests/unit/test_router_order_fallback.py similarity index 100% rename from tests/test_litellm/test_router_order_fallback.py rename to tests/unit/test_router_order_fallback.py diff --git a/tests/test_litellm/test_router_per_deployment_num_retries.py b/tests/unit/test_router_per_deployment_num_retries.py similarity index 100% rename from tests/test_litellm/test_router_per_deployment_num_retries.py rename to tests/unit/test_router_per_deployment_num_retries.py diff --git a/tests/test_litellm/test_router_redis_init.py b/tests/unit/test_router_redis_init.py similarity index 100% rename from tests/test_litellm/test_router_redis_init.py rename to tests/unit/test_router_redis_init.py diff --git a/tests/test_litellm/test_router_retry_backoff_headers.py b/tests/unit/test_router_retry_backoff_headers.py similarity index 100% rename from tests/test_litellm/test_router_retry_backoff_headers.py rename to tests/unit/test_router_retry_backoff_headers.py diff --git a/tests/test_litellm/test_router_retry_non_retryable_errors.py b/tests/unit/test_router_retry_non_retryable_errors.py similarity index 100% rename from tests/test_litellm/test_router_retry_non_retryable_errors.py rename to tests/unit/test_router_retry_non_retryable_errors.py diff --git a/tests/test_litellm/test_router_retry_policy_update.py b/tests/unit/test_router_retry_policy_update.py similarity index 100% rename from tests/test_litellm/test_router_retry_policy_update.py rename to tests/unit/test_router_retry_policy_update.py diff --git a/tests/test_litellm/test_router_silent_experiment.py b/tests/unit/test_router_silent_experiment.py similarity index 92% rename from tests/test_litellm/test_router_silent_experiment.py rename to tests/unit/test_router_silent_experiment.py index d62962da275..ab65e09e133 100644 --- a/tests/test_litellm/test_router_silent_experiment.py +++ b/tests/unit/test_router_silent_experiment.py @@ -388,47 +388,6 @@ async def test_shadow_of_a_shadow_is_not_launched(recording_logger): assert model_groups == ["shadow-a"] -def test_silent_experiment_completion_direct(): - """ - Test _silent_experiment_completion directly (for router code coverage). - Mocks router.completion to avoid real API call. - """ - model_list = [ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake-key"}, - }, - ] - router = Router(model_list=model_list) - messages = [{"role": "user", "content": "hi"}] - with patch.object(router, "acompletion", new_callable=AsyncMock, return_value=None): - router._silent_experiment_completion( - silent_model="gpt-3.5-turbo", - messages=messages, - ) - - -@pytest.mark.asyncio -async def test_silent_experiment_acompletion_direct(): - """ - Test _silent_experiment_acompletion directly (for router code coverage). - Mocks router.acompletion to avoid real API call. - """ - model_list = [ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake-key"}, - }, - ] - router = Router(model_list=model_list) - messages = [{"role": "user", "content": "hi"}] - with patch.object(router, "acompletion", new_callable=AsyncMock, return_value=None): - await router._silent_experiment_acompletion( - silent_model="gpt-3.5-turbo", - messages=messages, - ) - - @pytest.mark.asyncio async def test_run_silent_experiment_drains_stream_so_callbacks_fire(recording_logger): router = Router(model_list=_streaming_model_list(None)) @@ -602,3 +561,44 @@ def test_router_silent_experiment_completion(): assert silent_call[1]["model"] == "openai/gpt-4" # Verify model_group is set to the silent model name for correct metric attribution assert silent_call[1]["metadata"]["model_group"] == "silent-model" + + +SILENT_EXPERIMENT_RUNNERS: Final = ( + pytest.param(lambda router, **kwargs: router._silent_experiment_completion(**kwargs), id="sync"), + pytest.param(lambda router, **kwargs: asyncio.run(router._silent_experiment_acompletion(**kwargs)), id="async"), +) + + +@pytest.mark.parametrize("run_silent_experiment", SILENT_EXPERIMENT_RUNNERS) +def test_silent_experiment_sends_shadow_request_attributed_to_the_silent_model(run_silent_experiment): + router = Router(model_list=_streaming_model_list(["shadow-a"])) + primary_metadata: Final = {"model_group": "primary-model"} + with patch.object(router, "acompletion", new_callable=AsyncMock, return_value=None) as acompletion: + run_silent_experiment( + router, + silent_model="shadow-a", + messages=[{"role": "user", "content": "hi"}], + metadata=primary_metadata, + ) + + acompletion.assert_awaited_once() + shadow_call: Final = acompletion.await_args.kwargs + assert shadow_call["model"] == "shadow-a" + assert shadow_call["messages"] == [{"role": "user", "content": "hi"}] + assert shadow_call["metadata"]["model_group"] == "shadow-a" + assert shadow_call["metadata"]["is_silent_experiment"] is True + assert primary_metadata == {"model_group": "primary-model"} + + +@pytest.mark.parametrize("run_silent_experiment", SILENT_EXPERIMENT_RUNNERS) +def test_silent_experiment_does_not_launch_from_a_shadow_request(run_silent_experiment): + router = Router(model_list=_streaming_model_list(["shadow-a"])) + with patch.object(router, "acompletion", new_callable=AsyncMock, return_value=None) as acompletion: + run_silent_experiment( + router, + silent_model="shadow-a", + messages=[{"role": "user", "content": "hi"}], + metadata={"is_silent_experiment": True}, + ) + + acompletion.assert_not_awaited() diff --git a/tests/test_litellm/test_router_streaming_fallback_metadata.py b/tests/unit/test_router_streaming_fallback_metadata.py similarity index 100% rename from tests/test_litellm/test_router_streaming_fallback_metadata.py rename to tests/unit/test_router_streaming_fallback_metadata.py diff --git a/tests/test_litellm/test_router_weighted_failover.py b/tests/unit/test_router_weighted_failover.py similarity index 100% rename from tests/test_litellm/test_router_weighted_failover.py rename to tests/unit/test_router_weighted_failover.py diff --git a/tests/test_litellm/test_ruff_strict_gate.py b/tests/unit/test_ruff_strict_gate.py similarity index 100% rename from tests/test_litellm/test_ruff_strict_gate.py rename to tests/unit/test_ruff_strict_gate.py diff --git a/tests/test_litellm/test_sambanova_model_metadata.py b/tests/unit/test_sambanova_model_metadata.py similarity index 100% rename from tests/test_litellm/test_sambanova_model_metadata.py rename to tests/unit/test_sambanova_model_metadata.py diff --git a/tests/test_litellm/test_secret_redaction.py b/tests/unit/test_secret_redaction.py similarity index 100% rename from tests/test_litellm/test_secret_redaction.py rename to tests/unit/test_secret_redaction.py diff --git a/tests/test_litellm/test_select_ui_test_scope.py b/tests/unit/test_select_ui_test_scope.py similarity index 100% rename from tests/test_litellm/test_select_ui_test_scope.py rename to tests/unit/test_select_ui_test_scope.py diff --git a/tests/test_litellm/test_service_logger.py b/tests/unit/test_service_logger.py similarity index 100% rename from tests/test_litellm/test_service_logger.py rename to tests/unit/test_service_logger.py diff --git a/tests/test_litellm/test_setup_wizard.py b/tests/unit/test_setup_wizard.py similarity index 100% rename from tests/test_litellm/test_setup_wizard.py rename to tests/unit/test_setup_wizard.py diff --git a/tests/test_litellm/test_shared_session_integration.py b/tests/unit/test_shared_session_integration.py similarity index 100% rename from tests/test_litellm/test_shared_session_integration.py rename to tests/unit/test_shared_session_integration.py diff --git a/tests/test_litellm/test_ssl_verify_unit.py b/tests/unit/test_ssl_verify_unit.py similarity index 83% rename from tests/test_litellm/test_ssl_verify_unit.py rename to tests/unit/test_ssl_verify_unit.py index c39362c01a2..f47cdf3e6cd 100644 --- a/tests/test_litellm/test_ssl_verify_unit.py +++ b/tests/unit/test_ssl_verify_unit.py @@ -50,41 +50,6 @@ class TestBaseAWSLLMSSLVerify: # Result depends on environment, just verify it doesn't crash assert result is not None or result is None # Can be None, True, False, or path - @patch("boto3.client") - def test_get_credentials_propagates_ssl_verify(self, mock_boto_client): - """Test that get_credentials propagates ssl_verify to boto3 clients.""" - base_llm = BaseAWSLLM() - - # Mock the boto3 client - mock_sts_client = Mock() - mock_sts_client.assume_role.return_value = { - "Credentials": { - "AccessKeyId": "test_key", - "SecretAccessKey": "test_secret", - "SessionToken": "test_token", - "Expiration": "2026-01-20T00:00:00Z", - } - } - mock_boto_client.return_value = mock_sts_client - - # Call get_credentials with ssl_verify parameter - cert_path = "/path/to/cert.pem" - try: - base_llm.get_credentials( - aws_access_key_id="test_key", - aws_secret_access_key="test_secret", - aws_region_name="us-east-1", - ssl_verify=cert_path, - ) - except Exception: - # May fail due to missing credentials, but we're checking the call - pass - - # Verify boto3.client was called with verify parameter - # Note: This test verifies the parameter is accepted, actual propagation - # is tested in integration tests - assert True # If we got here without error, parameter was accepted - class TestAimGuardrailSSLVerify: """Test SSL verification parameter handling in AimGuardrail.""" diff --git a/tests/test_litellm/test_stream_chunk_builder_annotations.py b/tests/unit/test_stream_chunk_builder_annotations.py similarity index 100% rename from tests/test_litellm/test_stream_chunk_builder_annotations.py rename to tests/unit/test_stream_chunk_builder_annotations.py diff --git a/tests/test_litellm/test_stream_chunk_builder_citations.py b/tests/unit/test_stream_chunk_builder_citations.py similarity index 100% rename from tests/test_litellm/test_stream_chunk_builder_citations.py rename to tests/unit/test_stream_chunk_builder_citations.py diff --git a/tests/test_litellm/test_stream_chunk_builder_images.py b/tests/unit/test_stream_chunk_builder_images.py similarity index 100% rename from tests/test_litellm/test_stream_chunk_builder_images.py rename to tests/unit/test_stream_chunk_builder_images.py diff --git a/tests/test_litellm/test_streaming_connection_cleanup.py b/tests/unit/test_streaming_connection_cleanup.py similarity index 100% rename from tests/test_litellm/test_streaming_connection_cleanup.py rename to tests/unit/test_streaming_connection_cleanup.py diff --git a/tests/test_litellm/test_sync_together_ai_models.py b/tests/unit/test_sync_together_ai_models.py similarity index 100% rename from tests/test_litellm/test_sync_together_ai_models.py rename to tests/unit/test_sync_together_ai_models.py diff --git a/tests/test_litellm/test_system_message_format_bug.py b/tests/unit/test_system_message_format_bug.py similarity index 100% rename from tests/test_litellm/test_system_message_format_bug.py rename to tests/unit/test_system_message_format_bug.py diff --git a/tests/test_litellm/test_test_quality_gate.py b/tests/unit/test_test_quality_gate.py similarity index 100% rename from tests/test_litellm/test_test_quality_gate.py rename to tests/unit/test_test_quality_gate.py diff --git a/tests/test_litellm/test_thinking_enabled.py b/tests/unit/test_thinking_enabled.py similarity index 100% rename from tests/test_litellm/test_thinking_enabled.py rename to tests/unit/test_thinking_enabled.py diff --git a/tests/test_litellm/test_together_ai_model_metadata.py b/tests/unit/test_together_ai_model_metadata.py similarity index 100% rename from tests/test_litellm/test_together_ai_model_metadata.py rename to tests/unit/test_together_ai_model_metadata.py diff --git a/tests/test_litellm/test_type_check_gate.py b/tests/unit/test_type_check_gate.py similarity index 100% rename from tests/test_litellm/test_type_check_gate.py rename to tests/unit/test_type_check_gate.py diff --git a/tests/test_litellm/test_type_discipline_gate.py b/tests/unit/test_type_discipline_gate.py similarity index 100% rename from tests/test_litellm/test_type_discipline_gate.py rename to tests/unit/test_type_discipline_gate.py diff --git a/tests/test_litellm/test_typesafe_model_metadata.py b/tests/unit/test_typesafe_model_metadata.py similarity index 100% rename from tests/test_litellm/test_typesafe_model_metadata.py rename to tests/unit/test_typesafe_model_metadata.py diff --git a/tests/test_litellm/test_unit_shard_missing_paths.py b/tests/unit/test_unit_shard_missing_paths.py similarity index 96% rename from tests/test_litellm/test_unit_shard_missing_paths.py rename to tests/unit/test_unit_shard_missing_paths.py index b91c2cff764..e464402c9d8 100644 --- a/tests/test_litellm/test_unit_shard_missing_paths.py +++ b/tests/unit/test_unit_shard_missing_paths.py @@ -36,8 +36,10 @@ def _run_shard(tmp_path: Path, test_path: str, workers: str) -> subprocess.Compl **os.environ, **_SHARD_ENV, "PATH": f"{shim_dir}{os.pathsep}{os.environ['PATH']}", + "GITHUB_OUTPUT": str(tmp_path / "github_output"), "TEST_PATH": test_path, "WORKERS": workers, + "UNIT_FLAG": "", }, capture_output=True, text=True, diff --git a/tests/test_litellm/test_unit_shard_per_test_timeout.py b/tests/unit/test_unit_shard_per_test_timeout.py similarity index 100% rename from tests/test_litellm/test_unit_shard_per_test_timeout.py rename to tests/unit/test_unit_shard_per_test_timeout.py diff --git a/tests/test_litellm/test_utils.py b/tests/unit/test_utils.py similarity index 100% rename from tests/test_litellm/test_utils.py rename to tests/unit/test_utils.py diff --git a/tests/test_litellm/test_utils_module_docstring.py b/tests/unit/test_utils_module_docstring.py similarity index 100% rename from tests/test_litellm/test_utils_module_docstring.py rename to tests/unit/test_utils_module_docstring.py diff --git a/tests/test_litellm/test_uuid_helper.py b/tests/unit/test_uuid_helper.py similarity index 100% rename from tests/test_litellm/test_uuid_helper.py rename to tests/unit/test_uuid_helper.py diff --git a/tests/test_litellm/test_vcr_safe_body_matcher.py b/tests/unit/test_vcr_safe_body_matcher.py similarity index 98% rename from tests/test_litellm/test_vcr_safe_body_matcher.py rename to tests/unit/test_vcr_safe_body_matcher.py index 712ecf09911..cf4e4a1c276 100644 --- a/tests/test_litellm/test_vcr_safe_body_matcher.py +++ b/tests/unit/test_vcr_safe_body_matcher.py @@ -52,14 +52,6 @@ def test_safe_body_matcher_accepts_str_bytes_equivalent(): _safe_body_matcher(_req("hello"), _req(b"hello")) -def test_safe_body_matcher_handles_jsonl_without_crashing(): - jsonl = ( - b'{"recordId": "request-1", "modelInput": {}}\n' - b'{"recordId": "request-2", "modelInput": {}}\n' - ) - _safe_body_matcher(_req(jsonl), _req(jsonl)) - - def test_safe_body_matcher_rejects_different_jsonl_bodies(): a = b'{"recordId": "request-1"}\n{"recordId": "request-2"}\n' b = b'{"recordId": "request-1"}\n{"recordId": "request-3"}\n' diff --git a/tests/test_litellm/test_vertex_ai_xai_grok_prompt_caching_metadata.py b/tests/unit/test_vertex_ai_xai_grok_prompt_caching_metadata.py similarity index 100% rename from tests/test_litellm/test_vertex_ai_xai_grok_prompt_caching_metadata.py rename to tests/unit/test_vertex_ai_xai_grok_prompt_caching_metadata.py diff --git a/tests/test_litellm/test_video_generation.py b/tests/unit/test_video_generation.py similarity index 100% rename from tests/test_litellm/test_video_generation.py rename to tests/unit/test_video_generation.py diff --git a/tests/test_litellm/test_with_dashboard_node.py b/tests/unit/test_with_dashboard_node.py similarity index 100% rename from tests/test_litellm/test_with_dashboard_node.py rename to tests/unit/test_with_dashboard_node.py diff --git a/tests/test_litellm/test_xai_grok_4_3_model_metadata.py b/tests/unit/test_xai_grok_4_3_model_metadata.py similarity index 100% rename from tests/test_litellm/test_xai_grok_4_3_model_metadata.py rename to tests/unit/test_xai_grok_4_3_model_metadata.py diff --git a/tests/test_litellm/test_xai_responses_auto_routing.py b/tests/unit/test_xai_responses_auto_routing.py similarity index 100% rename from tests/test_litellm/test_xai_responses_auto_routing.py rename to tests/unit/test_xai_responses_auto_routing.py diff --git a/tests/test_litellm/types/test_completion.py b/tests/unit/types/test_completion.py similarity index 99% rename from tests/test_litellm/types/test_completion.py rename to tests/unit/types/test_completion.py index cd51913c5dd..4971a0c7e0a 100644 --- a/tests/test_litellm/types/test_completion.py +++ b/tests/unit/types/test_completion.py @@ -5,7 +5,7 @@ This test suite validates the CompletionRequest model and its compatibility with OpenAI ChatCompletion API message formats. Usage: - pytest tests/test_litellm/types/test_completion.py -v + pytest tests/unit/types/test_completion.py -v """ import dataclasses diff --git a/tests/test_litellm/types/test_guardrails_case_normalization.py b/tests/unit/types/test_guardrails_case_normalization.py similarity index 100% rename from tests/test_litellm/types/test_guardrails_case_normalization.py rename to tests/unit/types/test_guardrails_case_normalization.py diff --git a/tests/test_litellm/types/test_mcp.py b/tests/unit/types/test_mcp.py similarity index 100% rename from tests/test_litellm/types/test_mcp.py rename to tests/unit/types/test_mcp.py diff --git a/tests/test_litellm/types/test_presidio_entity_expansion.py b/tests/unit/types/test_presidio_entity_expansion.py similarity index 100% rename from tests/test_litellm/types/test_presidio_entity_expansion.py rename to tests/unit/types/test_presidio_entity_expansion.py diff --git a/tests/test_litellm/types/test_prometheus_label_value_sanitize.py b/tests/unit/types/test_prometheus_label_value_sanitize.py similarity index 100% rename from tests/test_litellm/types/test_prometheus_label_value_sanitize.py rename to tests/unit/types/test_prometheus_label_value_sanitize.py diff --git a/tests/test_litellm/types/test_prometheus_latency_buckets.py b/tests/unit/types/test_prometheus_latency_buckets.py similarity index 100% rename from tests/test_litellm/types/test_prometheus_latency_buckets.py rename to tests/unit/types/test_prometheus_latency_buckets.py diff --git a/tests/test_litellm/types/test_router.py b/tests/unit/types/test_router.py similarity index 100% rename from tests/test_litellm/types/test_router.py rename to tests/unit/types/test_router.py diff --git a/tests/test_litellm/types/test_types_utils.py b/tests/unit/types/test_types_utils.py similarity index 100% rename from tests/test_litellm/types/test_types_utils.py rename to tests/unit/types/test_types_utils.py diff --git a/tests/test_litellm/types/test_uk_pii_entities.py b/tests/unit/types/test_uk_pii_entities.py similarity index 100% rename from tests/test_litellm/types/test_uk_pii_entities.py rename to tests/unit/types/test_uk_pii_entities.py diff --git a/tests/unit/vector_stores/__init__.py b/tests/unit/vector_stores/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/vector_stores/test_main.py b/tests/unit/vector_stores/test_main.py similarity index 100% rename from tests/test_litellm/vector_stores/test_main.py rename to tests/unit/vector_stores/test_main.py diff --git a/tests/test_litellm/vector_stores/test_vector_store_create_provider_logic.py b/tests/unit/vector_stores/test_vector_store_create_provider_logic.py similarity index 100% rename from tests/test_litellm/vector_stores/test_vector_store_create_provider_logic.py rename to tests/unit/vector_stores/test_vector_store_create_provider_logic.py diff --git a/tests/test_litellm/vector_stores/test_vector_store_registry.py b/tests/unit/vector_stores/test_vector_store_registry.py similarity index 100% rename from tests/test_litellm/vector_stores/test_vector_store_registry.py rename to tests/unit/vector_stores/test_vector_store_registry.py